0.1.1 暂存
This commit is contained in:
@@ -7,7 +7,7 @@ DEEPSEEK_BASE_URL=https://api.deepseek.com
|
||||
|
||||
# 模型配置
|
||||
# 主生成模型
|
||||
MODEL_PRIMARY=deepseek-reasoner
|
||||
MODEL_PRIMARY=deepseek-chat
|
||||
TEMPERATURE=0 # 确定性模式:每次生成相同结果
|
||||
MAX_TOKENS=4096
|
||||
# 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等
|
||||
@@ -41,7 +41,7 @@ ENABLE_SQL_VALIDATION=true
|
||||
BLOCK_DANGEROUS_SQL=true
|
||||
|
||||
# ========== 重试配置 ==========
|
||||
MAX_RETRY=1 # 确定性模式:只尝试一次,不重试
|
||||
MAX_RETRY=2 # 确定性模式:只尝试一次,不重试
|
||||
RETRY_DELAY=1.0
|
||||
|
||||
# ========== Few-shot 配置 ==========
|
||||
@@ -57,3 +57,5 @@ FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl
|
||||
# ========== 日志配置 ==========
|
||||
LOG_LEVEL=INFO
|
||||
LOG_FILE=./logs/text2sql.log
|
||||
|
||||
database_url = mssql+pymssql://sa:123456@192.168.3.201:1433/G3SB_PROD_HW
|
||||
|
||||
@@ -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. 配置变更
|
||||
|
||||
- 无。
|
||||
@@ -95,22 +95,26 @@
|
||||
|
||||
```
|
||||
text2sql_agent_camel/
|
||||
├── agents/ # Agent 层
|
||||
│ ├── orchestrator.py # 主编排器(协调 Schema Linker → Generator → Validator)
|
||||
│ ├── schema_linker.py # 表筛选 Agent(粗筛 + LLM 精筛 + 外键扩展)
|
||||
│ ├── sql_generator.py # SQL 生成 Agent(集成 Few-Shot 注入)
|
||||
│ └── validator.py # SQL 验证 Agent(程序 + LLM 双重验证)
|
||||
├── config/
|
||||
│ ├── prompts.py # 系统与用户 Prompt 模板
|
||||
│ └── settings.py # 配置类(基于 Pydantic)
|
||||
├── schema/ # Schema 管理层
|
||||
│ ├── models.py # 数据模型(Table, Column, ForeignKey, DatabaseSchema)
|
||||
│ ├── loader.py # Schema 加载器(JSON / G3SB 格式解析)
|
||||
│ ├── manager.py # Schema 管理器(查询、过滤、转字符串)
|
||||
│ └── indexer.py # 向量索引构建器(ChromaDB 封装)
|
||||
├── llm/
|
||||
│ └── deepseek_client.py # DeepSeek API 客户端(chat / validate_sql / select_tables)
|
||||
├── utils/
|
||||
├── api_server.py # FastAPI NL 网关(根目录入口;启动前将 backend 加入 sys.path)
|
||||
├── backend/ # Python 后端(除 api_server 外的编排与工具)
|
||||
│ ├── main.py # CLI 入口(交互模式:python backend/main.py)
|
||||
│ ├── nl_lite_store.py # 轻量会话/收藏内存存储(API 演示用)
|
||||
│ ├── agents/ # Agent 层
|
||||
│ │ ├── orchestrator.py # 主编排器(协调 Schema Linker → Generator → Validator)
|
||||
│ │ ├── schema_linker.py # 表筛选 Agent(粗筛 + LLM 精筛 + 外键扩展)
|
||||
│ │ ├── sql_generator.py # SQL 生成 Agent(集成 Few-Shot 注入)
|
||||
│ │ └── validator.py # SQL 验证 Agent(程序 + LLM 双重验证)
|
||||
│ ├── config/
|
||||
│ │ ├── prompts.py # 系统与用户 Prompt 模板
|
||||
│ │ └── settings.py # 配置类(基于 Pydantic)
|
||||
│ ├── schema/ # Schema 管理层
|
||||
│ │ ├── models.py # 数据模型(Table, Column, ForeignKey, DatabaseSchema)
|
||||
│ │ ├── loader.py # Schema 加载器(JSON / G3SB 格式解析)
|
||||
│ │ ├── manager.py # Schema 管理器(查询、过滤、转字符串)
|
||||
│ │ └── indexer.py # 向量索引构建器(ChromaDB 封装)
|
||||
│ ├── llm/
|
||||
│ │ └── deepseek_client.py # DeepSeek API 客户端(chat / validate_sql / select_tables)
|
||||
│ └── utils/
|
||||
│ ├── embedding.py # Qwen3-Embedding 封装(本地 / 远程 API)
|
||||
│ ├── sql_parser.py # sqlglot 工具(语法验证、方言转换、规范化)
|
||||
│ ├── validators.py # 验证逻辑(危险操作检测、完整验证流水线)
|
||||
@@ -134,7 +138,6 @@ text2sql_agent_camel/
|
||||
├── logs/
|
||||
│ └── text2sql.log # 运行日志(默认)
|
||||
├── .env # 环境变量配置(API Key、路径、参数)
|
||||
├── main.py # CLI 入口(单次查询 / 批量 / 交互模式)
|
||||
├── pyproject.toml # 项目元数据与依赖
|
||||
├── README.md # 本文档
|
||||
└── SQL_GENERATION_LOGIC.md # SQL 生成逻辑详解(流程图、示例)
|
||||
@@ -230,11 +233,14 @@ Schema JSON 格式示例:
|
||||
### 4. 构建向量索引(首次运行)
|
||||
|
||||
```bash
|
||||
# 自动构建:首次查询时会自动创建
|
||||
python main.py "测试查询"
|
||||
# 自动构建:首次在交互模式里提问时会自动创建
|
||||
python backend/main.py
|
||||
|
||||
# 或手动预构建(推荐,加速首次查询)
|
||||
# 或手动预构建(推荐,加速首次查询;请在仓库根目录执行)
|
||||
python -c "
|
||||
import sys
|
||||
from pathlib import Path
|
||||
sys.path.insert(0, str(Path.cwd() / 'backend'))
|
||||
from agents.orchestrator import Text2SQLOrchestrator
|
||||
from schema.manager import SchemaManager
|
||||
|
||||
@@ -247,30 +253,13 @@ orch.build_vector_index(force_rebuild=True)
|
||||
### 5. 运行演示
|
||||
|
||||
```bash
|
||||
# 单次查询(默认 T-SQL / SQL Server)
|
||||
python main.py "查询2024年1月的销售额"
|
||||
# 启动后进入交互式问答(默认 T-SQL / SQL Server)
|
||||
python backend/main.py
|
||||
|
||||
# 交互模式
|
||||
python main.py --interactive
|
||||
|
||||
# 批量查询(每行一个问题)
|
||||
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
|
||||
# 可选:启动时附带参数(均在进入交互前生效),例如指定 Schema、方言、Few-Shot、日志等
|
||||
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
|
||||
```
|
||||
|
||||
## 🎯 使用示例
|
||||
@@ -279,6 +268,12 @@ python main.py "查询持仓余额" --verbose
|
||||
|
||||
```python
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# 在仓库根目录运行时,将 backend 加入模块搜索路径(请先 cd 到项目根)
|
||||
sys.path.insert(0, str(Path.cwd() / "backend"))
|
||||
|
||||
from agents.orchestrator import Text2SQLOrchestrator
|
||||
from schema.manager import SchemaManager
|
||||
from llm.deepseek_client import DeepSeekConfig
|
||||
@@ -376,7 +371,7 @@ VECTOR_DB_PATH=./data/embeddings/chroma
|
||||
### CLI 参数速查
|
||||
|
||||
```bash
|
||||
python main.py [问题] [选项]
|
||||
python backend/main.py [选项]
|
||||
|
||||
# 核心选项
|
||||
--schema PATH Schema 文件路径
|
||||
@@ -398,8 +393,6 @@ python main.py [问题] [选项]
|
||||
--embedding-model PATH Embedding 模型路径
|
||||
|
||||
# 其他
|
||||
--interactive, -i 交互模式
|
||||
--batch FILE 批量文件路径
|
||||
--verbose, -v 详细日志
|
||||
```
|
||||
|
||||
@@ -564,8 +557,11 @@ WHERE i.AccountID = 'ACC001'
|
||||
# 检查索引文件
|
||||
ls -la data/embeddings/chroma/
|
||||
|
||||
# 强制重建索引
|
||||
# 强制重建索引(在仓库根目录执行)
|
||||
python -c "
|
||||
import sys
|
||||
from pathlib import Path
|
||||
sys.path.insert(0, str(Path.cwd() / 'backend'))
|
||||
from agents.orchestrator import Text2SQLOrchestrator
|
||||
from schema.manager import SchemaManager
|
||||
|
||||
@@ -587,6 +583,9 @@ orch.build_vector_index(force_rebuild=True)
|
||||
```bash
|
||||
# 检查数据文件
|
||||
python -c "
|
||||
import sys
|
||||
from pathlib import Path
|
||||
sys.path.insert(0, str(Path.cwd() / 'backend'))
|
||||
from utils.fewshot_selector import FewShotSelector
|
||||
s = FewShotSelector('./data/experiences/all_samples.jsonl')
|
||||
print(f'样本数: {len(s.samples)}')
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+581
@@ -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"
|
||||
)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -34,10 +34,11 @@ class Text2SQLOrchestrator:
|
||||
Text2SQL 多智能体编排器
|
||||
|
||||
工作流程:
|
||||
1. Schema Linker:粗筛 + LLM精筛,选出相关表
|
||||
2. 外键扩展:自动包含关联表
|
||||
3. SQL Generator:生成SQL
|
||||
4. Validator:验证SQL,不通过则重试(最多max_retry次)
|
||||
1. 粗筛候选表
|
||||
2. Schema Linker:LLM 精筛表
|
||||
3. 外键扩展 → 拼 Schema 子集
|
||||
4. SQL Generator:生成 SQL
|
||||
5. Validator:验证 SQL
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -54,6 +55,7 @@ class Text2SQLOrchestrator:
|
||||
fewshot_samples_path: Optional[str] = None,
|
||||
fewshot_top_k: int = 3,
|
||||
fewshot_min_rating: int = 7,
|
||||
translate_english_to_zh: bool = True,
|
||||
):
|
||||
"""
|
||||
初始化编排器
|
||||
@@ -66,10 +68,12 @@ class Text2SQLOrchestrator:
|
||||
vector_db_path: 向量数据库路径
|
||||
max_retry: 最大重试次数(包含首次生成)
|
||||
use_vector_search: 是否使用向量检索粗筛
|
||||
translate_english_to_zh: 无中日韩字符的英文问句是否先译为中文再走检索与生成
|
||||
"""
|
||||
self.schema_manager = schema_manager
|
||||
self.max_retry = max_retry
|
||||
self.use_vector_search = use_vector_search
|
||||
self.translate_english_to_zh = translate_english_to_zh
|
||||
|
||||
# 初始化DeepSeek客户端
|
||||
if deepseek_config:
|
||||
@@ -112,6 +116,7 @@ class Text2SQLOrchestrator:
|
||||
f"[OK] Text2SQLOrchestrator初始化完成: "
|
||||
f"max_retry={max_retry}, use_vector_search={use_vector_search}"
|
||||
+ (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:
|
||||
@@ -167,7 +172,7 @@ class Text2SQLOrchestrator:
|
||||
self,
|
||||
question: str,
|
||||
candidate_tables: List[str],
|
||||
max_tables: int = 5
|
||||
max_tables: int = 5,
|
||||
) -> Tuple[List[str], str]:
|
||||
"""
|
||||
阶段2:LLM精筛(Schema Linker Agent)
|
||||
@@ -193,7 +198,7 @@ class Text2SQLOrchestrator:
|
||||
# 使用DeepSeek客户端调用(而非CAMEL Agent,更直接可控)
|
||||
response = self.deepseek.select_tables(
|
||||
question=question,
|
||||
table_list=table_list_str
|
||||
table_list=table_list_str,
|
||||
)
|
||||
|
||||
relevant_tables = response.get("relevant_tables", [])
|
||||
@@ -284,7 +289,8 @@ class Text2SQLOrchestrator:
|
||||
self,
|
||||
question: str,
|
||||
schema_str: str,
|
||||
dialect: str = "tsql"
|
||||
dialect: str = "tsql",
|
||||
validation_feedback: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
SQL生成(SQL Generator Agent)
|
||||
@@ -293,6 +299,7 @@ class Text2SQLOrchestrator:
|
||||
question: 用户问题
|
||||
schema_str: Schema描述字符串
|
||||
dialect: SQL方言
|
||||
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
|
||||
|
||||
Returns:
|
||||
SQL语句
|
||||
@@ -336,6 +343,15 @@ class Text2SQLOrchestrator:
|
||||
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
|
||||
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
|
||||
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
|
||||
"\n**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;"
|
||||
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。"
|
||||
)
|
||||
|
||||
if validation_feedback:
|
||||
user_content += (
|
||||
"\n\n【上次校验未通过】请根据下列错误修正 SQL,并输出完整可执行查询;"
|
||||
"表名、列名必须与「当前Schema」中完全一致,不要臆造字段名。\n"
|
||||
f"{validation_feedback}"
|
||||
)
|
||||
|
||||
messages = [
|
||||
@@ -343,7 +359,8 @@ class Text2SQLOrchestrator:
|
||||
{"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()
|
||||
|
||||
# 清理可能的markdown代码块
|
||||
@@ -362,20 +379,25 @@ class Text2SQLOrchestrator:
|
||||
sql: str,
|
||||
schema_str: str,
|
||||
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:
|
||||
sql: SQL语句
|
||||
schema_str: Schema描述
|
||||
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
|
||||
question: 用户自然语言(探针为 0 时用于生成补充说明)
|
||||
|
||||
Returns:
|
||||
(是否通过, 错误列表, 警告列表)
|
||||
(是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
|
||||
探针:``1`` 至少一行数据;``0`` 执行成功但行数为 0;``-1`` 执行失败;``None`` 未配置库或未跑探针
|
||||
"""
|
||||
errors = []
|
||||
warnings = []
|
||||
db_execution_status: Optional[int] = None
|
||||
empty_feedback: Optional[str] = None
|
||||
|
||||
# === 阶段1:程序验证(确定性规则) ===
|
||||
from utils.sql_parser import validate_sql_syntax, validate_schema_consistency
|
||||
@@ -393,12 +415,43 @@ class Text2SQLOrchestrator:
|
||||
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)
|
||||
if not danger_ok:
|
||||
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:
|
||||
llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str)
|
||||
|
||||
@@ -426,8 +479,28 @@ class Text2SQLOrchestrator:
|
||||
except Exception as 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
|
||||
return is_valid, errors, warnings
|
||||
return is_valid, errors, warnings, db_execution_status, empty_feedback
|
||||
|
||||
def generate(
|
||||
self,
|
||||
@@ -448,11 +521,34 @@ class Text2SQLOrchestrator:
|
||||
Returns:
|
||||
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]}...")
|
||||
|
||||
attempt = 0
|
||||
last_sql = None
|
||||
last_errors = []
|
||||
last_db_execution_status: Optional[int] = None
|
||||
filtered_schema_str = ""
|
||||
tables_used = []
|
||||
|
||||
@@ -465,7 +561,10 @@ class Text2SQLOrchestrator:
|
||||
candidate_tables = self._coarse_filter(question, top_k=top_k_candidates)
|
||||
|
||||
# 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)
|
||||
|
||||
# 1.3 外键扩展
|
||||
@@ -480,12 +579,41 @@ class Text2SQLOrchestrator:
|
||||
)
|
||||
logger.info(f" 选中表:{relevant_tables},扩展后:{expanded_tables}")
|
||||
else:
|
||||
# 重试时复用之前的Schema
|
||||
# 重试:在子 Schema 中并入「上次失败 SQL」实际引用到的表,并对齐程序校验与生成上下文
|
||||
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生成 ===
|
||||
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
|
||||
except Exception as e:
|
||||
last_errors = [f"SQL生成失败: {str(e)}"]
|
||||
@@ -493,12 +621,29 @@ class Text2SQLOrchestrator:
|
||||
continue
|
||||
|
||||
# === Step 3: 验证 ===
|
||||
is_valid, errors, warnings = self._validate_sql(
|
||||
sql, filtered_schema_str, dialect=dialect
|
||||
is_valid, errors, warnings, db_probe, empty_feedback = self._validate_sql(
|
||||
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(
|
||||
sql=sql,
|
||||
valid=True,
|
||||
@@ -507,24 +652,24 @@ class Text2SQLOrchestrator:
|
||||
tables_used=tables_used,
|
||||
attempts=attempt + 1,
|
||||
reasoning=reasoning if attempt == 0 else None,
|
||||
metadata=meta,
|
||||
)
|
||||
if include_schema_in_result:
|
||||
result.metadata["schema"] = filtered_schema_str
|
||||
return result
|
||||
|
||||
# 验证失败,准备重试
|
||||
last_errors = errors
|
||||
logger.warning(f" [FAIL] 验证失败:{errors}")
|
||||
attempt += 1
|
||||
|
||||
# 达到最大重试次数
|
||||
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(
|
||||
sql=last_sql or "",
|
||||
valid=False,
|
||||
errors=last_errors,
|
||||
tables_used=tables_used,
|
||||
attempts=attempt,
|
||||
metadata=fail_meta,
|
||||
)
|
||||
|
||||
def build_vector_index(self, force_rebuild: bool = False) -> bool:
|
||||
BIN
Binary file not shown.
@@ -56,6 +56,7 @@ SQL_GENERATOR_SYSTEM = """你是一个精通SQL的数据库专家,有10年以
|
||||
**硬性约束(必须遵守)**:
|
||||
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`),**不要**留空或写占位符。
|
||||
2b. **禁止在 SQL 字符串字面量中出现中文(CJK)**:用户问题里的中文业务词(如「过户费」「未结算」「活跃」)**禁止**写成 `'…中文…'` 或 `N'…中文…'` 去和代码型列(如 `FeeNatureID`、`SettleStatus`、`State`)比较。必须根据 **Schema 字段注释** 写成库内真实**代码/单字母/数字**(如 `State = 'A'`、`SettleStatus = 'U'`);若业务词对应维表或码表,应 **JOIN 维表** 用其键列或英文名列过滤,**不得**用中文当字面量。
|
||||
3. **输出版式与别名风格(统一规范)**:除遵守目标方言语法外,SQL **排版与命名**须与下方「标准版式范例」一致:
|
||||
- **关键字**:`SELECT`、`FROM`、`JOIN`/`LEFT JOIN`、`ON`、`WHERE`、`AND`、`GROUP BY`、`ORDER BY`、`HAVING` 等使用**大写**。
|
||||
- **换行与缩进**:`SELECT` 后换行;每个输出列**独占一行**,行首 **4 个空格**,列表达式之间用**行尾逗号**分隔(最后一列无逗号)。
|
||||
@@ -267,6 +268,21 @@ VALIDATOR_USER = """需要验证的SQL:
|
||||
|
||||
请输出验证结果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_EXAMPLES: Dict[str, str] = {
|
||||
@@ -319,3 +335,18 @@ WHERE ValueDate >= '2024-01-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_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:
|
||||
env_file = ".env"
|
||||
case_sensitive = False
|
||||
@@ -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.
Binary file not shown.
Binary file not shown.
@@ -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]
|
||||
@@ -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
@@ -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
|
||||
Binary file not shown.
@@ -138,7 +138,8 @@ class DeepSeekClient:
|
||||
return json.loads(content)
|
||||
except json.JSONDecodeError as 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(
|
||||
self,
|
||||
@@ -173,6 +174,8 @@ class DeepSeekClient:
|
||||
}
|
||||
]
|
||||
|
||||
kwargs.setdefault("temperature", 0.0)
|
||||
kwargs.setdefault("top_p", 1.0)
|
||||
response = self.chat(messages, **kwargs)
|
||||
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)
|
||||
|
||||
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(
|
||||
self,
|
||||
question: str,
|
||||
@@ -243,13 +278,35 @@ class DeepSeekClient:
|
||||
"role": "user",
|
||||
"content": SCHEMA_LINKER_USER.format(
|
||||
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)
|
||||
|
||||
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:
|
||||
"""
|
||||
+59
-54
@@ -1,11 +1,8 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Text2SQL 多智能体系统 - CLI演示入口
|
||||
Text2SQL 多智能体系统 - CLI 入口
|
||||
|
||||
用法:
|
||||
python main.py "查询2024年1月销售额最高的前5个产品"
|
||||
python main.py --question "查询所有状态为Active的账户数量"
|
||||
python main.py --interactive # 交互式模式
|
||||
用法:在项目根目录执行 python backend/main.py
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -17,21 +14,30 @@ from typing import Optional
|
||||
|
||||
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(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_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():
|
||||
"""加载项目根目录 .env(与 main.py 同目录),供后续 os.getenv 使用。"""
|
||||
load_dotenv(Path(__file__).resolve().parent / ".env")
|
||||
"""加载项目根目录 .env,供后续 os.getenv 使用。"""
|
||||
load_dotenv(_repo_root() / ".env")
|
||||
|
||||
|
||||
def _embedding_model_path() -> str:
|
||||
@@ -149,6 +155,15 @@ def create_orchestrator(schema_mgr, args):
|
||||
from agents.orchestrator import Text2SQLOrchestrator
|
||||
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()
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
@@ -176,6 +191,7 @@ def create_orchestrator(schema_mgr, args):
|
||||
fewshot_enabled=not args.no_fewshot,
|
||||
fewshot_top_k=args.fewshot_top_k,
|
||||
fewshot_min_rating=args.fewshot_min_rating,
|
||||
translate_english_to_zh=translate_en,
|
||||
)
|
||||
|
||||
return orchestrator
|
||||
@@ -185,8 +201,30 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
|
||||
"""单次查询"""
|
||||
import time
|
||||
|
||||
from agents.orchestrator import GenerationResult
|
||||
from utils.dialog_classifier import DialogIntent, classify_dialog
|
||||
|
||||
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()
|
||||
result = orchestrator.generate(
|
||||
question=question,
|
||||
@@ -212,6 +250,11 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
|
||||
for w in result.warnings:
|
||||
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}")
|
||||
|
||||
return result
|
||||
@@ -221,11 +264,12 @@ def interactive_mode(orchestrator, dialect: str = "tsql"):
|
||||
"""交互式模式"""
|
||||
print("\n" + "=" * 60)
|
||||
print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)")
|
||||
print("提示:请描述业务数据查询需求;寒暄或「你是谁」等会由对话分类处理,不生成 SQL。")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
while True:
|
||||
try:
|
||||
question = input("❓ 请输入问题: ").strip()
|
||||
question = input("❓ 请输入业务查询问题: ").strip()
|
||||
if question.lower() in ('quit', 'exit', 'q'):
|
||||
print("再见!")
|
||||
break
|
||||
@@ -243,34 +287,9 @@ def interactive_mode(orchestrator, dialect: str = "tsql"):
|
||||
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():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Text2SQL 多智能体系统 - 自然语言生成SQL"
|
||||
)
|
||||
parser.add_argument(
|
||||
"question",
|
||||
nargs="?",
|
||||
help="自然语言问题(如不提供则进入交互模式)"
|
||||
description="Text2SQL:运行后进入交互式自然语言生成 SQL(python backend/main.py)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--schema", "-s",
|
||||
@@ -349,20 +368,16 @@ def main():
|
||||
default=int(os.getenv("FEWSHOT_MIN_RATING", "7")),
|
||||
help="few-shot示例最低评分(默认: 7)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--interactive", "-i",
|
||||
action="store_true",
|
||||
help="交互模式"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch", "-b",
|
||||
help="批量文件路径(每行一个问题)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--verbose", "-v",
|
||||
action="store_true",
|
||||
help="详细日志"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-translate-en",
|
||||
action="store_true",
|
||||
help="关闭英文问句自动译为中文(默认开启;也可用 TRANSLATE_EN_TO_ZH=false)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
args.dialect = resolve_sql_dialect(args.dialect)
|
||||
@@ -389,17 +404,7 @@ def main():
|
||||
logger.error(f"Orchestrator创建失败: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# 根据参数选择模式
|
||||
if args.interactive or (not args.question and not args.batch):
|
||||
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__":
|
||||
@@ -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()
|
||||
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
@@ -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:
|
||||
"""加载默认的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))
|
||||
|
||||
|
||||
@@ -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__)
|
||||
|
||||
# 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 = [
|
||||
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE",
|
||||
Binary file not shown.
@@ -16,6 +16,8 @@ safetensors>=0.4.0
|
||||
# SQL 处理
|
||||
sqlglot>=20.0.0
|
||||
sqlparse>=0.4.0
|
||||
sqlalchemy>=2.0.0
|
||||
pymssql>=2.2.0
|
||||
|
||||
# 配置与验证
|
||||
pydantic>=2.0.0
|
||||
@@ -28,3 +30,7 @@ structlog>=23.0.0
|
||||
tqdm>=4.65.0
|
||||
jieba>=0.42.0
|
||||
numpy>=1.24.0,<2
|
||||
|
||||
# API框架
|
||||
fastapi>=0.100.0
|
||||
uvicorn>=0.20.0
|
||||
@@ -4,6 +4,13 @@ Text2SQL 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
|
||||
|
||||
|
||||
@@ -143,7 +150,7 @@ def integrate_fewshot():
|
||||
```
|
||||
生成: data/experiences/all_samples.jsonl
|
||||
|
||||
**步骤2:修改 agents/orchestrator.py**
|
||||
**步骤2:修改 backend/agents/orchestrator.py**
|
||||
|
||||
在文件开头添加:
|
||||
```python
|
||||
@@ -198,7 +205,7 @@ def integrate_fewshot():
|
||||
**步骤4:测试效果**
|
||||
```bash
|
||||
# 对比测试
|
||||
python main.py "查询2024年1月的销售额" --verbose
|
||||
python backend/main.py "查询2024年1月的销售额" --verbose
|
||||
# 观察日志中的 "Few-shot已选择: Q31, Q6, Q44"
|
||||
```
|
||||
|
||||
|
||||
@@ -17,9 +17,9 @@ def patch_orchestrator(
|
||||
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():
|
||||
print(f"❌ 文件不存在: {orch_path}")
|
||||
@@ -224,7 +224,7 @@ def main():
|
||||
print(f"1. 确保经验数据集已生成:")
|
||||
print(f" python scripts/parse_examples.py")
|
||||
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" Few-shot已选择: Q31, Q6, Q44")
|
||||
print(f"\n4. 调整参数:")
|
||||
|
||||
@@ -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 生成需求。
|
||||
@@ -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.
Binary file not shown.
@@ -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]
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
Reference in New Issue
Block a user