0.1.1 暂存

This commit is contained in:
陈辅元
2026-04-14 10:28:22 +08:00
parent 4a0638ba2e
commit cc83fe963a
83 changed files with 4005 additions and 179 deletions
+4 -2
View File
@@ -7,7 +7,7 @@ DEEPSEEK_BASE_URL=https://api.deepseek.com
# 模型配置 # 模型配置
# 主生成模型 # 主生成模型
MODEL_PRIMARY=deepseek-reasoner MODEL_PRIMARY=deepseek-chat
TEMPERATURE=0 # 确定性模式:每次生成相同结果 TEMPERATURE=0 # 确定性模式:每次生成相同结果
MAX_TOKENS=4096 MAX_TOKENS=4096
# 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等 # 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等
@@ -41,7 +41,7 @@ ENABLE_SQL_VALIDATION=true
BLOCK_DANGEROUS_SQL=true BLOCK_DANGEROUS_SQL=true
# ========== 重试配置 ========== # ========== 重试配置 ==========
MAX_RETRY=1 # 确定性模式:只尝试一次,不重试 MAX_RETRY=2 # 确定性模式:只尝试一次,不重试
RETRY_DELAY=1.0 RETRY_DELAY=1.0
# ========== Few-shot 配置 ========== # ========== Few-shot 配置 ==========
@@ -57,3 +57,5 @@ FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl
# ========== 日志配置 ========== # ========== 日志配置 ==========
LOG_LEVEL=INFO LOG_LEVEL=INFO
LOG_FILE=./logs/text2sql.log LOG_FILE=./logs/text2sql.log
database_url = mssql+pymssql://sa:123456@192.168.3.201:1433/G3SB_PROD_HW
+109
View File
@@ -0,0 +1,109 @@
# Impact Analysis Report — NL2SQL 确定性(temperature)
## 1. 改动概览
- **背景与目标**:同一中文问题多次请求时,选表与生成 SQL 因 LLM 采样(默认 `temperature=0.3`)出现不一致。
- **涉及模块**:`backend/llm/deepseek_client.py`(选表、校验、便捷 `generate_sql`)、`backend/agents/orchestrator.py`(编排器内 SQL 生成)。
- **改动类型**:缺陷修复 / 行为稳定性增强(非功能扩展)。
## 2. 方法级改动
| 位置 | 原行为 | 新行为 |
|------|--------|--------|
| `DeepSeekClient.select_tables` | 使用配置默认 `temperature`(如 0.3) | `setdefault(temperature=0.0, top_p=1.0)` 后调用 `chat_with_json`;调用方仍可通过 `**kwargs` 覆盖 |
| `DeepSeekClient.validate_sql` | 同上 | 同上 |
| `DeepSeekClient.generate_sql` | 同上 | 同上 |
| `Text2SQLOrchestrator._generate_sql` | `chat(messages)` 使用默认温度 | 显式 `temperature=0.0, top_p=1.0` |
与原有逻辑差异:选表、校验、SQL 文本生成在相同 prompt 下更趋确定;若上游 API 在 `temperature=0` 下仍存在极小浮动,属于服务商实现范畴。
## 3. 调用方与影响范围
- **调用方**:`orchestrator._llm_select_tables` → `deepseek.select_tables`;`orchestrator._generate_sql` → `deepseek.chat`;校验链路 → `validate_sql`;脚本或其它代码若直接调用 `generate_sql`/`select_tables`/`validate_sql` 并传入自定义 `temperature`,**仍以调用方 kwargs 为准**(`setdefault` 不覆盖已传值)。
- **输入输出**:接口签名未变;返回 JSON/SQL 文本的**内容分布**更集中,极端情况下可能从「多种可接受 SQL」收敛到其中一种。
- **破坏性变更**:否(未改公开方法签名;仅默认采样参数与编排器单次 `chat` 参数)。
## 4. 风险与回滚
- **风险级别**:低。可能略微降低「多解探索」多样性;对 Text2SQL 通常可接受。
- **回滚**:恢复上述文件中的 `setdefault` / 显式 `temperature` 改动,或于环境变量/配置层为 DeepSeek 提高默认 `temperature`(若未来抽到配置项)。
- **回滚方式是否简单**:是(单提交回退即可)。
## 5. 验证与测试
- **建议**:对固定中文问题连续请求 5~10 次,比对日志中 `LLM精筛选中表` 与最终 SQL 是否一致。
- **说明**:Few-shot / 向量粗筛若存在非确定性实现,仍可能导致差异;本次改动仅消除 **LLM 采样** 主导的不一致。
## 6. 配置变更(temperature 项)
- 无新增环境变量或配置文件项;未改 `.env`。
---
# 增补:英文问句先译中文再走 Text2SQL(2026-04)
## 1. 改动概览
- **目标**:缓解「同一意图中英文提问选表/SQL 不一致」——在无 CJK 的英文问句进入粗筛、选表、few-shot、SQL 生成前,先经 DeepSeek 译为中文。
- **涉及模块**:`backend/utils/question_locale.py`(新建)、`backend/config/prompts.py`、`backend/llm/deepseek_client.py`、`backend/agents/orchestrator.py`、`backend/main.py`、`api_server.py`。
## 2. 方法级改动
| 符号 | 说明 |
|------|------|
| `looks_like_english_only` | 无中日韩/假名/韩文且拉丁字母≥3 时视为「可译英文」 |
| `DeepSeekClient.translate_nl_question_to_zh` | 一次 `chat(temperature=0)`,返回单行中文问句 |
| `Text2SQLOrchestrator.generate` | 开头若开启且命中启发式则译问句,后续全流程使用中文;`metadata` 含 `question_original`、`question_zh_normalized` |
| `_nl_dict_from_generation` | 若有译句,响应增加 `query_normalization: { original, zh }` |
## 3. 调用方与破坏性变更
- **调用方**:所有 `orchestrator.generate()`;API 层无签名变更。
- **破坏性变更**:否。默认开启;关闭后行为与旧版一致。英文请求多一次 LLM 调用(延迟与费用略增)。
## 4. 配置变更
| 变量 | 含义 | 默认 |
|------|------|------|
| `TRANSLATE_EN_TO_ZH` | `false`/`0`/`off` 关闭英译中 | 开启 |
CLI:`python backend/main.py --no-translate-en` 关闭。
## 5. 风险与回滚
- **风险**:翻译偏差导致表意偏移;专有名词若未保留可能被误译。
- **回滚**:`TRANSLATE_EN_TO_ZH=false` 或 `--no-translate-en`;或回退相关提交。
- **回滚方式是否简单**:是。
---
# 增补:库探针状态码 0 / 1 / -1 语义对齐(2026-04)
## 1. 改动概览
- **目标**:与业务约定一致——**1** 表示查询返回至少一行数据;**0** 表示执行成功但**行数为 0**(含「有列无行」的 SELECT);**-1** 表示执行失败并触发重新生成,且重试提示携带数据库错误摘要。
- **涉及模块**:`backend/db/dbhub_tools.py`、`backend/agents/orchestrator.py`、`backend/llm/deepseek_client.py`(注释)、`backend/main.py`(CLI 展示无行说明)、`backend/db/__init__.py`、`tools/dbhub_tools.py`、`tools/__init__.py`。
## 2. 方法级改动
| 符号 | 说明 |
|------|------|
| `probe_sql_execution_status_ex` | 新增;返回 `(status, error_message)`,`-1` 时第二项为异常字符串 |
| `probe_sql_execution_status` | 仍返回 `int \| None`,内部委托 `_ex`,**语义变更**:由「有列或行即 1」改为「**至少一行数据**为 1」 |
| `Text2SQLOrchestrator._validate_sql` | `-1` 错误信息附带库端详情;状态 `0` 时在 LLM 反馈外增加固定前缀,引导用户核对 SQL 并补充条件 |
| `single_query` | 验证通过且存在 `db_empty_feedback` 时打印「探针 0」说明块 |
## 3. 调用方与破坏性变更
- **调用方**:仅编排器直接调用探针;对外 API 仍通过 `GenerationResult.metadata` 的 `db_execution_status` / `db_empty_feedback` 表达结果。
- **破坏性变更**:是(**语义**)。此前「SELECT 返回 0 行但有列名」会被判为 **1**;现改为 **0**,会走无数据说明分支而非「直接可交付」。
## 4. 风险与回滚
- **风险**:低~中。更多查询会落入「无数据行」分支,多一次 `empty_result_user_feedback` LLM 调用;但更符合「无数据则请用户补充」的产品逻辑。
- **回滚**:恢复 `probe_sql_execution_status` 内基于「列或行非空」判定 1 的旧实现,并回退编排器与 CLI 改动。
- **回滚方式是否简单**:是。
## 5. 配置变更
- 无。
+46 -47
View File
@@ -95,22 +95,26 @@
``` ```
text2sql_agent_camel/ text2sql_agent_camel/
├── agents/ # Agent 层 ├── api_server.py # FastAPI NL 网关(根目录入口;启动前将 backend 加入 sys.path)
│ ├── orchestrator.py # 主编排器(协调 Schema Linker → Generator → Validator) ├── backend/ # Python 后端(除 api_server 外的编排与工具)
│ ├── schema_linker.py # 表筛选 Agent(粗筛 + LLM 精筛 + 外键扩展) │ ├── main.py # CLI 入口(交互模式:python backend/main.py)
│ ├── sql_generator.py # SQL 生成 Agent(集成 Few-Shot 注入) │ ├── nl_lite_store.py # 轻量会话/收藏内存存储(API 演示用)
│ └── validator.py # SQL 验证 Agent(程序 + LLM 双重验证) │ ├── agents/ # Agent 层
├── config/ │ │ ├── orchestrator.py # 主编排器(协调 Schema Linker → Generator → Validator)
│ ├── prompts.py # 系统与用户 Prompt 模板 │ │ ├── schema_linker.py # 表筛选 Agent(粗筛 + LLM 精筛 + 外键扩展)
│ └── settings.py # 配置类(基于 Pydantic) │ │ ├── sql_generator.py # SQL 生成 Agent(集成 Few-Shot 注入)
├── schema/ # Schema 管理层 │ │ └── validator.py # SQL 验证 Agent(程序 + LLM 双重验证)
│ ├── models.py # 数据模型(Table, Column, ForeignKey, DatabaseSchema) │ ├── config/
│ ├── loader.py # Schema 加载器(JSON / G3SB 格式解析) │ │ ├── prompts.py # 系统与用户 Prompt 模板
│ ├── manager.py # Schema 管理器(查询、过滤、转字符串) │ │ └── settings.py # 配置类(基于 Pydantic)
│ └── indexer.py # 向量索引构建器(ChromaDB 封装) │ ├── schema/ # Schema 管理层
├── llm/ │ │ ├── models.py # 数据模型(Table, Column, ForeignKey, DatabaseSchema)
│ └── deepseek_client.py # DeepSeek API 客户端(chat / validate_sql / select_tables) │ │ ├── loader.py # Schema 加载器(JSON / G3SB 格式解析)
├── utils/ │ │ ├── manager.py # Schema 管理器(查询、过滤、转字符串)
│ │ └── indexer.py # 向量索引构建器(ChromaDB 封装)
│ ├── llm/
│ │ └── deepseek_client.py # DeepSeek API 客户端(chat / validate_sql / select_tables)
│ └── utils/
│ ├── embedding.py # Qwen3-Embedding 封装(本地 / 远程 API) │ ├── embedding.py # Qwen3-Embedding 封装(本地 / 远程 API)
│ ├── sql_parser.py # sqlglot 工具(语法验证、方言转换、规范化) │ ├── sql_parser.py # sqlglot 工具(语法验证、方言转换、规范化)
│ ├── validators.py # 验证逻辑(危险操作检测、完整验证流水线) │ ├── validators.py # 验证逻辑(危险操作检测、完整验证流水线)
@@ -134,7 +138,6 @@ text2sql_agent_camel/
├── logs/ ├── logs/
│ └── text2sql.log # 运行日志(默认) │ └── text2sql.log # 运行日志(默认)
├── .env # 环境变量配置(API Key、路径、参数) ├── .env # 环境变量配置(API Key、路径、参数)
├── main.py # CLI 入口(单次查询 / 批量 / 交互模式)
├── pyproject.toml # 项目元数据与依赖 ├── pyproject.toml # 项目元数据与依赖
├── README.md # 本文档 ├── README.md # 本文档
└── SQL_GENERATION_LOGIC.md # SQL 生成逻辑详解(流程图、示例) └── SQL_GENERATION_LOGIC.md # SQL 生成逻辑详解(流程图、示例)
@@ -230,11 +233,14 @@ Schema JSON 格式示例:
### 4. 构建向量索引(首次运行) ### 4. 构建向量索引(首次运行)
```bash ```bash
# 自动构建:首次查询时会自动创建 # 自动构建:首次在交互模式里提问时会自动创建
python main.py "测试查询" python backend/main.py
# 或手动预构建(推荐,加速首次查询) # 或手动预构建(推荐,加速首次查询;请在仓库根目录执行)
python -c " python -c "
import sys
from pathlib import Path
sys.path.insert(0, str(Path.cwd() / 'backend'))
from agents.orchestrator import Text2SQLOrchestrator from agents.orchestrator import Text2SQLOrchestrator
from schema.manager import SchemaManager from schema.manager import SchemaManager
@@ -247,30 +253,13 @@ orch.build_vector_index(force_rebuild=True)
### 5. 运行演示 ### 5. 运行演示
```bash ```bash
# 单次查询(默认 T-SQL / SQL Server) # 启动后进入交互式问答(默认 T-SQL / SQL Server)
python main.py "查询2024年1月的销售额" python backend/main.py
# 交互模式 # 可选:启动时附带参数(均在进入交互前生效),例如指定 Schema、方言、Few-Shot、日志等
python main.py --interactive python backend/main.py --schema ./data/schemas/custom_schema.json --dialect postgresql --verbose
python backend/main.py --no-fewshot
# 批量查询(每行一个问题) python backend/main.py --fewshot-top-k 5 --fewshot-min-rating 8
python main.py --batch queries.txt
# 指定 Schema 和方言
python main.py "统计每个市场的未结算交易量" \
--schema ./data/schemas/custom_schema.json \
--dialect postgresql
# 禁用 Few-Shot(对比实验)
python main.py "查询活跃账户数" --no-fewshot
# 调整 Few-Shot 参数
python main.py "查询2024年1月的销售额" \
--fewshot-top-k 5 \
--fewshot-min-rating 8
# 详细日志(调试用)
python main.py "查询持仓余额" --verbose
``` ```
## 🎯 使用示例 ## 🎯 使用示例
@@ -279,6 +268,12 @@ python main.py "查询持仓余额" --verbose
```python ```python
import os import os
import sys
from pathlib import Path
# 在仓库根目录运行时,将 backend 加入模块搜索路径(请先 cd 到项目根)
sys.path.insert(0, str(Path.cwd() / "backend"))
from agents.orchestrator import Text2SQLOrchestrator from agents.orchestrator import Text2SQLOrchestrator
from schema.manager import SchemaManager from schema.manager import SchemaManager
from llm.deepseek_client import DeepSeekConfig from llm.deepseek_client import DeepSeekConfig
@@ -376,7 +371,7 @@ VECTOR_DB_PATH=./data/embeddings/chroma
### CLI 参数速查 ### CLI 参数速查
```bash ```bash
python main.py [问题] [选项] python backend/main.py [选项]
# 核心选项 # 核心选项
--schema PATH Schema 文件路径 --schema PATH Schema 文件路径
@@ -398,8 +393,6 @@ python main.py [问题] [选项]
--embedding-model PATH Embedding 模型路径 --embedding-model PATH Embedding 模型路径
# 其他 # 其他
--interactive, -i 交互模式
--batch FILE 批量文件路径
--verbose, -v 详细日志 --verbose, -v 详细日志
``` ```
@@ -564,8 +557,11 @@ WHERE i.AccountID = 'ACC001'
# 检查索引文件 # 检查索引文件
ls -la data/embeddings/chroma/ ls -la data/embeddings/chroma/
# 强制重建索引 # 强制重建索引(在仓库根目录执行)
python -c " python -c "
import sys
from pathlib import Path
sys.path.insert(0, str(Path.cwd() / 'backend'))
from agents.orchestrator import Text2SQLOrchestrator from agents.orchestrator import Text2SQLOrchestrator
from schema.manager import SchemaManager from schema.manager import SchemaManager
@@ -587,6 +583,9 @@ orch.build_vector_index(force_rebuild=True)
```bash ```bash
# 检查数据文件 # 检查数据文件
python -c " python -c "
import sys
from pathlib import Path
sys.path.insert(0, str(Path.cwd() / 'backend'))
from utils.fewshot_selector import FewShotSelector from utils.fewshot_selector import FewShotSelector
s = FewShotSelector('./data/experiences/all_samples.jsonl') s = FewShotSelector('./data/experiences/all_samples.jsonl')
print(f'样本数: {len(s.samples)}') print(f'样本数: {len(s.samples)}')
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+581
View File
@@ -0,0 +1,581 @@
#!/usr/bin/env python3
"""
FastAPI 服务入口 - Text2SQL NL Chat API
包装 Text2SQLOrchestrator 为 REST API 服务,供前端调用
"""
from __future__ import annotations
import os
import sys
import json
import html as html_lib
import asyncio
import logging
from pathlib import Path
from typing import Optional, List, Dict, Any, AsyncIterator
from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException, Body, Path as FPath
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse, JSONResponse
from pydantic import BaseModel, Field
from dotenv import load_dotenv
_REPO_DIR = Path(__file__).resolve().parent
_BACKEND_DIR = _REPO_DIR / "backend"
if str(_BACKEND_DIR) not in sys.path:
sys.path.insert(0, str(_BACKEND_DIR))
from main import setup_environment, load_schema, create_orchestrator, resolve_sql_dialect
from agents.orchestrator import GenerationResult
from nl_lite_store import lite_nl_store
from utils.dialog_classifier import DialogIntent, classify_dialog
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
datefmt='%Y-%m-%d %H:%M:%S'
)
logger = logging.getLogger(__name__)
load_dotenv(Path(__file__).resolve().parent / ".env")
orchestrator = None
schema_manager = None
def get_orchestrator():
"""获取或初始化 orchestrator"""
global orchestrator, schema_manager
if orchestrator is not None:
return orchestrator
if not setup_environment():
raise RuntimeError("环境检查失败,请检查配置")
schema_path = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
schema_meta_path = os.getenv("SCHEMA_META_PATH", None)
schema_manager = load_schema(schema_path, schema_meta_path)
class Args:
api_key = os.getenv("DEEPSEEK_API_KEY")
model = os.getenv("MODEL_PRIMARY", "deepseek-chat")
temperature = float(os.getenv("TEMPERATURE", "0.3"))
max_tokens = int(os.getenv("MAX_TOKENS", "4096"))
max_retry = int(os.getenv("MAX_RETRY", "2"))
embedding_model = None
vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma")
no_vector_search = False
no_fewshot = not os.getenv("FEWSHOT_ENABLED", "true").lower() in ("true", "1", "yes")
fewshot_top_k = int(os.getenv("FEWSHOT_TOP_K", "3"))
fewshot_min_rating = int(os.getenv("FEWSHOT_MIN_RATING", "7"))
no_translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() in (
"0",
"false",
"no",
"off",
)
orchestrator = create_orchestrator(schema_manager, Args())
logger.info("[OK] Orchestrator 初始化完成")
return orchestrator
class NLChatRequest(BaseModel):
message: str = Field(..., description="用户输入的自然语言问题")
service_code: Optional[str] = Field(None, description="模型路由: openai | deepseek")
lang_code: Optional[str] = Field("auto", description="语言: zh | en | tc | auto")
taskId: Optional[str] = Field(None, description="任务ID")
session_id: Optional[str] = Field(None, description="会话ID")
visitor_biz_id: Optional[str] = Field(None, description="访客业务ID")
user_id: Optional[str] = Field(None, description="用户ID")
streaming_throttle: Optional[int] = Field(None, description="流式节流参数")
class IntentPayload(BaseModel):
intent: str
confidence: Optional[float] = None
reason: Optional[str] = None
class DataQueryResult(BaseModel):
sql: str = Field(..., description="生成的SQL语句")
columns: List[str] = Field(default_factory=list, description="查询结果列名")
rows: List[Dict[str, Any]] = Field(default_factory=list, description="查询结果行数据")
row_count: int = Field(0, description="结果行数")
truncated: bool = Field(False, description="是否截断")
sql_explain: Optional[str] = Field(None, description="SQL自然语言说明")
can_export: Optional[bool] = Field(True, description="是否允许导出")
class NLChatSuccessData(BaseModel):
intent: IntentPayload
branch_result: DataQueryResult
stream_narrative: Optional[str] = Field(None, description="流式叙述(SQL块上方说明)")
stream_narrative_after: Optional[str] = Field(None, description="流式叙述(SQL块下方说明)")
class ApiEnvelope(BaseModel):
code: int = Field(200, description="状态码,200表示成功")
msg: str = Field("", description="消息")
data: Optional[NLChatSuccessData] = Field(None, description="响应数据")
class ErrorResponse(BaseModel):
code: int
msg: str
data: Optional[Any] = None
class SessionCreateBody(BaseModel):
title: Optional[str] = None
user_id: Optional[str] = None
visitor_biz_id: Optional[str] = None
class SessionTitlePatchBody(BaseModel):
title: Optional[str] = None
user_id: Optional[str] = None
visitor_biz_id: Optional[str] = None
class SessionMessagePatchBody(BaseModel):
content: str
user_id: Optional[str] = None
visitor_biz_id: Optional[str] = None
class SqlExecuteBody(BaseModel):
sql: str
max_rows: Optional[int] = None
user_id: Optional[str] = None
visitor_biz_id: Optional[str] = None
chat_session_id: Optional[str] = None
chat_message_id: Optional[int] = None
class FavoriteCreateBody(BaseModel):
fav_type: str = Field(..., description="sql | function | report")
name: str = ""
desc: Optional[str] = None
user_id: Optional[str] = None
visitor_biz_id: Optional[str] = None
sql: Optional[str] = None
sql_explain: Optional[str] = None
path: Optional[str] = None
reportPath: Optional[str] = None
params: Optional[str] = None
def _sse_data(obj: Dict[str, Any]) -> bytes:
return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n".encode("utf-8")
def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
"""寒暄/元问题等:与前端 BUSINESS_MANUAL + branch_result.answer 一致。"""
r = (reply or "").strip() or "您好。"
return {
"intent": {"intent": "BUSINESS_MANUAL", "confidence": 1.0, "reason": "conversation"},
"branch_result": {"answer": r},
}
async def _append_session_if_needed(request: NLChatRequest, user_text: str, data: Dict[str, Any]) -> None:
if not (request.session_id and request.session_id.strip()):
return
try:
assistant_json = json.dumps(data, ensure_ascii=False)
await lite_nl_store.append_exchange(
request.user_id,
request.visitor_biz_id,
request.session_id.strip(),
user_text,
assistant_json,
)
except Exception as e:
logger.warning(f"[API] 会话落库跳过: {e}")
def _nl_dict_from_generation(result: GenerationResult) -> Dict[str, Any]:
sql = (result.sql or "").strip()
explain_parts: List[str] = []
if result.metadata.get("db_empty_feedback"):
explain_parts.append(str(result.metadata["db_empty_feedback"]))
if result.warnings:
explain_parts.extend(str(w) for w in result.warnings)
if result.errors:
explain_parts.extend(str(e) for e in result.errors)
explain = "; ".join(explain_parts) if explain_parts else None
branch_result: Dict[str, Any] = {
"sql": sql,
"columns": [],
"rows": [],
"row_count": 0,
"truncated": False,
}
if explain:
branch_result["sql_explain"] = explain
dbs = result.metadata.get("db_execution_status")
if dbs is not None:
branch_result["db_execution_status"] = dbs
dbe = result.metadata.get("db_empty_feedback")
if dbe:
branch_result["db_empty_feedback"] = dbe
conf = 1.0 if result.valid else 0.0
if result.valid:
reason = f"使用了 {len(result.tables_used)} 张表"
else:
reason = result.errors[0] if result.errors else "SQL生成未通过验证"
payload: Dict[str, Any] = {
"intent": {"intent": "DATA_QUERY", "confidence": conf, "reason": reason},
"branch_result": branch_result,
}
qo = result.metadata.get("question_original")
qz = result.metadata.get("question_zh_normalized")
if qo and qz:
payload["query_normalization"] = {"original": qo, "zh": qz}
return payload
def _sql_gen_stream_html(result: GenerationResult) -> str:
sql = (result.sql or "").strip()
inner = json.dumps({"sql": sql}, ensure_ascii=False)
parts = [f"<data>{inner}</data>"]
explain = ""
if result.metadata.get("db_empty_feedback"):
explain = str(result.metadata["db_empty_feedback"])
elif result.warnings:
explain = str(result.warnings[0])
elif result.errors:
explain = "; ".join(str(e) for e in result.errors[:5])
if explain.strip():
parts.append(f'<div class="sql-explain-body">{html_lib.escape(explain)}</div>')
return "".join(parts)
async def _run_generate(question: str, top_k: int = 20) -> GenerationResult:
orch = get_orchestrator()
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
def _call() -> GenerationResult:
return orch.generate(
question=question.strip(),
dialect=dialect,
top_k_candidates=top_k,
)
return await asyncio.to_thread(_call)
async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
if not request.message or not request.message.strip():
yield _sse_data({"code": 400, "msg": "消息内容不能为空", "data": None})
return
text = request.message.strip()
classified = classify_dialog(text)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
logger.info("[API/stream] 对话意图: conversation(跳过 Text2SQL,与 CLI single_query 一致)")
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "CHAT"})
yield _sse_data({"stage": "chat", "stream_kind": "content", "content": reply})
data_dict = _conversation_nl_dict(reply)
yield _sse_data({"code": 200, "msg": "success", "data": data_dict})
await _append_session_if_needed(request, text, data_dict)
return
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "DATA_QUERY"})
try:
result = await _run_generate(text)
except Exception as e:
logger.error(f"[API/stream] 生成异常: {e}")
yield _sse_data({"code": 500, "msg": str(e), "data": None})
return
yield _sse_data({"stage": "sql_gen", "stream_kind": "content", "content": _sql_gen_stream_html(result)})
data_dict = _nl_dict_from_generation(result)
msg = "success" if result.valid else "partial"
yield _sse_data({"code": 200, "msg": msg, "data": data_dict})
await _append_session_if_needed(request, text, data_dict)
@asynccontextmanager
async def lifespan(app: FastAPI):
logger.info("=" * 60)
logger.info("Text2SQL API Server 启动中...")
logger.info("=" * 60)
try:
get_orchestrator()
logger.info("[OK] 服务已就绪")
except Exception as e:
logger.error(f"服务初始化失败: {e}")
raise
yield
logger.info("Text2SQL API Server 已关闭")
app = FastAPI(
title="Text2SQL NL Chat API",
description="自然语言转SQL的对话接口",
version="1.0.0",
lifespan=lifespan
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/health")
async def health_check():
"""健康检查接口"""
return {"status": "ok", "message": "Text2SQL API is running"}
@app.post("/g3sb/api/nl/chat")
async def nl_chat(request: NLChatRequest):
"""
自然语言对话接口(非流式),响应结构与流式结束包一致。
寒暄/致谢等先经 classify_dialog,与 CLI 一致。
"""
try:
if not request.message or not request.message.strip():
raise HTTPException(status_code=400, detail="消息内容不能为空")
text = request.message.strip()
logger.info(f"[API] 问题: {text[:100]}...")
classified = classify_dialog(text)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
data_dict = _conversation_nl_dict(reply)
await _append_session_if_needed(request, text, data_dict)
return {"code": 200, "msg": "success", "data": data_dict}
result = await _run_generate(text)
data_dict = _nl_dict_from_generation(result)
response_data = NLChatSuccessData.model_validate(data_dict)
msg = "success" if result.valid else "partial"
if result.valid:
logger.info(f"[API] 生成成功: {(result.sql or '')[:80]}...")
else:
logger.warning(f"[API] 生成失败: {result.errors}")
await _append_session_if_needed(request, text, data_dict)
return ApiEnvelope(code=200, msg=msg, data=response_data)
except HTTPException:
raise
except Exception as e:
logger.error(f"[API] 异常: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.post("/g3sb/api/nl/chat/stream")
async def nl_chat_stream(request: NLChatRequest):
"""
自然语言对话流式接口(SSE)。
事件体为 JSON:分片 delta `{stage, stream_kind, content}` 或结束包 `{code, msg, data}`。
"""
return StreamingResponse(
_chat_stream_events(request),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
},
)
@app.post("/g3sb/api/nl/sessions")
async def nl_create_session(body: SessionCreateBody):
row = await lite_nl_store.create_session(body.user_id, body.visitor_biz_id, body.title)
return {"code": 200, "msg": "success", "data": row}
@app.get("/g3sb/api/nl/sessions")
async def nl_list_sessions(
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
limit: int = 100,
offset: int = 0,
):
data = await lite_nl_store.list_sessions(user_id, visitor_biz_id, limit, offset)
return {"code": 200, "msg": "success", "data": data}
@app.get("/g3sb/api/nl/sessions/{session_id}/messages")
async def nl_get_messages(
session_id: str = FPath(..., description="会话ID"),
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
limit: int = 500,
offset: int = 0,
):
data = await lite_nl_store.get_messages(user_id, visitor_biz_id, session_id, limit, offset)
if data is None:
return JSONResponse(status_code=200, content={"code": 404, "msg": "session not found", "data": None})
return {"code": 200, "msg": "success", "data": data}
@app.patch("/g3sb/api/nl/sessions/{session_id}/title")
async def nl_patch_session_title(session_id: str, body: SessionTitlePatchBody):
row = await lite_nl_store.update_session_title(
body.user_id, body.visitor_biz_id, session_id, body.title
)
if not row:
return JSONResponse(status_code=200, content={"code": 404, "msg": "session not found", "data": None})
return {"code": 200, "msg": "success", "data": row}
@app.patch("/g3sb/api/nl/sessions/{session_id}/messages/{message_id}")
async def nl_patch_session_message(
session_id: str,
message_id: int,
body: SessionMessagePatchBody,
):
row = await lite_nl_store.patch_message(
body.user_id, body.visitor_biz_id, session_id, message_id, body.content
)
if not row:
return JSONResponse(status_code=200, content={"code": 404, "msg": "message not found", "data": None})
return {"code": 200, "msg": "success", "data": row}
@app.delete("/g3sb/api/nl/sessions/{session_id}")
async def nl_delete_session(
session_id: str,
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
):
row = await lite_nl_store.delete_session(user_id, visitor_biz_id, session_id)
if not row:
return JSONResponse(status_code=200, content={"code": 404, "msg": "session not found", "data": None})
return {"code": 200, "msg": "success", "data": row}
@app.post("/g3sb/api/nl/sql/execute")
async def nl_sql_execute(_body: SqlExecuteBody):
return {
"code": 501,
"msg": "本机 Text2SQL 演示服务未接数据库,无法执行 SQL。请配置独立 NL 网关或移除「运行 SQL」依赖。",
"data": None,
}
@app.post("/g3sb/api/nl/operation-logs")
async def nl_post_operation_log(_body: Dict[str, Any] = Body(...)):
return {"code": 200, "msg": "success", "data": None}
@app.get("/g3sb/api/nl/operation-logs")
async def nl_list_operation_logs(
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
limit: int = 50,
offset: int = 0,
q: Optional[str] = None,
op_type: Optional[str] = None,
result: Optional[str] = None,
date_from: Optional[str] = None,
date_to: Optional[str] = None,
):
_ = (user_id, visitor_biz_id, q, op_type, result, date_from, date_to)
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}
@app.get("/g3sb/api/nl/favorites")
async def nl_list_favorites(
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
):
grouped = await lite_nl_store.get_favorites_grouped(user_id, visitor_biz_id)
return {"code": 200, "msg": "success", "data": grouped}
@app.post("/g3sb/api/nl/favorites")
async def nl_create_favorite(body: FavoriteCreateBody):
row = await lite_nl_store.add_favorite(body.user_id, body.visitor_biz_id, body.model_dump(exclude_none=True))
return {"code": 200, "msg": "success", "data": row}
@app.patch("/g3sb/api/nl/favorites/{fav_id}")
async def nl_patch_favorite(
fav_id: str,
body: Dict[str, Any] = Body(default_factory=dict),
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
):
patch = {k: v for k, v in body.items() if k not in ("user_id", "visitor_biz_id")}
ok = await lite_nl_store.patch_favorite(user_id, visitor_biz_id, fav_id, patch)
if not ok:
return JSONResponse(status_code=200, content={"code": 404, "msg": "favorite not found", "data": None})
return {"code": 200, "msg": "success", "data": {}}
@app.delete("/g3sb/api/nl/favorites/{fav_id}")
async def nl_delete_favorite(
fav_id: str,
user_id: Optional[str] = None,
visitor_biz_id: Optional[str] = None,
):
ok = await lite_nl_store.delete_favorite(user_id, visitor_biz_id, fav_id)
if not ok:
return JSONResponse(status_code=200, content={"code": 404, "msg": "favorite not found", "data": None})
return {"code": 200, "msg": "success", "data": {"ok": True}}
@app.get("/g3sb/api/nl/knowledge-docs/uploads")
async def nl_list_knowledge_uploads(
doc_type: Optional[str] = None,
limit: int = 50,
offset: int = 0,
):
_ = doc_type
return {"code": 200, "msg": "success", "data": {"items": [], "limit": limit, "offset": offset}}
@app.post("/g3sb/api/nl/knowledge-docs/upload")
async def nl_upload_knowledge_doc():
return JSONResponse(
status_code=200,
content={"code": 501, "msg": "knowledge upload not supported on lite API", "data": None},
)
@app.get("/g3sb/api/nl/admin/visibility")
async def admin_visibility():
"""管理端可见性接口"""
return {
"code": 200,
"msg": "success",
"data": {
"admin_visible": True
}
}
if __name__ == "__main__":
import uvicorn
port = int(os.getenv("API_PORT", "8041"))
host = os.getenv("API_HOST", "0.0.0.0")
logger.info(f"启动服务: http://{host}:{port}")
logger.info(f"API文档: http://{host}:{port}/docs")
uvicorn.run(
"api_server:app",
host=host,
port=port,
reload=False,
log_level="info"
)
View File
Binary file not shown.
Binary file not shown.
@@ -34,10 +34,11 @@ class Text2SQLOrchestrator:
Text2SQL 多智能体编排器 Text2SQL 多智能体编排器
工作流程: 工作流程:
1. Schema Linker:粗筛 + LLM精筛,选出相关表 1. 粗筛候选表
2. 外键扩展:自动包含关联表 2. Schema Linker:LLM 精筛表
3. SQL Generator:生成SQL 3. 外键扩展 → 拼 Schema 子集
4. Validator:验证SQL,不通过则重试(最多max_retry次) 4. SQL Generator:生成 SQL
5. Validator:验证 SQL
""" """
def __init__( def __init__(
@@ -54,6 +55,7 @@ class Text2SQLOrchestrator:
fewshot_samples_path: Optional[str] = None, fewshot_samples_path: Optional[str] = None,
fewshot_top_k: int = 3, fewshot_top_k: int = 3,
fewshot_min_rating: int = 7, fewshot_min_rating: int = 7,
translate_english_to_zh: bool = True,
): ):
""" """
初始化编排器 初始化编排器
@@ -66,10 +68,12 @@ class Text2SQLOrchestrator:
vector_db_path: 向量数据库路径 vector_db_path: 向量数据库路径
max_retry: 最大重试次数(包含首次生成) max_retry: 最大重试次数(包含首次生成)
use_vector_search: 是否使用向量检索粗筛 use_vector_search: 是否使用向量检索粗筛
translate_english_to_zh: 无中日韩字符的英文问句是否先译为中文再走检索与生成
""" """
self.schema_manager = schema_manager self.schema_manager = schema_manager
self.max_retry = max_retry self.max_retry = max_retry
self.use_vector_search = use_vector_search self.use_vector_search = use_vector_search
self.translate_english_to_zh = translate_english_to_zh
# 初始化DeepSeek客户端 # 初始化DeepSeek客户端
if deepseek_config: if deepseek_config:
@@ -112,6 +116,7 @@ class Text2SQLOrchestrator:
f"[OK] Text2SQLOrchestrator初始化完成: " f"[OK] Text2SQLOrchestrator初始化完成: "
f"max_retry={max_retry}, use_vector_search={use_vector_search}" f"max_retry={max_retry}, use_vector_search={use_vector_search}"
+ (f", fewshot=on" if self.fewshot_enabled else "") + (f", fewshot=on" if self.fewshot_enabled else "")
+ (", en→zh=on" if self.translate_english_to_zh else ", en→zh=off")
) )
def _get_vector_index(self) -> SchemaIndexer: def _get_vector_index(self) -> SchemaIndexer:
@@ -167,7 +172,7 @@ class Text2SQLOrchestrator:
self, self,
question: str, question: str,
candidate_tables: List[str], candidate_tables: List[str],
max_tables: int = 5 max_tables: int = 5,
) -> Tuple[List[str], str]: ) -> Tuple[List[str], str]:
""" """
阶段2:LLM精筛(Schema Linker Agent) 阶段2:LLM精筛(Schema Linker Agent)
@@ -193,7 +198,7 @@ class Text2SQLOrchestrator:
# 使用DeepSeek客户端调用(而非CAMEL Agent,更直接可控) # 使用DeepSeek客户端调用(而非CAMEL Agent,更直接可控)
response = self.deepseek.select_tables( response = self.deepseek.select_tables(
question=question, question=question,
table_list=table_list_str table_list=table_list_str,
) )
relevant_tables = response.get("relevant_tables", []) relevant_tables = response.get("relevant_tables", [])
@@ -284,7 +289,8 @@ class Text2SQLOrchestrator:
self, self,
question: str, question: str,
schema_str: str, schema_str: str,
dialect: str = "tsql" dialect: str = "tsql",
validation_feedback: Optional[str] = None,
) -> str: ) -> str:
""" """
SQL生成(SQL Generator Agent) SQL生成(SQL Generator Agent)
@@ -293,6 +299,7 @@ class Text2SQLOrchestrator:
question: 用户问题 question: 用户问题
schema_str: Schema描述字符串 schema_str: Schema描述字符串
dialect: SQL方言 dialect: SQL方言
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
Returns: Returns:
SQL语句 SQL语句
@@ -336,6 +343,15 @@ class Text2SQLOrchestrator:
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。" "**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。" "条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。" "排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
"\n**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。"
)
if validation_feedback:
user_content += (
"\n\n【上次校验未通过】请根据下列错误修正 SQL,并输出完整可执行查询;"
"表名、列名必须与「当前Schema」中完全一致,不要臆造字段名。\n"
f"{validation_feedback}"
) )
messages = [ messages = [
@@ -343,7 +359,8 @@ class Text2SQLOrchestrator:
{"role": "user", "content": user_content}, {"role": "user", "content": user_content},
] ]
response = self.deepseek.chat(messages) # 与选表一致:生成阶段默认贪心解码,减少同一中文问题多次 SQL 不一致
response = self.deepseek.chat(messages, temperature=0.0, top_p=1.0)
sql = response.content.strip() sql = response.content.strip()
# 清理可能的markdown代码块 # 清理可能的markdown代码块
@@ -362,20 +379,25 @@ class Text2SQLOrchestrator:
sql: str, sql: str,
schema_str: str, schema_str: str,
dialect: str = "tsql", dialect: str = "tsql",
) -> Tuple[bool, List[str], List[str]]: question: str = "",
) -> Tuple[bool, List[str], List[str], Optional[int], Optional[str]]:
""" """
SQL验证(Validator Agent + 程序验证) SQL验证(程序验证 + 库探针 + 按探针分支的 LLM)
Args: Args:
sql: SQL语句 sql: SQL语句
schema_str: Schema描述 schema_str: Schema描述
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql) dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
question: 用户自然语言(探针为 0 时用于生成补充说明)
Returns: Returns:
(是否通过, 错误列表, 警告列表) (是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
探针:``1`` 至少一行数据;``0`` 执行成功但行数为 0;``-1`` 执行失败;``None`` 未配置库或未跑探针
""" """
errors = [] errors = []
warnings = [] warnings = []
db_execution_status: Optional[int] = None
empty_feedback: Optional[str] = None
# === 阶段1:程序验证(确定性规则) === # === 阶段1:程序验证(确定性规则) ===
from utils.sql_parser import validate_sql_syntax, validate_schema_consistency from utils.sql_parser import validate_sql_syntax, validate_schema_consistency
@@ -393,12 +415,43 @@ class Text2SQLOrchestrator:
errors.extend(schema_errors) errors.extend(schema_errors)
# 危险操作检查 # 危险操作检查
from utils.validators import check_dangerous_operations from utils.validators import check_dangerous_operations, check_no_cjk_in_sql_string_literals
danger_ok, danger_errors = check_dangerous_operations(sql) danger_ok, danger_errors = check_dangerous_operations(sql)
if not danger_ok: if not danger_ok:
errors.extend(danger_errors) errors.extend(danger_errors)
# === 阶段2:LLM语义验证 === # T-SQL:禁止中文等业务词出现在字符串字面量(如 FeeNatureID = '过户费')
if dialect == "tsql":
cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql)
if not cjk_ok:
errors.extend(cjk_errors)
# === 阶段1.5:数据库试执行(仅程序校验全部通过时;需配置 database_url) ===
if len(errors) == 0:
from db.dbhub_tools import probe_sql_execution_status_ex
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
if db_execution_status == -1:
msg = (
"数据库执行验证失败:SQL 在目标库执行报错(探针状态 -1),"
"将据此重新生成 SQL。"
)
if db_probe_err:
msg += f" 数据库返回:{db_probe_err}"
errors.append(msg)
elif db_execution_status is None:
warnings.append(
"未配置 database_url,已跳过数据库执行探针"
)
# 探针 1:库上至少有一行数据,跳过 Validator LLM,直接将 SQL 视为可交付
# 探针 0:执行成功但行数为 0,跳过 Validator LLM,另调 LLM 生成说明并引导用户补充条件
# 探针 -1:执行失败,跳过 Validator LLM,走重试
# 探针 None:走完整 Validator LLM
skip_validator_llm = db_execution_status in (-1, 0, 1)
# === 阶段2:LLM 语义验证(仅未命中库探针 0/1/-1 时) ===
if not skip_validator_llm:
try: try:
llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str) llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str)
@@ -426,8 +479,28 @@ class Text2SQLOrchestrator:
except Exception as e: except Exception as e:
logger.warning(f"LLM验证失败(降级为仅程序验证): {e}") logger.warning(f"LLM验证失败(降级为仅程序验证): {e}")
# === 阶段2b:探针 0 时生成用户可读补充说明(仍返回 SQL,由 API/CLI 一并展示) ===
if db_execution_status == 0 and len(errors) == 0:
prefix = (
"该 SQL 已在数据库成功执行,但返回的数据行数为 0(未查到匹配记录)。"
"请将下方 SQL 与说明一并核对;若不符合预期,请补充或调整条件后再次提问。"
)
try:
llm_fb = self.deepseek.empty_result_user_feedback(
question=question,
sql=sql,
schema=schema_str,
)
empty_feedback = f"{prefix}\n\n【分析与建议】\n{llm_fb}"
except Exception as e:
logger.warning(f"无数据说明生成失败: {e}")
empty_feedback = (
f"{prefix}\n\n【分析与建议】\n"
"未能自动生成详细分析。请补充时间范围、筛选条件或业务对象后重新提问。"
)
is_valid = len(errors) == 0 is_valid = len(errors) == 0
return is_valid, errors, warnings return is_valid, errors, warnings, db_execution_status, empty_feedback
def generate( def generate(
self, self,
@@ -448,11 +521,34 @@ class Text2SQLOrchestrator:
Returns: Returns:
GenerationResult对象 GenerationResult对象
""" """
from utils.question_locale import looks_like_english_only
original_question = (question or "").strip()
translation_meta: Dict = {}
work_question = original_question
if self.translate_english_to_zh and looks_like_english_only(original_question):
try:
zh = self.deepseek.translate_nl_question_to_zh(original_question).strip()
if zh and len(zh) >= 2:
work_question = zh
translation_meta["question_original"] = original_question
translation_meta["question_zh_normalized"] = zh
logger.info(
"[GEN] 英文已译为中文:%s",
zh[:120] + ("…" if len(zh) > 120 else ""),
)
else:
logger.warning("[GEN] 英译中结果为空或过短,使用原文")
except Exception as e:
logger.warning("[GEN] 英译中失败,使用原文: %s", e)
question = work_question
logger.info(f"[GEN] 开始生成SQL:{question[:50]}...") logger.info(f"[GEN] 开始生成SQL:{question[:50]}...")
attempt = 0 attempt = 0
last_sql = None last_sql = None
last_errors = [] last_errors = []
last_db_execution_status: Optional[int] = None
filtered_schema_str = "" filtered_schema_str = ""
tables_used = [] tables_used = []
@@ -465,7 +561,10 @@ class Text2SQLOrchestrator:
candidate_tables = self._coarse_filter(question, top_k=top_k_candidates) candidate_tables = self._coarse_filter(question, top_k=top_k_candidates)
# 1.2 LLM精筛 # 1.2 LLM精筛
relevant_tables, reasoning = self._llm_select_tables(question, candidate_tables) relevant_tables, reasoning = self._llm_select_tables(
question,
candidate_tables,
)
relevant_tables = self._prioritize_broker_tables(question, relevant_tables) relevant_tables = self._prioritize_broker_tables(question, relevant_tables)
# 1.3 外键扩展 # 1.3 外键扩展
@@ -480,12 +579,41 @@ class Text2SQLOrchestrator:
) )
logger.info(f" 选中表:{relevant_tables},扩展后:{expanded_tables}") logger.info(f" 选中表:{relevant_tables},扩展后:{expanded_tables}")
else: else:
# 重试时复用之前的Schema # 重试:在子 Schema 中并入「上次失败 SQL」实际引用到的表,并对齐程序校验与生成上下文
logger.info(f" 重试使用之前的Schema({len(tables_used)}张表)") logger.info(f" 重试使用之前的Schema({len(tables_used)}张表)")
if last_sql:
from utils.sql_parser import extract_tables_from_sql
extra = [
t
for t in extract_tables_from_sql(last_sql, dialect=dialect)
if self.schema_manager.get_table(t)
]
merged = list(dict.fromkeys([*(tables_used or []), *extra]))
tables_used = self._expand_relations(merged)
filtered_schema_str = self.schema_manager.to_compact_string(
table_names=tables_used,
include_columns=True,
max_columns_per_table=20,
)
if extra:
logger.info(
" 重试:合并失败SQL中的表 %s,外键扩展后:%s",
extra,
tables_used,
)
# === Step 2: SQL生成 === # === Step 2: SQL生成 ===
try: try:
sql = self._generate_sql(question, filtered_schema_str, dialect) feedback: Optional[str] = None
if attempt > 0 and last_errors:
feedback = "\n".join(f"- {e}" for e in last_errors[:20])
sql = self._generate_sql(
question,
filtered_schema_str,
dialect,
validation_feedback=feedback,
)
last_sql = sql last_sql = sql
except Exception as e: except Exception as e:
last_errors = [f"SQL生成失败: {str(e)}"] last_errors = [f"SQL生成失败: {str(e)}"]
@@ -493,12 +621,29 @@ class Text2SQLOrchestrator:
continue continue
# === Step 3: 验证 === # === Step 3: 验证 ===
is_valid, errors, warnings = self._validate_sql( is_valid, errors, warnings, db_probe, empty_feedback = self._validate_sql(
sql, filtered_schema_str, dialect=dialect sql,
filtered_schema_str,
dialect=dialect,
question=question,
) )
if db_probe is not None:
last_db_execution_status = db_probe
if not is_valid:
last_errors = errors
logger.warning(f" [FAIL] 验证失败:{errors}")
attempt += 1
continue
logger.info(f"[OK] SQL生成与验证通过({attempt + 1}次尝试)")
meta: Dict = dict(translation_meta)
if db_probe is not None:
meta["db_execution_status"] = db_probe
if empty_feedback:
meta["db_empty_feedback"] = empty_feedback
if is_valid:
logger.info(f"[OK] SQL生成并验证通过({attempt + 1}次尝试)")
result = GenerationResult( result = GenerationResult(
sql=sql, sql=sql,
valid=True, valid=True,
@@ -507,24 +652,24 @@ class Text2SQLOrchestrator:
tables_used=tables_used, tables_used=tables_used,
attempts=attempt + 1, attempts=attempt + 1,
reasoning=reasoning if attempt == 0 else None, reasoning=reasoning if attempt == 0 else None,
metadata=meta,
) )
if include_schema_in_result: if include_schema_in_result:
result.metadata["schema"] = filtered_schema_str result.metadata["schema"] = filtered_schema_str
return result return result
# 验证失败,准备重试
last_errors = errors
logger.warning(f" [FAIL] 验证失败:{errors}")
attempt += 1
# 达到最大重试次数 # 达到最大重试次数
logger.error(f"[FAIL] 达到最大重试次数({self.max_retry}),生成失败") logger.error(f"[FAIL] 达到最大重试次数({self.max_retry}),生成失败")
fail_meta: Dict = dict(translation_meta)
if last_db_execution_status is not None:
fail_meta["db_execution_status"] = last_db_execution_status
return GenerationResult( return GenerationResult(
sql=last_sql or "", sql=last_sql or "",
valid=False, valid=False,
errors=last_errors, errors=last_errors,
tables_used=tables_used, tables_used=tables_used,
attempts=attempt, attempts=attempt,
metadata=fail_meta,
) )
def build_vector_index(self, force_rebuild: bool = False) -> bool: def build_vector_index(self, force_rebuild: bool = False) -> bool:
@@ -56,6 +56,7 @@ SQL_GENERATOR_SYSTEM = """你是一个精通SQL的数据库专家,有10年以
**硬性约束(必须遵守)**: **硬性约束(必须遵守)**:
1. **合理推断业务语义**:用户问题中的时间范围(如"2024年1月")、状态含义(如"活跃"对应Active)、常见业务默认值(如"当前"指近期),应根据Schema中的字段注释和常见业务逻辑进行合理推断并转化为WHERE条件;但禁止编造问题中未提及的过滤维度或指标。 1. **合理推断业务语义**:用户问题中的时间范围(如"2024年1月")、状态含义(如"活跃"对应Active)、常见业务默认值(如"当前"指近期),应根据Schema中的字段注释和常见业务逻辑进行合理推断并转化为WHERE条件;但禁止编造问题中未提及的过滤维度或指标。
2. **禁止虚构值与占位符**:不得使用 `'[日期]'`、`TODO`、`xxx`、空泛占位等冒充具体字面量。若用户未给出具体日期、代码或 ID,应根据问题上下文推断合理值(如"2024年1月" → `ValueDate >= '2024-01-01' AND ValueDate < '2024-02-01'`),或使用Schema中常见的枚举值(如状态字段的`A/D/X`),**不要**留空或写占位符。 2. **禁止虚构值与占位符**:不得使用 `'[日期]'`、`TODO`、`xxx`、空泛占位等冒充具体字面量。若用户未给出具体日期、代码或 ID,应根据问题上下文推断合理值(如"2024年1月" → `ValueDate >= '2024-01-01' AND ValueDate < '2024-02-01'`),或使用Schema中常见的枚举值(如状态字段的`A/D/X`),**不要**留空或写占位符。
2b. **禁止在 SQL 字符串字面量中出现中文(CJK)**:用户问题里的中文业务词(如「过户费」「未结算」「活跃」)**禁止**写成 `'…中文…'` 或 `N'…中文…'` 去和代码型列(如 `FeeNatureID`、`SettleStatus`、`State`)比较。必须根据 **Schema 字段注释** 写成库内真实**代码/单字母/数字**(如 `State = 'A'`、`SettleStatus = 'U'`);若业务词对应维表或码表,应 **JOIN 维表** 用其键列或英文名列过滤,**不得**用中文当字面量。
3. **输出版式与别名风格(统一规范)**:除遵守目标方言语法外,SQL **排版与命名**须与下方「标准版式范例」一致: 3. **输出版式与别名风格(统一规范)**:除遵守目标方言语法外,SQL **排版与命名**须与下方「标准版式范例」一致:
- **关键字**:`SELECT`、`FROM`、`JOIN`/`LEFT JOIN`、`ON`、`WHERE`、`AND`、`GROUP BY`、`ORDER BY`、`HAVING` 等使用**大写**。 - **关键字**:`SELECT`、`FROM`、`JOIN`/`LEFT JOIN`、`ON`、`WHERE`、`AND`、`GROUP BY`、`ORDER BY`、`HAVING` 等使用**大写**。
- **换行与缩进**:`SELECT` 后换行;每个输出列**独占一行**,行首 **4 个空格**,列表达式之间用**行尾逗号**分隔(最后一列无逗号)。 - **换行与缩进**:`SELECT` 后换行;每个输出列**独占一行**,行首 **4 个空格**,列表达式之间用**行尾逗号**分隔(最后一列无逗号)。
@@ -267,6 +268,21 @@ VALIDATOR_USER = """需要验证的SQL:
请输出验证结果JSON:""" 请输出验证结果JSON:"""
EMPTY_RESULT_FEEDBACK_SYSTEM = """你是数据分析助手。用户的自然语言问题已转成 SQL,且在目标库执行成功,但**当前结果在列与行上均为空**(无可用结果集)。
请用 2~5 句简洁中文说明可能原因(如条件过严、时间范围无数据、对象不存在等),并**友好引导用户补充**时间、筛选条件、业务对象等,便于下次提问更精确。
不要编造 Schema 中不存在的表或字段;不要重复输出整段 SQL;不要输出 JSON 或 Markdown 代码块。"""
EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
{question}
已执行的 SQL:
{sql}
相关 Schema(节选):
{schema}
请直接输出给终端用户阅读的说明文字(纯文本)。"""
# ========== Few-Shot 示例 ========== # ========== Few-Shot 示例 ==========
FEW_SHOT_EXAMPLES: Dict[str, str] = { FEW_SHOT_EXAMPLES: Dict[str, str] = {
@@ -319,3 +335,18 @@ WHERE ValueDate >= '2024-01-01'
AND ValueDate < '2024-02-01'; AND ValueDate < '2024-02-01';
""", """,
} }
# ========== NL 英译中(检索 / Text2SQL 前归一化)==========
TRANSLATE_NL_TO_ZH_SYSTEM = """你是证券/期货类数据仓库领域的翻译助手。
将用户给出的英文(或主要为拉丁字母的)分析需求翻译成**一句简洁的中文自然语言问题**,供后续中文向量检索与 Text2SQL 使用。
规则:
1. 语义忠实,使用业内常用中文表述(如 market value→市值、single holding→单一持仓 等)。
2. 保留阿拉伯数字、日期、币种代码、证券代码;「10 million」等与中文习惯一致时可译为「一千万」「1000万」等。
3. 若原句中出现明确的英文表名、字段名,保持英文不译。
4. **只输出中文问句本身**,不要引号、不要「翻译如下」等前后缀。"""
TRANSLATE_NL_TO_ZH_USER = """原句:
{question}
仅输出一句中文:"""
@@ -39,6 +39,10 @@ class Settings(BaseSettings):
log_level: str = "INFO" log_level: str = "INFO"
log_file: Optional[str] = None log_file: Optional[str] = None
# 业务库(backend/db:execute_sql / search_objects);与 .env 中 database_url 一致
database_url: Optional[str] = None
sql_max_rows: int = 10000
class Config: class Config:
env_file = ".env" env_file = ".env"
case_sensitive = False case_sensitive = False
+31
View File
@@ -0,0 +1,31 @@
"""
数据库工具:对齐 DBHub 的 ``execute_sql`` / ``search_objects``。
使用前请将 ``backend`` 目录加入 ``sys.path``(与 ``api_server.py`` / ``backend/main.py`` 一致),并已在进程内加载根目录 ``.env``。
"""
from __future__ import annotations
from db.dbhub_tools import (
DbHubTools,
dbhub_tools,
execute_sql,
execute_sql_all,
execute_sql_count_only,
probe_sql_execution_status,
probe_sql_execution_status_ex,
search_objects,
)
from db.engine import get_engine
__all__ = [
"DbHubTools",
"dbhub_tools",
"execute_sql",
"execute_sql_all",
"execute_sql_count_only",
"get_engine",
"probe_sql_execution_status",
"probe_sql_execution_status_ex",
"search_objects",
]
Binary file not shown.
Binary file not shown.
Binary file not shown.
+87
View File
@@ -0,0 +1,87 @@
"""
从 DBHub allowed-keywords.ts 等价移植:只读 SQL 判定。
参见 dbhub/src/utils/allowed-keywords.ts
"""
from __future__ import annotations
import re
from typing import Literal
from db.dbhub_sql_parser import ConnectorType, strip_comments_and_strings
ALLOWED_KEYWORDS: dict[ConnectorType, list[str]] = {
"postgres": ["select", "with", "explain", "show"],
"mysql": ["select", "with", "explain", "show", "describe", "desc"],
"mariadb": ["select", "with", "explain", "show", "describe", "desc"],
"sqlite": ["select", "with", "explain", "pragma"],
"sqlserver": ["select", "with", "explain", "showplan"],
}
_MUTATING = [
"insert",
"update",
"delete",
"drop",
"alter",
"create",
"truncate",
"merge",
"grant",
"revoke",
"rename",
]
_mutating_pattern = re.compile(rf"\b(?:{'|'.join(_MUTATING)})\b", re.IGNORECASE)
_mutating_pattern_with_replace = re.compile(
rf"\b(?:{'|'.join(_MUTATING)}|replace\s+(?:(?:low_priority|delayed)\s+)?into)\b",
re.IGNORECASE,
)
_MUTATING_PATTERNS: dict[ConnectorType, re.Pattern[str]] = {
"postgres": _mutating_pattern,
"mysql": _mutating_pattern_with_replace,
"mariadb": _mutating_pattern_with_replace,
"sqlite": _mutating_pattern_with_replace,
"sqlserver": _mutating_pattern,
}
_SELECT_INTO_PATTERN = re.compile(r"\bselect\b[\s\S]+\binto\b", re.IGNORECASE)
_EXPLAIN_ANALYZE_PATTERN = re.compile(
r"^explain\s+(?:\([^)]*\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)[^)]*\)|\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)(?:\s+verbose\b)?)",
re.IGNORECASE,
)
def _check_read_only(cleaned_sql: str, connector_type: ConnectorType | str) -> bool:
if not cleaned_sql:
return False
m = re.search(r"\S+", cleaned_sql)
first_word = m.group(0) if m else ""
keyword_list = ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
if first_word not in keyword_list:
return False
if first_word == "with":
pat = _MUTATING_PATTERNS.get(connector_type, _mutating_pattern) # type: ignore[arg-type]
if pat.search(cleaned_sql):
return False
if first_word in ("select", "with") and _SELECT_INTO_PATTERN.search(cleaned_sql):
return False
if first_word == "explain":
em = _EXPLAIN_ANALYZE_PATTERN.match(cleaned_sql)
if em:
after_explain = cleaned_sql[em.end() :].strip()
if after_explain and not _check_read_only(after_explain, connector_type):
return False
return True
def is_read_only_sql(sql: str, connector_type: ConnectorType | str) -> bool:
"""Check if a SQL query is read-only (DBHub-compatible)."""
cleaned = strip_comments_and_strings(sql, connector_type if connector_type in ALLOWED_KEYWORDS else None)
cleaned = cleaned.strip().lower()
return _check_read_only(cleaned, connector_type)
def allowed_keywords_list(connector_type: ConnectorType | str) -> list[str]:
return ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
+268
View File
@@ -0,0 +1,268 @@
"""
从 DBHub sql-parser.ts 等价移植:按方言剥离注释/字符串、切分语句。
参见 dbhub/src/utils/sql-parser.ts
"""
from __future__ import annotations
import re
from typing import Callable, Literal, TypedDict
ConnectorType = Literal["postgres", "mysql", "mariadb", "sqlite", "sqlserver"]
class _Token(TypedDict):
type: int # 0 Plain, 1 Comment, 2 QuotedBlock
end: int
_TOKEN_PLAIN = 0
_TOKEN_COMMENT = 1
_TOKEN_QUOTED = 2
def _plain_token(i: int) -> _Token:
return {"type": _TOKEN_PLAIN, "end": i + 1}
def _scan_single_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "-" or sql[i + 1] != "-":
return None
j = i
while j < len(sql) and sql[j] != "\n":
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
while j + 1 < len(sql) and not (sql[j] == "*" and sql[j + 1] == "/"):
j += 1
if j + 1 < len(sql):
j += 2
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment_mysql(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
nxt = sql[i + 2] if i + 2 < len(sql) else ""
nxt2 = sql[i + 3] if i + 3 < len(sql) else ""
if nxt == "!" or (nxt == "M" and nxt2 == "!"):
return None
return _scan_multi_line_comment(sql, i)
def _scan_nested_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
depth = 1
while j < len(sql) and depth > 0:
if j + 1 < len(sql) and sql[j] == "/" and sql[j + 1] == "*":
depth += 1
j += 2
elif j + 1 < len(sql) and sql[j] == "*" and sql[j + 1] == "/":
depth -= 1
j += 2
else:
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_single_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != "'":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "'" and sql[j + 1] == "'":
j += 2
elif sql[j] == "'":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_double_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != '"':
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == '"' and sql[j + 1] == '"':
j += 2
elif sql[j] == '"':
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
_dollar_quote_open_regex = re.compile(r"^\$([a-zA-Z_]\w*)?\$")
def _scan_dollar_quoted_block(sql: str, i: int) -> _Token | None:
if sql[i] != "$":
return None
nxt = sql[i + 1] if i + 1 < len(sql) else ""
if nxt.isdigit():
return None
remaining = sql[i:]
m = _dollar_quote_open_regex.match(remaining)
if not m:
return None
tag = m.group(0)
body_start = i + len(tag)
close_idx = sql.find(tag, body_start)
end = close_idx + len(tag) if close_idx != -1 else len(sql)
return {"type": _TOKEN_QUOTED, "end": end}
def _scan_backtick_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "`":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "`" and sql[j + 1] == "`":
j += 2
elif sql[j] == "`":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_bracket_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "[":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "]" and sql[j + 1] == "]":
j += 2
elif sql[j] == "]":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_token_ansi(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _plain_token(i)
)
def _scan_token_postgres(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_nested_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_dollar_quoted_block(sql, i)
or _plain_token(i)
)
def _scan_token_mysql(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment_mysql(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlite(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlserver(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
_DIALECT_SCANNERS: dict[ConnectorType, Callable[[str, int], _Token]] = {
"postgres": _scan_token_postgres,
"mysql": _scan_token_mysql,
"mariadb": _scan_token_mysql,
"sqlite": _scan_token_sqlite,
"sqlserver": _scan_token_sqlserver,
}
def _get_scanner(dialect: ConnectorType | None) -> Callable[[str, int], _Token]:
if dialect and dialect in _DIALECT_SCANNERS:
return _DIALECT_SCANNERS[dialect]
return _scan_token_ansi
def strip_comments_and_strings(sql: str, dialect: ConnectorType | None = None) -> str:
"""Replace comments, string literals, and dialect-specific quoted blocks with a single space each."""
scan_token = _get_scanner(dialect)
parts: list[str] = []
plain_start = -1
i = 0
n = len(sql)
while i < n:
token = scan_token(sql, i)
if token["type"] == _TOKEN_PLAIN:
if plain_start == -1:
plain_start = i
else:
if plain_start != -1:
parts.append(sql[plain_start:i])
plain_start = -1
parts.append(" ")
i = token["end"]
if plain_start != -1:
parts.append(sql[plain_start:])
return "".join(parts)
def split_sql_statements(sql: str, dialect: ConnectorType | None = None) -> list[str]:
"""Split SQL into individual statements, handling semicolons inside quoted contexts."""
scan_token = _get_scanner(dialect)
statements: list[str] = []
stmt_start = 0
i = 0
n = len(sql)
while i < n:
if sql[i] == ";":
trimmed = sql[stmt_start:i].strip()
if trimmed:
statements.append(trimmed)
stmt_start = i + 1
i += 1
continue
token = scan_token(sql, i)
i = token["end"]
trimmed = sql[stmt_start:].strip()
if trimmed:
statements.append(trimmed)
return statements
File diff suppressed because it is too large Load Diff
+46
View File
@@ -0,0 +1,46 @@
"""
SQLAlchemy 引擎:使用环境变量 ``database_url`` / ``DATABASE_URL``(与项目根目录 ``.env`` 一致)。
"""
from __future__ import annotations
import logging
import os
from typing import Optional
from sqlalchemy import create_engine
from sqlalchemy.engine import Engine
logger = logging.getLogger(__name__)
_engine: Engine | None = None
def _database_url_from_env() -> str:
for key in ("database_url", "DATABASE_URL"):
v = os.getenv(key)
if v is not None and str(v).strip():
return str(v).strip()
return ""
def get_engine(*, url: Optional[str] = None, reset: bool = False) -> Engine:
"""
返回默认业务库引擎(单例)。未配置 ``database_url`` 时抛出 ``ValueError``。
:param url: 若传入,则忽略单例并为此 URL 新建引擎(便于测试)。
:param reset: 为 True 时丢弃已缓存的单例,下次再按环境变量创建。
"""
global _engine
if reset:
_engine = None
if url is not None:
return create_engine(url, pool_pre_ping=True)
if _engine is not None:
return _engine
u = _database_url_from_env()
if not u:
raise ValueError("database_url 未配置")
_engine = create_engine(u, pool_pre_ping=True)
logger.info("SQLAlchemy engine initialized from database_url")
return _engine
@@ -138,7 +138,8 @@ class DeepSeekClient:
return json.loads(content) return json.loads(content)
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
logger.warning(f"JSON解析失败,返回原始内容: {e}") logger.warning(f"JSON解析失败,返回原始内容: {e}")
return {"raw_content": content} # 勿仅用 raw_content 判失败:空串时下游 `not raw.get("raw_content")` 会误判为成功
return {"_json_decode_failed": True, "raw_content": content}
def generate_sql( def generate_sql(
self, self,
@@ -173,6 +174,8 @@ class DeepSeekClient:
} }
] ]
kwargs.setdefault("temperature", 0.0)
kwargs.setdefault("top_p", 1.0)
response = self.chat(messages, **kwargs) response = self.chat(messages, **kwargs)
content = response.content.strip() content = response.content.strip()
@@ -216,8 +219,40 @@ class DeepSeekClient:
} }
] ]
kwargs.setdefault("temperature", 0.0)
kwargs.setdefault("top_p", 1.0)
return self.chat_with_json(messages, **kwargs) return self.chat_with_json(messages, **kwargs)
def empty_result_user_feedback(
self,
question: str,
sql: str,
schema: str,
*,
max_schema_chars: int = 8000,
**kwargs: Any,
) -> str:
"""
库探针为 0(执行成功但结果行数为 0)时,生成面向用户的中文补充说明,引导用户完善问题。
"""
from config.prompts import EMPTY_RESULT_FEEDBACK_SYSTEM, EMPTY_RESULT_FEEDBACK_USER
schema_snip = (schema or "")[:max_schema_chars]
messages = [
{"role": "system", "content": EMPTY_RESULT_FEEDBACK_SYSTEM},
{
"role": "user",
"content": EMPTY_RESULT_FEEDBACK_USER.format(
question=question or "(无)",
sql=sql,
schema=schema_snip,
),
},
]
msg = self.chat(messages, temperature=0.4, max_tokens=512, **kwargs)
text = (msg.content or "").strip()
return text
def select_tables( def select_tables(
self, self,
question: str, question: str,
@@ -243,13 +278,35 @@ class DeepSeekClient:
"role": "user", "role": "user",
"content": SCHEMA_LINKER_USER.format( "content": SCHEMA_LINKER_USER.format(
question=question, question=question,
table_list=table_list table_list=table_list,
) ),
} },
] ]
# 选表为结构化决策:默认 temperature=0,避免同一问题多次选不同表/SQL 上下文
kwargs.setdefault("temperature", 0.0)
kwargs.setdefault("top_p", 1.0)
return self.chat_with_json(messages, **kwargs) return self.chat_with_json(messages, **kwargs)
def translate_nl_question_to_zh(self, question: str) -> str:
"""
将主要为英文的自然语言分析问题译为中文,便于与中文 Schema 注释 / 向量索引对齐。
"""
from config.prompts import TRANSLATE_NL_TO_ZH_SYSTEM, TRANSLATE_NL_TO_ZH_USER
q = (question or "").strip()
if not q:
return ""
messages = [
{"role": "system", "content": TRANSLATE_NL_TO_ZH_SYSTEM},
{"role": "user", "content": TRANSLATE_NL_TO_ZH_USER.format(question=q)},
]
msg = self.chat(messages, temperature=0.0, top_p=1.0, max_tokens=512)
text = (msg.content or "").strip()
# 只取首行,避免模型附加说明
line = text.splitlines()[0].strip() if text else ""
return line.strip("「」\"'“”")
class AsyncDeepSeekClient: class AsyncDeepSeekClient:
""" """
+59 -54
View File
@@ -1,11 +1,8 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
""" """
Text2SQL 多智能体系统 - CLI演示入口 Text2SQL 多智能体系统 - CLI 入口
用法: 用法:在项目根目录执行 python backend/main.py
python main.py "查询2024年1月销售额最高的前5个产品"
python main.py --question "查询所有状态为Active的账户数量"
python main.py --interactive # 交互式模式
""" """
import os import os
@@ -17,21 +14,30 @@ from typing import Optional
from dotenv import load_dotenv from dotenv import load_dotenv
# 从仓库根目录运行 python backend/main.py 时,将 backend 加入模块搜索路径
_backend_dir = Path(__file__).resolve().parent
if str(_backend_dir) not in sys.path:
sys.path.insert(0, str(_backend_dir))
# 配置日志 # 配置日志
logging.basicConfig( logging.basicConfig(
level=logging.INFO, level=logging.INFO,
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s', format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
datefmt='%Y-%m-%d %H:%M:%S' datefmt='%Y-%m-%d %H:%M:%S'
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_DEFAULT_EMBEDDING_PATH = "./data/models/Qwen3-Embedding-0.6B" _DEFAULT_EMBEDDING_PATH = "./data/models/Qwen3-Embedding-0.6B"
def _repo_root() -> Path:
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
return Path(__file__).resolve().parent.parent
def _load_project_env(): def _load_project_env():
"""加载项目根目录 .env(与 main.py 同目录),供后续 os.getenv 使用。""" """加载项目根目录 .env,供后续 os.getenv 使用。"""
load_dotenv(Path(__file__).resolve().parent / ".env") load_dotenv(_repo_root() / ".env")
def _embedding_model_path() -> str: def _embedding_model_path() -> str:
@@ -149,6 +155,15 @@ def create_orchestrator(schema_mgr, args):
from agents.orchestrator import Text2SQLOrchestrator from agents.orchestrator import Text2SQLOrchestrator
from llm.deepseek_client import DeepSeekConfig from llm.deepseek_client import DeepSeekConfig
translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() not in (
"0",
"false",
"no",
"off",
)
if getattr(args, "no_translate_en", False):
translate_en = False
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip() api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
if not api_key: if not api_key:
raise ValueError( raise ValueError(
@@ -176,6 +191,7 @@ def create_orchestrator(schema_mgr, args):
fewshot_enabled=not args.no_fewshot, fewshot_enabled=not args.no_fewshot,
fewshot_top_k=args.fewshot_top_k, fewshot_top_k=args.fewshot_top_k,
fewshot_min_rating=args.fewshot_min_rating, fewshot_min_rating=args.fewshot_min_rating,
translate_english_to_zh=translate_en,
) )
return orchestrator return orchestrator
@@ -185,8 +201,30 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
"""单次查询""" """单次查询"""
import time import time
from agents.orchestrator import GenerationResult
from utils.dialog_classifier import DialogIntent, classify_dialog
logger.info(f"[Q] 问题: {question}") logger.info(f"[Q] 问题: {question}")
classified = classify_dialog(question)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
logger.info("[Q] 意图: conversation(跳过 SQL 生成)")
print("\n" + "=" * 60)
print("对话 / 非查询输入(未触发 SQL 生成)")
print("=" * 60)
print(reply)
print(f"\n使用表: []")
return GenerationResult(
sql="",
valid=False,
errors=[],
warnings=[],
tables_used=[],
attempts=0,
metadata={"dialog_intent": DialogIntent.CONVERSATION.value},
)
start = time.time() start = time.time()
result = orchestrator.generate( result = orchestrator.generate(
question=question, question=question,
@@ -212,6 +250,11 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
for w in result.warnings: for w in result.warnings:
print(f" - {w}") print(f" - {w}")
dbe = result.metadata.get("db_empty_feedback")
if result.valid and dbe:
print("\n[DB 探针 0 — 无数据行] 说明:")
print(dbe)
print(f"\n使用表: {result.tables_used}") print(f"\n使用表: {result.tables_used}")
return result return result
@@ -221,11 +264,12 @@ def interactive_mode(orchestrator, dialect: str = "tsql"):
"""交互式模式""" """交互式模式"""
print("\n" + "=" * 60) print("\n" + "=" * 60)
print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)") print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)")
print("提示:请描述业务数据查询需求;寒暄或「你是谁」等会由对话分类处理,不生成 SQL。")
print("=" * 60 + "\n") print("=" * 60 + "\n")
while True: while True:
try: try:
question = input("❓ 请输入问题: ").strip() question = input("❓ 请输入业务查询问题: ").strip()
if question.lower() in ('quit', 'exit', 'q'): if question.lower() in ('quit', 'exit', 'q'):
print("再见!") print("再见!")
break break
@@ -243,34 +287,9 @@ def interactive_mode(orchestrator, dialect: str = "tsql"):
logger.error(f"查询失败: {e}") logger.error(f"查询失败: {e}")
def batch_mode(orchestrator, questions: list, dialect: str = "tsql"):
"""批量查询模式"""
print(f"\n批量模式:共 {len(questions)} 个问题\n")
results = []
for i, question in enumerate(questions, 1):
print(f"[{i}/{len(questions)}] {question}")
result = single_query(orchestrator, question, dialect)
results.append(result)
print()
# 统计
success_count = sum(1 for r in results if r.valid)
print("=" * 60)
print(f"统计: {success_count}/{len(questions)} 成功 "
f"({success_count/len(questions)*100:.1f}%)")
return results
def main(): def main():
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
description="Text2SQL 多智能体系统 - 自然语言生成SQL" description="Text2SQL:运行后进入交互式自然语言生成 SQL(python backend/main.py)"
)
parser.add_argument(
"question",
nargs="?",
help="自然语言问题(如不提供则进入交互模式)"
) )
parser.add_argument( parser.add_argument(
"--schema", "-s", "--schema", "-s",
@@ -349,20 +368,16 @@ def main():
default=int(os.getenv("FEWSHOT_MIN_RATING", "7")), default=int(os.getenv("FEWSHOT_MIN_RATING", "7")),
help="few-shot示例最低评分(默认: 7)" help="few-shot示例最低评分(默认: 7)"
) )
parser.add_argument(
"--interactive", "-i",
action="store_true",
help="交互模式"
)
parser.add_argument(
"--batch", "-b",
help="批量文件路径(每行一个问题)"
)
parser.add_argument( parser.add_argument(
"--verbose", "-v", "--verbose", "-v",
action="store_true", action="store_true",
help="详细日志" help="详细日志"
) )
parser.add_argument(
"--no-translate-en",
action="store_true",
help="关闭英文问句自动译为中文(默认开启;也可用 TRANSLATE_EN_TO_ZH=false)",
)
args = parser.parse_args() args = parser.parse_args()
args.dialect = resolve_sql_dialect(args.dialect) args.dialect = resolve_sql_dialect(args.dialect)
@@ -389,17 +404,7 @@ def main():
logger.error(f"Orchestrator创建失败: {e}") logger.error(f"Orchestrator创建失败: {e}")
sys.exit(1) sys.exit(1)
# 根据参数选择模式
if args.interactive or (not args.question and not args.batch):
interactive_mode(orchestrator, args.dialect) interactive_mode(orchestrator, args.dialect)
elif args.batch:
with open(args.batch, 'r', encoding='utf-8') as f:
questions = [line.strip() for line in f if line.strip()]
batch_mode(orchestrator, questions, args.dialect)
elif args.question:
single_query(orchestrator, args.question, args.dialect)
else:
parser.print_help()
if __name__ == "__main__": if __name__ == "__main__":
+287
View File
@@ -0,0 +1,287 @@
"""
进程内 NL 附属数据(会话 / 消息 / 收藏),字段与 web 端 ApiEnvelope、SessionRow、MessageRow 对齐。
仅用于本地或 demo 联调,重启后数据丢失。
"""
from __future__ import annotations
import asyncio
import uuid
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Tuple
def utc_ts() -> str:
return datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
class LiteNlStore:
def __init__(self) -> None:
self._lock = asyncio.Lock()
self._sessions: Dict[Tuple[str, str], List[Dict[str, Any]]] = {}
self._messages: Dict[Tuple[str, str, str], List[Dict[str, Any]]] = {}
self._msg_counters: Dict[Tuple[str, str, str], int] = {}
self._fav: Dict[Tuple[str, str], Dict[str, List[Any]]] = {}
@staticmethod
def _visitor_key(vid: Optional[str]) -> str:
return (vid or "").strip()
def _scope(self, user_id: Optional[str], visitor_biz_id: Optional[str]) -> Tuple[str, str]:
uid = (user_id or "anonymous").strip() or "anonymous"
return uid, self._visitor_key(visitor_biz_id)
async def create_session(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
title: Optional[str],
) -> Dict[str, Any]:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
sid = str(uuid.uuid4())
now = utc_ts()
row: Dict[str, Any] = {
"session_id": sid,
"title": (title or "").strip() or None,
"user_id": sk[0],
"visitor_biz_id": visitor_biz_id or None,
"created_at": now,
"updated_at": now,
}
self._sessions.setdefault(sk, []).insert(0, row)
self._messages[(*sk, sid)] = []
self._msg_counters[(*sk, sid)] = 0
return {"session_id": sid, "title": row["title"]}
async def list_sessions(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
limit: int,
offset: int,
) -> Dict[str, Any]:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
items = list(self._sessions.get(sk, []))
items.sort(key=lambda r: str(r.get("updated_at") or ""), reverse=True)
sl = items[offset : offset + limit]
return {"items": sl, "limit": limit, "offset": offset}
def _bump_session(self, sk: Tuple[str, str], session_id: str, user_first_line: str) -> None:
now = utc_ts()
for r in self._sessions.get(sk, []):
if r.get("session_id") != session_id:
continue
r["updated_at"] = now
if not (r.get("title") or "").strip() and user_first_line.strip():
r["title"] = user_first_line.strip()[:80]
break
async def append_exchange(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
session_id: str,
user_text: str,
assistant_content_json: str,
) -> None:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
sk3 = (*sk, session_id)
if sk3 not in self._messages:
self._messages[sk3] = []
self._msg_counters[sk3] = 0
now = utc_ts()
def next_id() -> int:
n = self._msg_counters.get(sk3, 0) + 1
self._msg_counters[sk3] = n
return n
self._messages[sk3].append(
{
"id": next_id(),
"role": "user",
"content": user_text,
"created_at": now,
"llm_total_tokens": None,
}
)
self._messages[sk3].append(
{
"id": next_id(),
"role": "assistant",
"content": assistant_content_json,
"created_at": now,
"llm_total_tokens": None,
}
)
self._bump_session(sk, session_id, user_text)
async def get_messages(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
session_id: str,
limit: int,
offset: int,
) -> Optional[Dict[str, Any]]:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
sk3 = (*sk, session_id)
if sk3 not in self._messages:
found = any(
r.get("session_id") == session_id for r in self._sessions.get(sk, [])
)
if not found:
return None
self._messages[sk3] = []
self._msg_counters[sk3] = 0
rows = list(self._messages[sk3])
rows = rows[offset : offset + limit]
return {"session_id": session_id, "items": rows, "limit": limit, "offset": offset}
async def update_session_title(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
session_id: str,
title: Optional[str],
) -> Optional[Dict[str, Any]]:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
for r in self._sessions.get(sk, []):
if r.get("session_id") == session_id:
r["title"] = title
r["updated_at"] = utc_ts()
return {"session_id": session_id, "title": r.get("title")}
return None
async def delete_session(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
session_id: str,
) -> Optional[Dict[str, Any]]:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
lst = self._sessions.get(sk, [])
sk3 = (*sk, session_id)
n = len(self._messages.get(sk3, []))
new_lst = [x for x in lst if x.get("session_id") != session_id]
if len(new_lst) == len(lst):
return None
self._sessions[sk] = new_lst
self._messages.pop(sk3, None)
self._msg_counters.pop(sk3, None)
return {"session_id": session_id, "deleted_messages": n}
async def patch_message(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
session_id: str,
message_id: int,
content: str,
) -> Optional[Dict[str, Any]]:
async with self._lock:
sk3 = (*self._scope(user_id, visitor_biz_id), session_id)
for m in self._messages.get(sk3, []):
if m.get("id") == message_id:
m["content"] = content
m["created_at"] = utc_ts()
return {
"id": message_id,
"session_id": session_id,
"role": m.get("role"),
"content": content,
"created_at": m.get("created_at"),
"llm_total_tokens": m.get("llm_total_tokens"),
}
return None
async def get_favorites_grouped(
self, user_id: Optional[str], visitor_biz_id: Optional[str]
) -> Dict[str, List[Any]]:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
g = self._fav.get(sk) or {"sql": [], "function": [], "report": []}
return {"sql": list(g["sql"]), "function": list(g["function"]), "report": list(g["report"])}
async def add_favorite(
self, user_id: Optional[str], visitor_biz_id: Optional[str], body: Dict[str, Any]
) -> Any:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
g = self._fav.setdefault(sk, {"sql": [], "function": [], "report": []})
fav_type = str(body.get("fav_type") or "sql")
fid = f"{uuid.uuid4().hex[:12]}"
if fav_type == "sql":
row = {
"id": fid,
"name": str(body.get("name") or ""),
"desc": str(body.get("desc") or ""),
"sql": str(body.get("sql") or ""),
}
if body.get("sql_explain"):
row["sql_explain"] = str(body["sql_explain"])
g["sql"].insert(0, row)
return row
if fav_type == "function":
row = {
"id": fid,
"name": str(body.get("name") or ""),
"desc": str(body.get("desc") or ""),
"path": str(body.get("path") or ""),
}
g["function"].insert(0, row)
return row
row = {
"id": fid,
"name": str(body.get("name") or ""),
"desc": str(body.get("desc") or ""),
"reportPath": str(body.get("reportPath") or ""),
"params": str(body.get("params") or ""),
}
g["report"].insert(0, row)
return row
async def patch_favorite(
self,
user_id: Optional[str],
visitor_biz_id: Optional[str],
fav_id: str,
patch: Dict[str, Any],
) -> bool:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
g = self._fav.get(sk)
if not g:
return False
for bucket in ("sql", "function", "report"):
lst = g.get(bucket, [])
for i, it in enumerate(lst):
if str(it.get("id")) == fav_id:
lst[i] = {**it, **patch}
return True
return False
async def delete_favorite(
self, user_id: Optional[str], visitor_biz_id: Optional[str], fav_id: str
) -> bool:
async with self._lock:
sk = self._scope(user_id, visitor_biz_id)
g = self._fav.get(sk)
if not g:
return False
removed = False
for bucket in ("sql", "function", "report"):
before = len(g.get(bucket, []))
g[bucket] = [x for x in g.get(bucket, []) if str(x.get("id")) != fav_id]
if len(g[bucket]) < before:
removed = True
return removed
lite_nl_store = LiteNlStore()
+135
View File
@@ -0,0 +1,135 @@
"""
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
"""
from __future__ import annotations
import logging
import re
import unicodedata
from enum import Enum
from typing import NamedTuple, Optional
logger = logging.getLogger(__name__)
class DialogIntent(str, Enum):
TEXT2SQL = "text2sql"
CONVERSATION = "conversation"
class DialogClassifyResult(NamedTuple):
intent: DialogIntent
"""若为 CONVERSATION,可展示给用户的引导文案;TEXT2SQL 时为 None。"""
reply_suggestion: Optional[str] = None
DEFAULT_CONVERSATION_REPLY = (
"您好,我是业务库 Text2SQL 助手。\n"
"请用自然语言描述要查询或统计的内容(例如:查询某账户可用余额、按经纪商汇总未结算交易笔数)。\n"
"输入 quit 或 exit 可退出。"
)
_EMPTY_INPUT_REPLY = "请输入具体的业务查询问题,或输入 quit 退出。"
# 一旦出现,倾向于按「要查数据」处理(含常见业务词,避免误判)
_SQL_OR_QUERY_HINT_RE = re.compile(
r"(查|查询|查出|检索|统计|列出|汇总|求和|平均|分组|排序|排名|显示|导出|筛选|过滤|"
r"多少|几个|几张|哪些|占比|同比|环比|"
r"余额|交易|账户|持仓|报表|结算|合约|订单|流水|经纪商|对手方|证券|资金|"
r"query|select|list|show|count|sum|avg|how\s+many|statistics|\bfrom\b|\bwhere\b|\btable\b)",
re.IGNORECASE,
)
_CHITCHAT_PHRASES = frozenset(
{
"你好",
"您好",
"嗨",
"哈喽",
"hello",
"hi",
"hey",
"早上好",
"下午好",
"晚上好",
"在吗",
"在不在",
"谢谢",
"多谢",
"感谢",
"thanks",
"thank you",
"thx",
"再见",
"拜拜",
"bye",
"goodbye",
"哈哈",
"哈哈哈",
"嗯",
"嗯嗯",
"好的",
"好",
"ok",
"okay",
"行",
"收到",
"👋",
"😀",
"哈哈谢谢",
}
)
_CHITCHAT_KEYS = frozenset(p.casefold() for p in _CHITCHAT_PHRASES)
_META_QUESTION_RE = re.compile(
r"(你是谁|你是什么|你能(做|干)什么|你会什么|怎么用|如何使用|使用说明|帮助|help\b|"
r"什么功能|干啥的)",
re.IGNORECASE,
)
def _normalize(text: str) -> str:
t = unicodedata.normalize("NFKC", text or "").strip()
t = re.sub(r"\s+", " ", t)
return t
def _strip_trailing_punct(t: str) -> str:
return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip()
def classify_dialog(user_text: str) -> DialogClassifyResult:
"""
对用户一轮输入做粗分类。
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 CONVERSATION;
其余默认 TEXT2SQL,避免漏判真实查询。
"""
t = _normalize(user_text)
if not t:
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=_EMPTY_INPUT_REPLY
)
if _SQL_OR_QUERY_HINT_RE.search(t):
logger.debug("[dialog] intent=text2sql (query/business hint)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
core = _strip_trailing_punct(t)
if core.casefold() in _CHITCHAT_KEYS:
logger.debug("[dialog] intent=conversation (chitchat phrase)")
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
)
if _META_QUESTION_RE.search(t):
logger.debug("[dialog] intent=conversation (meta question)")
return DialogClassifyResult(
DialogIntent.CONVERSATION,
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
)
logger.debug("[dialog] intent=text2sql (default)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
@@ -317,7 +317,8 @@ class FewShotSelector:
def load_fewshot_selector() -> FewShotSelector: def load_fewshot_selector() -> FewShotSelector:
"""加载默认的few-shot选择器""" """加载默认的few-shot选择器"""
default_path = Path(__file__).resolve().parent.parent / "data" / "experiences" / "all_samples.jsonl" # __file__ = backend/utils/fewshot_selector.py → 仓库根为 parents[2]
default_path = Path(__file__).resolve().parents[2] / "data" / "experiences" / "all_samples.jsonl"
return FewShotSelector(str(default_path)) return FewShotSelector(str(default_path))
+23
View File
@@ -0,0 +1,23 @@
"""自然语言问题语种启发式(用于是否走「英译中」再 Text2SQL)。"""
import re
_CJK_RE = re.compile(r"[\u4e00-\u9fff]")
_HANGUL_RE = re.compile(r"[\uac00-\ud7af]")
_KANA_RE = re.compile(r"[\u3040-\u30ff]")
def looks_like_english_only(text: str) -> bool:
"""
判断问题是否主要为英文(无中日韩表意文字),适合先译成中文再走检索/选表。
含中文、日文假名、韩文时不翻译,避免破坏中英混合问句。
"""
s = (text or "").strip()
if not s or len(s) < 2:
return False
if _CJK_RE.search(s) or _HANGUL_RE.search(s) or _KANA_RE.search(s):
return False
latin = sum(1 for c in s if ("a" <= c <= "z") or ("A" <= c <= "Z"))
return latin >= 3
@@ -8,6 +8,90 @@ from typing import Tuple, List, Dict
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# CJK Unified Ideographs + 兼容扩展(用于禁止中文业务词出现在 SQL 字符串字面量中)
_CJK_IN_STRING_RE = re.compile(
r"[\u3000-\u303f\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff\uff00-\uffef]"
)
def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
"""
扫描 SQL 中单引号字符串(含 T-SQL N'…'),若字面量内出现 CJK 则判失败。
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本,避免误报。
"""
errors: List[str] = []
i = 0
n = len(sql)
in_line_comment = False
in_block_comment = False
def _read_single_quoted_string(start: int) -> Tuple[str, int]:
"""从 start 指向的 opening `'` 之后开始读,返回 (内容, 闭合引号后下标)。"""
j = start
parts: List[str] = []
while j < n:
ch = sql[j]
if ch == "'":
if j + 1 < n and sql[j + 1] == "'":
parts.append("'")
j += 2
continue
return "".join(parts), j + 1
parts.append(ch)
j += 1
return "".join(parts), j
while i < n:
if in_line_comment:
if sql[i] == "\n":
in_line_comment = False
i += 1
continue
if in_block_comment:
if i + 1 < n and sql[i : i + 2] == "*/":
in_block_comment = False
i += 2
else:
i += 1
continue
two = sql[i : i + 2]
if two == "--":
in_line_comment = True
i += 2
continue
if two == "/*":
in_block_comment = True
i += 2
continue
# N' 或 n' 前缀的 Unicode 字面量
if i + 1 < n and sql[i] in "Nn" and sql[i + 1] == "'":
body, i = _read_single_quoted_string(i + 2)
if _CJK_IN_STRING_RE.search(body):
prev = body[:48] + ("…" if len(body) > 48 else "")
errors.append(
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
)
continue
if sql[i] == "'":
body, i = _read_single_quoted_string(i + 1)
if _CJK_IN_STRING_RE.search(body):
prev = body[:48] + ("…" if len(body) > 48 else "")
errors.append(
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
)
continue
i += 1
return len(errors) == 0, errors
# 危险操作关键词(除非明确允许) # 危险操作关键词(除非明确允许)
DANGEROUS_KEYWORDS = [ DANGEROUS_KEYWORDS = [
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE", "DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE",
Binary file not shown.
+6
View File
@@ -16,6 +16,8 @@ safetensors>=0.4.0
# SQL 处理 # SQL 处理
sqlglot>=20.0.0 sqlglot>=20.0.0
sqlparse>=0.4.0 sqlparse>=0.4.0
sqlalchemy>=2.0.0
pymssql>=2.2.0
# 配置与验证 # 配置与验证
pydantic>=2.0.0 pydantic>=2.0.0
@@ -28,3 +30,7 @@ structlog>=23.0.0
tqdm>=4.65.0 tqdm>=4.65.0
jieba>=0.42.0 jieba>=0.42.0
numpy>=1.24.0,<2 numpy>=1.24.0,<2
# API框架
fastapi>=0.100.0
uvicorn>=0.20.0
+9 -2
View File
@@ -4,6 +4,13 @@ Text2SQL Few-shot 集成示例
展示如何在现有系统中集成经验数据集few-shot功能 展示如何在现有系统中集成经验数据集few-shot功能
""" """
import sys
from pathlib import Path
_root = Path(__file__).resolve().parent.parent
if str(_root / "backend") not in sys.path:
sys.path.insert(0, str(_root / "backend"))
from typing import Optional, List from typing import Optional, List
@@ -143,7 +150,7 @@ def integrate_fewshot():
``` ```
生成: data/experiences/all_samples.jsonl 生成: data/experiences/all_samples.jsonl
**步骤2:修改 agents/orchestrator.py** **步骤2:修改 backend/agents/orchestrator.py**
在文件开头添加: 在文件开头添加:
```python ```python
@@ -198,7 +205,7 @@ def integrate_fewshot():
**步骤4:测试效果** **步骤4:测试效果**
```bash ```bash
# 对比测试 # 对比测试
python main.py "查询2024年1月的销售额" --verbose python backend/main.py "查询2024年1月的销售额" --verbose
# 观察日志中的 "Few-shot已选择: Q31, Q6, Q44" # 观察日志中的 "Few-shot已选择: Q31, Q6, Q44"
``` ```
+3 -3
View File
@@ -17,9 +17,9 @@ def patch_orchestrator(
data_path: str = "data/experiences/all_samples.jsonl" data_path: str = "data/experiences/all_samples.jsonl"
): ):
""" """
修改 agents/orchestrator.py 添加few-shot支持 修改 backend/agents/orchestrator.py 添加few-shot支持
""" """
orch_path = Path(__file__).parent.parent / "agents" / "orchestrator.py" orch_path = Path(__file__).parent.parent / "backend" / "agents" / "orchestrator.py"
if not orch_path.exists(): if not orch_path.exists():
print(f"❌ 文件不存在: {orch_path}") print(f"❌ 文件不存在: {orch_path}")
@@ -224,7 +224,7 @@ def main():
print(f"1. 确保经验数据集已生成:") print(f"1. 确保经验数据集已生成:")
print(f" python scripts/parse_examples.py") print(f" python scripts/parse_examples.py")
print(f"\n2. 测试系统:") print(f"\n2. 测试系统:")
print(f" python main.py \"查询2024年1月的销售额\" --verbose") print(f" python backend/main.py \"查询2024年1月的销售额\" --verbose")
print(f"\n3. 查看日志中的few-shot选择:") print(f"\n3. 查看日志中的few-shot选择:")
print(f" Few-shot已选择: Q31, Q6, Q44") print(f" Few-shot已选择: Q31, Q6, Q44")
print(f"\n4. 调整参数:") print(f"\n4. 调整参数:")
+242
View File
@@ -0,0 +1,242 @@
# Text2SQL 算法技术方案
## 算法架构概述
Text2SQL 是一个将自然语言转换为 SQL 查询语句的智能算法系统,采用多智能体协作架构,结合向量检索和大语言模型技术,实现高效准确的 SQL 生成。
## 智能体组成
Text2SQL 系统由以下三个核心智能体组成:
| 智能体名称 | 主要职责 | 核心功能 |
| ------------- | ------ | ------------------------ |
| Schema Linker | 表选择 | 向量检索粗筛、LLM 精筛、外键扩展 |
| SQL Generator | SQL 生成 | Few-shot 示例增强、LLM 生成 SQL |
| Validator | SQL 验证 | 程序验证、库执行探针(0/1/-1)、LLM 语义验证、结果评估 |
这些智能体协同工作,形成完整的自然语言到 SQL 的转换流程。
## 核心算法流程
```mermaid
flowchart TD
subgraph 输入层
A[自然语言问题]
end
subgraph 核心处理层
A --> B[意图分类
Dialog Classifier]
B -->|非查询意图| C[返回对话回复]
subgraph "智能体 1: Schema Linker"
D[表选择]
D1[向量检索粗筛
SchemaIndexer]
D2[LLM 精筛
DeepSeek]
D3[外键扩展
关联表处理]
D --> D1
D1 --> D2
D2 --> D3
end
B -->|查询意图| D
subgraph "智能体 2: SQL Generator"
E[SQL 生成]
E1[Few-shot 示例增强]
E2[LLM 生成 SQL
DeepSeek]
E --> E1
E1 --> E2
end
D3 --> E
subgraph "智能体 3: Validator"
F[SQL 验证]
F1[程序验证
语法检查]
F1b[数据库试执行
行列探针 0/1/-1]
F2[LLM 语义验证
未探针时]
F3[结果评估]
F --> F1
F1 --> F1b
F1b --> F2
F2 --> F3
end
E2 --> F
F3 -->|验证通过| H[返回有效 SQL]
F3 -->|验证失败| I[重试逻辑
最多 max_retry 次]
I -->|重试| E
end
subgraph 数据层
J[Schema 管理
SchemaManager]
K[向量数据库
Chroma]
L[Few-shot 示例库
JSONL]
M[DeepSeek API
大语言模型]
N[业务数据库
database_url]
J --> D
K --> D1
L --> E1
M --> D2
M --> E2
M --> F2
N --> F1b
end
```
## 详细算法流程说明
### 1. 意图分类
- **Dialog Classifier**:对用户输入的自然语言进行意图分类
- **分类结果**:非查询意图直接返回对话回复,查询意图进入 SQL 生成流程
### 2. Schema Linker 算法
- **向量检索粗筛**:使用 `SchemaIndexer` 对用户问题进行向量检索,获取相关表的候选列表
- 利用预训练的嵌入模型将表名和描述转换为向量
- 使用 Chroma DB 进行相似度搜索
- 返回 top-k 个最相关的表
* **LLM 精筛**:调用 DeepSeek 模型对候选表进行精筛,选择最相关的表
- 构造包含表名和描述的提示
- 让 LLM 基于用户问题选择最相关的表
* **外键扩展**:自动添加与选中表相关的关联表,确保查询完整性
- 分析表之间的外键关系
- 自动添加被引用和引用当前表的关联表
### 3. SQL Generator 算法
- **Few-shot 示例增强**:根据用户问题从示例库中选择相似的示例,增强 SQL 生成质量
- 计算用户问题与示例库中问题的相似度
- 选择 top-k 个最相似的高质量示例
- 将示例注入到提示中,指导 SQL 生成
- **LLM 生成 SQL**:调用 DeepSeek 模型,根据选中的表结构和示例生成 SQL 语句
- 构造包含表结构、用户问题和示例的提示
- 指导 LLM 生成符合特定 SQL 方言的语句
- 清理和规范化生成的 SQL
### 4. Validator 算法
- **程序验证**:检查 SQL 语法是否正确,表和列是否存在于 Schema 中,是否包含危险操作
- 使用 sqlglot 进行语法检查
- 验证表和列是否存在于 Schema 中
- 检查是否包含危险操作(如 DROP、DELETE 等)
- **数据库试执行(探针)**:在程序验证全部通过后,于已配置的业务库上对 SQL 做只读试执行,结果用单一状态码表示,**不把具体行列数据交给 Validator LLM**:
- **1**:执行成功,且**有返回行或有返回列**(列名或数据行至少其一非空)→ **跳过** Validator 的语义审核 LLM,**直接将 SQL 交付用户**
- **0**:执行成功,但**既无列也无行** → 仍交付 SQL,**跳过** Validator 语义审核 LLM,另调 LLM 生成简短中文说明(可能原因 + 请用户补充条件),随响应一并返回
- **-1**:执行失败 → 视为本次 SQL 不可用,**不交付**,进入重试;**不调用** Validator 语义审核 LLM
- 未配置 `database_url` 时跳过探针(可选告警),此时走 **LLM 语义验证** 作为兜底
- **LLM 语义验证**:在未命中上述探针结果(即未配置库、未执行探针)时调用 DeepSeek,验证 SQL 语义是否符合用户意图
- 构造包含 SQL、表结构和用户问题的提示
- 让 LLM 评估 SQL 是否正确回答了用户问题
- 收集错误和警告信息
- **结果评估**:验证通过则返回有效 SQL,失败则进入重试逻辑
- 最多重试 max\_retry 次
- 每次重试使用相同的表结构,但重新生成 SQL
### 5. 数据层支持
- **Schema 管理**:管理数据库表结构和元数据
- 从 JSON 文件加载表结构
- 提供表和列的查询接口
- 分析表之间的外键关系
- **向量数据库**:存储表和列的向量表示,用于快速检索
- 使用 Chroma DB 存储向量
- 支持增量更新和查询
- **Few-shot 示例库**:存储高质量的自然语言到 SQL 的示例
- 从 JSONL 文件加载示例
- 支持基于相似度的示例检索
- **DeepSeek API**:提供大语言模型能力,用于表选择、SQL 生成和验证
- 调用 DeepSeek 聊天模型
- 支持不同的模型参数配置
## 技术栈
| 类别 | 技术/库 | 用途 |
| ------ | ------------ | ----------- |
| 编程语言 | Python | 算法实现 |
| LLM | DeepSeek API | 提供大语言模型能力 |
| 向量检索 | Chroma DB | 存储和检索表的向量表示 |
| SQL 处理 | sqlglot | SQL 语法解析和验证 |
| 环境管理 | dotenv | 管理环境变量 |
## 算法特点
1. **多智能体协作**:Schema Linker、SQL Generator、Validator 三个智能体协同工作,各负责专门任务,形成完整的处理流程
2. **向量检索增强**:Schema Linker 使用向量数据库快速筛选相关表,提高表选择效率
3. **Few-shot 学习**:SQL Generator 利用示例库增强 SQL 生成质量,学习最佳实践
4. **多层验证**:程序校验 + 库上探针;探针为 1/0/-1 时不再走 Validator 语义 LLM(1 直接交付、0 附带无数据说明、-1 重试),未配置库时仍以 LLM 语义验证兜底
5. **自动关联表扩展**:Schema Linker 通过外键关系自动扩展相关表,提高查询完整性
6. **可配置性**:支持多种配置参数,如温度、最大重试次数、向量检索开关等,适应不同场景需求
## Few-shot 学习详细说明
### 概念介绍
Few-shot 学习是一种机器学习方法,指通过少量示例来指导模型学习和执行任务。与传统的监督学习需要大量标注数据不同,Few-shot 学习仅需提供少量(通常为个位数)的示例,就能让模型理解任务的模式和要求。
### 在 Text2SQL 系统中的应用
在 Text2SQL 系统中,Few-shot 学习主要应用于 SQL Generator 智能体,具体流程如下:
1. **示例库构建**:系统维护一个存储高质量自然语言到 SQL 示例的库,从 JSONL 文件加载。这些示例包含各种类型的 SQL 查询场景,如简单查询、复杂连接、聚合操作等。
2. **相似度匹配**:当用户提出自然语言问题时,系统会计算该问题与示例库中问题的语义相似度。
3. **示例选择**:基于相似度排序,选择最相关的 top-k 个高质量示例。
4. **提示增强**:将这些示例注入到给大语言模型的提示中,指导模型生成更准确的 SQL 语句。
5. **生成指导**:模型参考示例的结构和风格,结合用户问题和表结构,生成符合要求的 SQL 语句。
### 技术实现
- **示例存储**:使用 JSONL 格式存储示例,每条示例包含自然语言问题和对应的 SQL 语句。
- **相似度计算**:利用预训练的嵌入模型将问题转换为向量,计算向量相似度。
- **示例选择**:根据相似度得分选择最相关的示例。
- **提示构造**:将选中的示例与用户问题、表结构一起构造提示,确保模型能理解任务要求。
### 优势
1. **减少数据需求**:不需要大量的标注数据,仅需少量高质量示例。
2. **提高生成质量**:通过示例指导,模型能生成更符合特定场景的 SQL 语句。
3. **学习最佳实践**:示例库可以包含领域专家编写的高质量 SQL,使模型学习到最佳实践。
4. **适应不同场景**:通过扩展示例库,可以适应不同领域和复杂度的 SQL 生成需求。
5. **灵活性**:可以根据具体应用场景调整示例库,提高系统的适应性。
### 应用效果
通过 Few-shot 学习,Text2SQL 系统能够:
- 处理更复杂的查询场景
- 生成更符合用户意图的 SQL 语句
- 减少生成错误
- 提高系统的泛化能力
## 性能优化
1. **向量索引预构建**:提前构建向量索引,加速首次查询
2. **缓存机制**:缓存常见查询的结果,提高响应速度
3. **并行处理**:对多个查询进行并行处理,提高系统吞吐量
4. **模型调优**:调整 LLM 参数,平衡生成质量和速度
## 扩展性
1. **支持多种数据库**:通过配置支持不同的 SQL 方言
2. **可插拔的 LLM**:支持替换不同的大语言模型
3. **自定义示例库**:可根据特定领域扩展示例库
## 总结
Text2SQL 算法通过多智能体协作、向量检索和大语言模型技术,实现了从自然语言到 SQL 的高效准确转换。算法流程清晰,逻辑完善,具有良好的可扩展性和可配置性,能够满足不同场景下的 SQL 生成需求。
+25
View File
@@ -0,0 +1,25 @@
"""DBHub 风格数据库工具;实现位于 ``backend/db``,经 ``dbhub_tools`` 转发。"""
from __future__ import annotations
from .dbhub_tools import (
DbHubTools,
dbhub_tools,
execute_sql,
execute_sql_all,
execute_sql_count_only,
probe_sql_execution_status,
probe_sql_execution_status_ex,
search_objects,
)
__all__ = [
"DbHubTools",
"dbhub_tools",
"execute_sql",
"execute_sql_all",
"execute_sql_count_only",
"probe_sql_execution_status",
"probe_sql_execution_status_ex",
"search_objects",
]
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+90
View File
@@ -0,0 +1,90 @@
"""
从 DBHub allowed-keywords.ts 等价移植:只读 SQL 判定。
参见 dbhub/src/utils/allowed-keywords.ts
与 ``backend/db/dbhub_allowed_keywords.py`` 保持一致;此处使用同目录 ``dbhub_sql_parser`` 导入,
便于在仅将 ``tools/`` 加入 ``sys.path`` 的脚本中单独使用。
"""
from __future__ import annotations
import re
from typing import Literal
from dbhub_sql_parser import ConnectorType, strip_comments_and_strings
ALLOWED_KEYWORDS: dict[ConnectorType, list[str]] = {
"postgres": ["select", "with", "explain", "show"],
"mysql": ["select", "with", "explain", "show", "describe", "desc"],
"mariadb": ["select", "with", "explain", "show", "describe", "desc"],
"sqlite": ["select", "with", "explain", "pragma"],
"sqlserver": ["select", "with", "explain", "showplan"],
}
_MUTATING = [
"insert",
"update",
"delete",
"drop",
"alter",
"create",
"truncate",
"merge",
"grant",
"revoke",
"rename",
]
_mutating_pattern = re.compile(rf"\b(?:{'|'.join(_MUTATING)})\b", re.IGNORECASE)
_mutating_pattern_with_replace = re.compile(
rf"\b(?:{'|'.join(_MUTATING)}|replace\s+(?:(?:low_priority|delayed)\s+)?into)\b",
re.IGNORECASE,
)
_MUTATING_PATTERNS: dict[ConnectorType, re.Pattern[str]] = {
"postgres": _mutating_pattern,
"mysql": _mutating_pattern_with_replace,
"mariadb": _mutating_pattern_with_replace,
"sqlite": _mutating_pattern_with_replace,
"sqlserver": _mutating_pattern,
}
_SELECT_INTO_PATTERN = re.compile(r"\bselect\b[\s\S]+\binto\b", re.IGNORECASE)
_EXPLAIN_ANALYZE_PATTERN = re.compile(
r"^explain\s+(?:\([^)]*\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)[^)]*\)|\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)(?:\s+verbose\b)?)",
re.IGNORECASE,
)
def _check_read_only(cleaned_sql: str, connector_type: ConnectorType | str) -> bool:
if not cleaned_sql:
return False
m = re.search(r"\S+", cleaned_sql)
first_word = m.group(0) if m else ""
keyword_list = ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
if first_word not in keyword_list:
return False
if first_word == "with":
pat = _MUTATING_PATTERNS.get(connector_type, _mutating_pattern) # type: ignore[arg-type]
if pat.search(cleaned_sql):
return False
if first_word in ("select", "with") and _SELECT_INTO_PATTERN.search(cleaned_sql):
return False
if first_word == "explain":
em = _EXPLAIN_ANALYZE_PATTERN.match(cleaned_sql)
if em:
after_explain = cleaned_sql[em.end() :].strip()
if after_explain and not _check_read_only(after_explain, connector_type):
return False
return True
def is_read_only_sql(sql: str, connector_type: ConnectorType | str) -> bool:
"""Check if a SQL query is read-only (DBHub-compatible)."""
cleaned = strip_comments_and_strings(sql, connector_type if connector_type in ALLOWED_KEYWORDS else None)
cleaned = cleaned.strip().lower()
return _check_read_only(cleaned, connector_type)
def allowed_keywords_list(connector_type: ConnectorType | str) -> list[str]:
return ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
+268
View File
@@ -0,0 +1,268 @@
"""
从 DBHub sql-parser.ts 等价移植:按方言剥离注释/字符串、切分语句。
参见 dbhub/src/utils/sql-parser.ts
"""
from __future__ import annotations
import re
from typing import Callable, Literal, TypedDict
ConnectorType = Literal["postgres", "mysql", "mariadb", "sqlite", "sqlserver"]
class _Token(TypedDict):
type: int # 0 Plain, 1 Comment, 2 QuotedBlock
end: int
_TOKEN_PLAIN = 0
_TOKEN_COMMENT = 1
_TOKEN_QUOTED = 2
def _plain_token(i: int) -> _Token:
return {"type": _TOKEN_PLAIN, "end": i + 1}
def _scan_single_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "-" or sql[i + 1] != "-":
return None
j = i
while j < len(sql) and sql[j] != "\n":
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
while j + 1 < len(sql) and not (sql[j] == "*" and sql[j + 1] == "/"):
j += 1
if j + 1 < len(sql):
j += 2
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment_mysql(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
nxt = sql[i + 2] if i + 2 < len(sql) else ""
nxt2 = sql[i + 3] if i + 3 < len(sql) else ""
if nxt == "!" or (nxt == "M" and nxt2 == "!"):
return None
return _scan_multi_line_comment(sql, i)
def _scan_nested_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
depth = 1
while j < len(sql) and depth > 0:
if j + 1 < len(sql) and sql[j] == "/" and sql[j + 1] == "*":
depth += 1
j += 2
elif j + 1 < len(sql) and sql[j] == "*" and sql[j + 1] == "/":
depth -= 1
j += 2
else:
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_single_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != "'":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "'" and sql[j + 1] == "'":
j += 2
elif sql[j] == "'":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_double_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != '"':
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == '"' and sql[j + 1] == '"':
j += 2
elif sql[j] == '"':
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
_dollar_quote_open_regex = re.compile(r"^\$([a-zA-Z_]\w*)?\$")
def _scan_dollar_quoted_block(sql: str, i: int) -> _Token | None:
if sql[i] != "$":
return None
nxt = sql[i + 1] if i + 1 < len(sql) else ""
if nxt.isdigit():
return None
remaining = sql[i:]
m = _dollar_quote_open_regex.match(remaining)
if not m:
return None
tag = m.group(0)
body_start = i + len(tag)
close_idx = sql.find(tag, body_start)
end = close_idx + len(tag) if close_idx != -1 else len(sql)
return {"type": _TOKEN_QUOTED, "end": end}
def _scan_backtick_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "`":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "`" and sql[j + 1] == "`":
j += 2
elif sql[j] == "`":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_bracket_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "[":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "]" and sql[j + 1] == "]":
j += 2
elif sql[j] == "]":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_token_ansi(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _plain_token(i)
)
def _scan_token_postgres(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_nested_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_dollar_quoted_block(sql, i)
or _plain_token(i)
)
def _scan_token_mysql(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment_mysql(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlite(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlserver(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
_DIALECT_SCANNERS: dict[ConnectorType, Callable[[str, int], _Token]] = {
"postgres": _scan_token_postgres,
"mysql": _scan_token_mysql,
"mariadb": _scan_token_mysql,
"sqlite": _scan_token_sqlite,
"sqlserver": _scan_token_sqlserver,
}
def _get_scanner(dialect: ConnectorType | None) -> Callable[[str, int], _Token]:
if dialect and dialect in _DIALECT_SCANNERS:
return _DIALECT_SCANNERS[dialect]
return _scan_token_ansi
def strip_comments_and_strings(sql: str, dialect: ConnectorType | None = None) -> str:
"""Replace comments, string literals, and dialect-specific quoted blocks with a single space each."""
scan_token = _get_scanner(dialect)
parts: list[str] = []
plain_start = -1
i = 0
n = len(sql)
while i < n:
token = scan_token(sql, i)
if token["type"] == _TOKEN_PLAIN:
if plain_start == -1:
plain_start = i
else:
if plain_start != -1:
parts.append(sql[plain_start:i])
plain_start = -1
parts.append(" ")
i = token["end"]
if plain_start != -1:
parts.append(sql[plain_start:])
return "".join(parts)
def split_sql_statements(sql: str, dialect: ConnectorType | None = None) -> list[str]:
"""Split SQL into individual statements, handling semicolons inside quoted contexts."""
scan_token = _get_scanner(dialect)
statements: list[str] = []
stmt_start = 0
i = 0
n = len(sql)
while i < n:
if sql[i] == ";":
trimmed = sql[stmt_start:i].strip()
if trimmed:
statements.append(trimmed)
stmt_start = i + 1
i += 1
continue
token = scan_token(sql, i)
i = token["end"]
trimmed = sql[stmt_start:].strip()
if trimmed:
statements.append(trimmed)
return statements
+41
View File
@@ -0,0 +1,41 @@
"""
DBHub 风格 SQL 执行与元数据探索。
**实现单一来源**:``backend/db/dbhub_tools.py``(含 ``execute_sql`` / ``execute_sql_all`` /
``execute_sql_count_only`` / ``search_objects`` / ``probe_sql_execution_status`` / ``probe_sql_execution_status_ex``)。
本文件将 ``backend`` 加入 ``sys.path`` 后从 ``db`` 包转发,避免 ``tools/`` 与 ``backend/db`` 双份漂移。
若你在本仓库内迭代 DB 工具逻辑,请直接修改 ``backend/db/dbhub_tools.py``。
"""
from __future__ import annotations
import sys
from pathlib import Path
_root = Path(__file__).resolve().parents[1]
_backend = _root / "backend"
if _backend.is_dir() and str(_backend) not in sys.path:
sys.path.insert(0, str(_backend))
from db.dbhub_tools import ( # noqa: E402
DbHubTools,
dbhub_tools,
execute_sql,
execute_sql_all,
execute_sql_count_only,
probe_sql_execution_status,
probe_sql_execution_status_ex,
search_objects,
)
__all__ = [
"DbHubTools",
"dbhub_tools",
"execute_sql",
"execute_sql_all",
"execute_sql_count_only",
"probe_sql_execution_status",
"probe_sql_execution_status_ex",
"search_objects",
]