diff --git a/.env b/.env
index c353dab..a40b2fc 100644
--- a/.env
+++ b/.env
@@ -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 配置 ==========
@@ -56,4 +56,6 @@ FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl
# ========== 日志配置 ==========
LOG_LEVEL=INFO
-LOG_FILE=./logs/text2sql.log
\ No newline at end of file
+LOG_FILE=./logs/text2sql.log
+
+database_url = mssql+pymssql://sa:123456@192.168.3.201:1433/G3SB_PROD_HW
diff --git a/IMPACT_ANALYSIS.md b/IMPACT_ANALYSIS.md
new file mode 100644
index 0000000..6d470c1
--- /dev/null
+++ b/IMPACT_ANALYSIS.md
@@ -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. 配置变更
+
+- 无。
diff --git a/README.md b/README.md
index c7097ca..92fdb44 100644
--- a/README.md
+++ b/README.md
@@ -95,26 +95,30 @@
```
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/
-│ ├── embedding.py # Qwen3-Embedding 封装(本地 / 远程 API)
-│ ├── sql_parser.py # sqlglot 工具(语法验证、方言转换、规范化)
-│ ├── validators.py # 验证逻辑(危险操作检测、完整验证流水线)
-│ └── fewshot_selector.py # Few-Shot 示例选择器(语义/关键词检索)
+├── 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 # 验证逻辑(危险操作检测、完整验证流水线)
+│ └── fewshot_selector.py # Few-Shot 示例选择器(语义/关键词检索)
├── data/
│ ├── schemas/ # Schema JSON 文件
│ │ ├── G3SB_MCDataDictionary_table_structure.json # 表结构
@@ -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)}')
diff --git a/__pycache__/api_server.cpython-312.pyc b/__pycache__/api_server.cpython-312.pyc
new file mode 100644
index 0000000..7db6e49
Binary files /dev/null and b/__pycache__/api_server.cpython-312.pyc differ
diff --git a/__pycache__/main.cpython-312.pyc b/__pycache__/main.cpython-312.pyc
index 5a6ff4e..93f4861 100644
Binary files a/__pycache__/main.cpython-312.pyc and b/__pycache__/main.cpython-312.pyc differ
diff --git a/__pycache__/nl_lite_store.cpython-312.pyc b/__pycache__/nl_lite_store.cpython-312.pyc
new file mode 100644
index 0000000..42118b7
Binary files /dev/null and b/__pycache__/nl_lite_store.cpython-312.pyc differ
diff --git a/agents/__pycache__/orchestrator.cpython-312.pyc b/agents/__pycache__/orchestrator.cpython-312.pyc
deleted file mode 100644
index cfbc8ce..0000000
Binary files a/agents/__pycache__/orchestrator.cpython-312.pyc and /dev/null differ
diff --git a/api_server.py b/api_server.py
new file mode 100644
index 0000000..b047936
--- /dev/null
+++ b/api_server.py
@@ -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"{inner}"]
+ 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'
{html_lib.escape(explain)}
')
+ 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"
+ )
diff --git a/__init__.py b/backend/__init__.py
similarity index 100%
rename from __init__.py
rename to backend/__init__.py
diff --git a/backend/__pycache__/main.cpython-312.pyc b/backend/__pycache__/main.cpython-312.pyc
new file mode 100644
index 0000000..947ec4b
Binary files /dev/null and b/backend/__pycache__/main.cpython-312.pyc differ
diff --git a/backend/__pycache__/nl_lite_store.cpython-312.pyc b/backend/__pycache__/nl_lite_store.cpython-312.pyc
new file mode 100644
index 0000000..f01ca6a
Binary files /dev/null and b/backend/__pycache__/nl_lite_store.cpython-312.pyc differ
diff --git a/agents/__init__.py b/backend/agents/__init__.py
similarity index 100%
rename from agents/__init__.py
rename to backend/agents/__init__.py
diff --git a/agents/__pycache__/__init__.cpython-312.pyc b/backend/agents/__pycache__/__init__.cpython-312.pyc
similarity index 100%
rename from agents/__pycache__/__init__.cpython-312.pyc
rename to backend/agents/__pycache__/__init__.cpython-312.pyc
diff --git a/backend/agents/__pycache__/orchestrator.cpython-312.pyc b/backend/agents/__pycache__/orchestrator.cpython-312.pyc
new file mode 100644
index 0000000..cfa8067
Binary files /dev/null and b/backend/agents/__pycache__/orchestrator.cpython-312.pyc differ
diff --git a/agents/__pycache__/schema_linker.cpython-312.pyc b/backend/agents/__pycache__/schema_linker.cpython-312.pyc
similarity index 100%
rename from agents/__pycache__/schema_linker.cpython-312.pyc
rename to backend/agents/__pycache__/schema_linker.cpython-312.pyc
diff --git a/agents/__pycache__/sql_generator.cpython-312.pyc b/backend/agents/__pycache__/sql_generator.cpython-312.pyc
similarity index 100%
rename from agents/__pycache__/sql_generator.cpython-312.pyc
rename to backend/agents/__pycache__/sql_generator.cpython-312.pyc
diff --git a/agents/__pycache__/validator.cpython-312.pyc b/backend/agents/__pycache__/validator.cpython-312.pyc
similarity index 100%
rename from agents/__pycache__/validator.cpython-312.pyc
rename to backend/agents/__pycache__/validator.cpython-312.pyc
diff --git a/agents/orchestrator.py b/backend/agents/orchestrator.py
similarity index 62%
rename from agents/orchestrator.py
rename to backend/agents/orchestrator.py
index cb50967..05887ef 100644
--- a/agents/orchestrator.py
+++ b/backend/agents/orchestrator.py
@@ -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,41 +415,92 @@ 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语义验证 ===
- try:
- llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str)
+ # 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)
- llm_errors = list(llm_result.get("errors", []))
- # 程序校验已通过表/列(含别名解析)时,LLM 仍常误报 unknown_*,避免误杀整次生成
- if schema_ok:
- llm_errors = [
- e
- for e in llm_errors
- if isinstance(e, str)
- and not (
- e.startswith("unknown_table:")
- or e.startswith("unknown_column:")
- )
- ]
+ # === 阶段1.5:数据库试执行(仅程序校验全部通过时;需配置 database_url) ===
+ if len(errors) == 0:
+ from db.dbhub_tools import probe_sql_execution_status_ex
- if not llm_result.get("valid", True):
- errors.extend(llm_errors)
+ 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,已跳过数据库执行探针"
+ )
- warnings.extend(llm_result.get("warnings", []))
- suggestions = llm_result.get("suggestions", [])
- if suggestions:
- logger.debug(f"优化建议:{suggestions}")
+ # 探针 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)
- except Exception as e:
- logger.warning(f"LLM验证失败(降级为仅程序验证): {e}")
+ # === 阶段2:LLM 语义验证(仅未命中库探针 0/1/-1 时) ===
+ if not skip_validator_llm:
+ try:
+ llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str)
+
+ llm_errors = list(llm_result.get("errors", []))
+ # 程序校验已通过表/列(含别名解析)时,LLM 仍常误报 unknown_*,避免误杀整次生成
+ if schema_ok:
+ llm_errors = [
+ e
+ for e in llm_errors
+ if isinstance(e, str)
+ and not (
+ e.startswith("unknown_table:")
+ or e.startswith("unknown_column:")
+ )
+ ]
+
+ if not llm_result.get("valid", True):
+ errors.extend(llm_errors)
+
+ warnings.extend(llm_result.get("warnings", []))
+ suggestions = llm_result.get("suggestions", [])
+ if suggestions:
+ logger.debug(f"优化建议:{suggestions}")
+
+ 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,38 +621,55 @@ 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 is_valid:
- logger.info(f"[OK] SQL生成并验证通过({attempt + 1}次尝试)")
- result = GenerationResult(
- sql=sql,
- valid=True,
- errors=[],
- warnings=warnings,
- tables_used=tables_used,
- attempts=attempt + 1,
- reasoning=reasoning if attempt == 0 else None,
- )
- if include_schema_in_result:
- result.metadata["schema"] = filtered_schema_str
- return result
+ if not is_valid:
+ last_errors = errors
+ logger.warning(f" [FAIL] 验证失败:{errors}")
+ attempt += 1
+ continue
- # 验证失败,准备重试
- last_errors = errors
- logger.warning(f" [FAIL] 验证失败:{errors}")
- attempt += 1
+ 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
+
+ result = GenerationResult(
+ sql=sql,
+ valid=True,
+ errors=[],
+ warnings=warnings,
+ 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
# 达到最大重试次数
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:
diff --git a/agents/schema_linker.py b/backend/agents/schema_linker.py
similarity index 100%
rename from agents/schema_linker.py
rename to backend/agents/schema_linker.py
diff --git a/agents/sql_generator.py b/backend/agents/sql_generator.py
similarity index 100%
rename from agents/sql_generator.py
rename to backend/agents/sql_generator.py
diff --git a/agents/validator.py b/backend/agents/validator.py
similarity index 100%
rename from agents/validator.py
rename to backend/agents/validator.py
diff --git a/config/__init__.py b/backend/config/__init__.py
similarity index 100%
rename from config/__init__.py
rename to backend/config/__init__.py
diff --git a/config/__pycache__/__init__.cpython-312.pyc b/backend/config/__pycache__/__init__.cpython-312.pyc
similarity index 100%
rename from config/__pycache__/__init__.cpython-312.pyc
rename to backend/config/__pycache__/__init__.cpython-312.pyc
diff --git a/config/__pycache__/prompts.cpython-312.pyc b/backend/config/__pycache__/prompts.cpython-312.pyc
similarity index 84%
rename from config/__pycache__/prompts.cpython-312.pyc
rename to backend/config/__pycache__/prompts.cpython-312.pyc
index 2c21da7..e2d8fe2 100644
Binary files a/config/__pycache__/prompts.cpython-312.pyc and b/backend/config/__pycache__/prompts.cpython-312.pyc differ
diff --git a/config/__pycache__/settings.cpython-312.pyc b/backend/config/__pycache__/settings.cpython-312.pyc
similarity index 100%
rename from config/__pycache__/settings.cpython-312.pyc
rename to backend/config/__pycache__/settings.cpython-312.pyc
diff --git a/config/prompts.py b/backend/config/prompts.py
similarity index 88%
rename from config/prompts.py
rename to backend/config/prompts.py
index 02bfd58..11d8d8d 100644
--- a/config/prompts.py
+++ b/backend/config/prompts.py
@@ -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}
+
+仅输出一句中文:"""
diff --git a/config/settings.py b/backend/config/settings.py
similarity index 87%
rename from config/settings.py
rename to backend/config/settings.py
index dbe793e..80fec32 100644
--- a/config/settings.py
+++ b/backend/config/settings.py
@@ -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
diff --git a/backend/db/__init__.py b/backend/db/__init__.py
new file mode 100644
index 0000000..49cc81c
--- /dev/null
+++ b/backend/db/__init__.py
@@ -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",
+]
diff --git a/backend/db/__pycache__/__init__.cpython-312.pyc b/backend/db/__pycache__/__init__.cpython-312.pyc
new file mode 100644
index 0000000..ceb063a
Binary files /dev/null and b/backend/db/__pycache__/__init__.cpython-312.pyc differ
diff --git a/backend/db/__pycache__/dbhub_allowed_keywords.cpython-312.pyc b/backend/db/__pycache__/dbhub_allowed_keywords.cpython-312.pyc
new file mode 100644
index 0000000..df0f441
Binary files /dev/null and b/backend/db/__pycache__/dbhub_allowed_keywords.cpython-312.pyc differ
diff --git a/backend/db/__pycache__/dbhub_sql_parser.cpython-312.pyc b/backend/db/__pycache__/dbhub_sql_parser.cpython-312.pyc
new file mode 100644
index 0000000..d7a558f
Binary files /dev/null and b/backend/db/__pycache__/dbhub_sql_parser.cpython-312.pyc differ
diff --git a/backend/db/__pycache__/dbhub_tools.cpython-312.pyc b/backend/db/__pycache__/dbhub_tools.cpython-312.pyc
new file mode 100644
index 0000000..8d8938e
Binary files /dev/null and b/backend/db/__pycache__/dbhub_tools.cpython-312.pyc differ
diff --git a/backend/db/__pycache__/engine.cpython-312.pyc b/backend/db/__pycache__/engine.cpython-312.pyc
new file mode 100644
index 0000000..61960b5
Binary files /dev/null and b/backend/db/__pycache__/engine.cpython-312.pyc differ
diff --git a/backend/db/dbhub_allowed_keywords.py b/backend/db/dbhub_allowed_keywords.py
new file mode 100644
index 0000000..35f8a59
--- /dev/null
+++ b/backend/db/dbhub_allowed_keywords.py
@@ -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]
diff --git a/backend/db/dbhub_sql_parser.py b/backend/db/dbhub_sql_parser.py
new file mode 100644
index 0000000..cdf8b3a
--- /dev/null
+++ b/backend/db/dbhub_sql_parser.py
@@ -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
diff --git a/backend/db/dbhub_tools.py b/backend/db/dbhub_tools.py
new file mode 100644
index 0000000..1499041
--- /dev/null
+++ b/backend/db/dbhub_tools.py
@@ -0,0 +1,1252 @@
+"""
+与 DBHub 对齐的数据库工具:search_objects(元数据探索)、execute_sql(SQL 执行,对齐 execute-sql.ts)、
+execute_sql_all(同 execute_sql 但不截断行数)、execute_sql_count_only(仅统计行数,不返回明细)。
+Text2SQL Validator 使用 ``probe_sql_execution_status`` / ``probe_sql_execution_status_ex`` 做库上探针。
+
+- 模块级 ``execute_sql`` / ``execute_sql_all`` / ``execute_sql_count_only`` / ``search_objects``:便捷入口,委托给 ``default_dbhub_tools``。
+- ``DbHubTools``:可实例化的工具类,方法语义与上述模块函数相同。
+- ``default_dbhub_tools``:默认单例实例,与模块级函数共用。
+"""
+
+from __future__ import annotations
+
+import json
+import logging
+import os
+import re
+from decimal import Decimal
+from typing import Any, Literal
+
+from sqlalchemy import inspect, text
+from sqlalchemy.engine import Engine
+from sqlalchemy.exc import SQLAlchemyError
+
+from db.dbhub_allowed_keywords import allowed_keywords_list, is_read_only_sql
+from db.dbhub_sql_parser import (
+ ConnectorType,
+ _get_scanner,
+ _TOKEN_PLAIN,
+ split_sql_statements,
+ strip_comments_and_strings,
+)
+from db.engine import get_engine
+
+log = logging.getLogger(__name__)
+
+
+def _get_optional_str(key: str) -> str | None:
+ """从环境变量读取(与根目录 .env / load_dotenv 一致);键名大小写不敏感。"""
+ v = os.getenv(key)
+ if v is not None and str(v).strip():
+ return str(v).strip()
+ v2 = os.getenv(key.upper())
+ if v2 is not None and str(v2).strip():
+ return str(v2).strip()
+ return None
+
+
+ObjectType = Literal["schema", "table", "column", "procedure", "function", "index"]
+DetailLevel = Literal["names", "summary", "full"]
+
+_SKIP_SCHEMAS = frozenset(
+ {
+ "information_schema",
+ "pg_catalog",
+ "pg_toast",
+ "mysql",
+ "performance_schema",
+ "sys",
+ },
+)
+
+
+def _like_pattern_to_regex(pattern: str) -> re.Pattern[str]:
+ """SQL LIKE → regex,与 DBHub likePatternToRegex 一致(% / _,其余字符 re.escape)。"""
+ parts: list[str] = []
+ for c in pattern:
+ if c == "%":
+ parts.append(".*")
+ elif c == "_":
+ parts.append(".")
+ else:
+ parts.append(re.escape(c))
+ return re.compile("^" + "".join(parts) + "$", re.IGNORECASE)
+
+
+def _bare_table_name(name: str) -> str:
+ return name.strip().rsplit(".", 1)[-1].strip()
+
+
+def _normalize_user_object_name(name: str) -> str:
+ """去掉首尾空白;若整段为 [Name] 形式则去括号(单层)。"""
+ t = name.strip()
+ if len(t) >= 2 and t[0] == "[" and t[-1] == "]":
+ return t[1:-1].replace("]]", "]")
+ return t
+
+
+def _resolve_table_or_view_bare_name(insp: Any, schema: str | None, user_table: str) -> str:
+ """
+ 将用户输入的表/视图名解析为与 SQLAlchemy inspect 一致、可在 get_columns/get_indexes 中使用的名称。
+
+ SQL Server 等库对标识符大小写不敏感,但反射时传入与系统目录不一致的大小写可能导致 get_columns 返回空;
+ 视图不在 get_table_names 中,仅传表名会漏匹配。本函数在指定 schema 下对表名与视图名做不区分大小写对齐。
+ """
+ want = _normalize_user_object_name(user_table)
+ if not want:
+ return user_table.strip()
+ key = want.lower()
+ for lister in (insp.get_table_names, insp.get_view_names):
+ try:
+ for raw in lister(schema=schema):
+ bare = _bare_table_name(str(raw))
+ if bare.lower() == key:
+ return bare
+ except Exception: # noqa: BLE001
+ continue
+ return want
+
+
+def _filter_schemas(raw: list[str]) -> list[str]:
+ return [s for s in raw if s and s.lower() not in _SKIP_SCHEMAS]
+
+
+def _qualified_table_sql(engine: Any, schema: str | None, bare_table: str) -> str:
+ prep = engine.dialect.identifier_preparer
+ tq = prep.quote(bare_table)
+ if schema:
+ return f"{prep.quote_schema(schema)}.{tq}"
+ return tq
+
+
+def _table_row_count(engine: Any, schema: str | None, bare_table: str) -> int | None:
+ q = _qualified_table_sql(engine, schema, bare_table)
+ try:
+ with engine.connect() as conn:
+ r = conn.execute(text(f"SELECT COUNT(*) AS c FROM {q}"))
+ row = r.fetchone()
+ if row is None:
+ return None
+ return int(row[0])
+ except Exception: # noqa: BLE001
+ return None
+
+
+def _pk_columns(insp: Any, schema: str | None, bare: str) -> tuple[str, ...]:
+ try:
+ pk = insp.get_pk_constraint(bare, schema=schema)
+ cols = pk.get("constrained_columns") or []
+ return tuple(str(c) for c in cols)
+ except Exception: # noqa: BLE001
+ return ()
+
+
+def _index_dicts(insp: Any, schema: str | None, bare: str) -> list[dict[str, Any]]:
+ pk_cols = _pk_columns(insp, schema, bare)
+ try:
+ raw_idx = list(insp.get_indexes(bare, schema=schema))
+ except Exception: # noqa: BLE001
+ return []
+ out: list[dict[str, Any]] = []
+ for idx in raw_idx:
+ cols = list(idx.get("column_names") or [])
+ unique = bool(idx.get("unique"))
+ primary = bool(pk_cols) and tuple(cols) == pk_cols
+ out.append(
+ {
+ "name": str(idx.get("name") or ""),
+ "columns": cols,
+ "unique": unique,
+ "primary": primary,
+ },
+ )
+ return out
+
+
+def _table_comment(insp: Any, schema: str | None, bare: str) -> str | None:
+ """获取表注释,若不支持或无注释则返回 None。"""
+ try:
+ tc = insp.get_table_comment(bare, schema=schema)
+ if isinstance(tc, dict):
+ text = (tc.get("text") or "").strip()
+ return text if text else None
+ except (NotImplementedError, AttributeError, TypeError):
+ pass
+ except Exception: # noqa: BLE001
+ pass
+ return None
+
+
+def _column_dicts(insp: Any, schema: str | None, bare: str) -> list[dict[str, Any]]:
+ try:
+ cols = list(insp.get_columns(bare, schema=schema))
+ except Exception: # noqa: BLE001
+ return []
+ out: list[dict[str, Any]] = []
+ for c in cols:
+ if not isinstance(c, dict):
+ continue
+ fname = str(c.get("name") or "").strip()
+ if not fname:
+ continue
+ ftype = c.get("type")
+ dtype = str(ftype).strip() if ftype is not None else ""
+ nullable = c.get("nullable")
+ null_b = nullable is True or (isinstance(nullable, str) and nullable.upper() == "YES")
+ entry: dict[str, Any] = {
+ "column_name": fname,
+ "data_type": dtype,
+ "is_nullable": "YES" if null_b else "NO",
+ "column_default": c.get("default"),
+ }
+ # 添加列描述(如果存在)
+ comment = c.get("comment")
+ if comment is not None:
+ comment_str = str(comment).strip()
+ if comment_str:
+ entry["description"] = comment_str
+ out.append(entry)
+ return out
+
+
+def _fetch_routines(
+ engine: Any,
+ dialect: str,
+ *,
+ schema_filter: str | None,
+ routine_sql_type: str | None,
+) -> list[tuple[str, str, str, str | None, str | None, str | None]]:
+ """
+ 返回 (schema, name, routine_type, definition_or_none, language_or_none, return_type_or_none)。
+ routine_sql_type: 'PROCEDURE' | 'FUNCTION' | None(两者都要)
+
+ 注意:INFORMATION_SCHEMA.ROUTINES 在不同数据库中包含的字段不同:
+ - SQL Server: 有 EXTERNAL_LANGUAGE(但通常为 NULL),无 RETURN_TYPE
+ - MySQL/MariaDB: 有 EXTERNAL_LANGUAGE、DTD_IDENTIFIER(返回类型)
+ - PostgreSQL: INFORMATION_SCHEMA 中存储过程支持有限,通常需要通过 pg_proc 查询
+ """
+ rows: list[tuple[str, str, str, str | None, str | None, str | None]] = []
+ if dialect not in ("mssql", "postgresql", "mysql", "mariadb"):
+ return rows
+
+ # 基础查询:所有数据库都支持的字段
+ sql = """
+ SELECT ROUTINE_SCHEMA, ROUTINE_NAME, ROUTINE_TYPE, ROUTINE_DEFINITION
+ FROM INFORMATION_SCHEMA.ROUTINES
+ WHERE 1=1
+ """
+ params: dict[str, Any] = {}
+ if schema_filter:
+ sql += " AND ROUTINE_SCHEMA = :sch"
+ params["sch"] = schema_filter
+ if routine_sql_type:
+ sql += " AND ROUTINE_TYPE = :rt"
+ params["rt"] = routine_sql_type
+
+ with engine.connect() as conn:
+ for r in conn.execute(text(sql), params):
+ defn = r[3]
+ rows.append(
+ (
+ str(r[0]),
+ str(r[1]),
+ str(r[2]),
+ None if defn is None else str(defn),
+ None, # language - INFORMATION_SCHEMA 中通常不可用
+ None, # return_type - 需要额外查询
+ ),
+ )
+
+ # 对于 MySQL/MariaDB,尝试获取更多信息
+ if dialect in ("mysql", "mariadb") and rows:
+ try:
+ extra_sql = """
+ SELECT ROUTINE_SCHEMA, ROUTINE_NAME, EXTERNAL_LANGUAGE, DTD_IDENTIFIER
+ FROM INFORMATION_SCHEMA.ROUTINES
+ WHERE 1=1
+ """
+ extra_params: dict[str, Any] = {}
+ if schema_filter:
+ extra_sql += " AND ROUTINE_SCHEMA = :sch"
+ extra_params["sch"] = schema_filter
+ if routine_sql_type:
+ extra_sql += " AND ROUTINE_TYPE = :rt"
+ extra_params["rt"] = routine_sql_type
+
+ extra_map: dict[tuple[str, str], tuple[str | None, str | None]] = {}
+ for r in conn.execute(text(extra_sql), extra_params):
+ key = (str(r[0]), str(r[1]))
+ lang = str(r[2]).strip() if r[2] else None
+ ret_type = str(r[3]).strip() if r[3] else None
+ extra_map[key] = (lang, ret_type)
+
+ # 更新已有行
+ updated_rows: list[tuple[str, str, str, str | None, str | None, str | None]] = []
+ for sch, name, rtype, defn, _, _ in rows:
+ key = (sch, name)
+ if key in extra_map:
+ lang, ret_type = extra_map[key]
+ updated_rows.append((sch, name, rtype, defn, lang, ret_type))
+ else:
+ updated_rows.append((sch, name, rtype, defn, None, None))
+ rows = updated_rows
+ except Exception: # noqa: BLE001
+ # 如果额外查询失败,保持原有数据
+ pass
+
+ return rows
+
+
+def sqlalchemy_dialect_to_connector(dialect_name: str) -> ConnectorType:
+ """
+ 将 SQLAlchemy dialect.name 映射到 DBHub ConnectorType。
+ 未知方言按 postgres 规则做只读校验(与 DBHub 移植约定一致)。
+ """
+ m: dict[str, ConnectorType] = {
+ "postgresql": "postgres",
+ "mysql": "mysql",
+ "mariadb": "mariadb",
+ "sqlite": "sqlite",
+ "mssql": "sqlserver",
+ }
+ return m.get(dialect_name, "postgres")
+
+
+_COUNT_WRAP_FIRST_WORDS = frozenset({"select", "with", "explain"})
+
+_AGGREGATE_SCALAR_PREFIX = re.compile(
+ r"(?is)^(?:distinct\s+)?(?:count|sum|avg|min|max|stdev|stdevp|string_agg|group_concat|"
+ r"variance|var_pop|var_samp|stddev|stddev_pop|stddev_samp)\s*\(",
+)
+
+
+def _parse_top_level_select_list(sql: str, connector: ConnectorType) -> tuple[int, str] | None:
+ """
+ 定位最外层 ``SELECT`` 与深度 0 上首个 ``FROM`` 之间的选择列表,返回 (列数, 列表原文断片)。
+ ``WITH`` 开头、``SELECT *``、或无法配对到 ``FROM`` 时返回 None。
+ """
+ if _first_sql_keyword(sql, connector) == "with":
+ return None
+ scan = _get_scanner(connector)
+ n = len(sql)
+ i = 0
+ list_paren = 0
+ col_commas = 0
+ state: Literal["LEADING", "IN_LIST"] = "LEADING"
+ list_start = 0
+ while i < n:
+ tok = scan(sql, i)
+ if tok["type"] != _TOKEN_PLAIN:
+ i = tok["end"]
+ continue
+ if state == "LEADING":
+ m = re.match(r"(?is)select\s+", sql[i:n])
+ if m:
+ state = "IN_LIST"
+ list_start = i + m.end()
+ i += m.end()
+ continue
+ i += 1
+ continue
+ if list_paren == 0:
+ mf = re.match(r"(?is)\bfrom\b", sql[i:n])
+ if mf:
+ frag = sql[list_start:i]
+ if re.match(r"(?is)^\s*\*\s*$", frag.strip()):
+ return None
+ return (col_commas + 1, frag)
+ c = sql[i]
+ if c == "(":
+ list_paren += 1
+ elif c == ")":
+ list_paren = max(0, list_paren - 1)
+ elif c == "," and list_paren == 0:
+ col_commas += 1
+ i += 1
+ return None
+
+
+def _has_top_level_group_by(sql: str, connector: ConnectorType) -> bool:
+ scan = _get_scanner(connector)
+ n = len(sql)
+ i = 0
+ depth = 0
+ while i < n:
+ tok = scan(sql, i)
+ if tok["type"] != _TOKEN_PLAIN:
+ i = tok["end"]
+ continue
+ if depth == 0:
+ m = re.match(r"(?is)\bgroup\s+by\b", sql[i:n])
+ if m:
+ return True
+ c = sql[i]
+ if c == "(":
+ depth += 1
+ elif c == ")":
+ depth = max(0, depth - 1)
+ i += 1
+ return False
+
+
+def _is_aggregate_scalar_fast_path(inner: str, connector: ConnectorType) -> bool:
+ """单列、无 GROUP BY、且选择列表以常见聚合函数开头时,结果集为标量,宜直接执行内层而非 COUNT(*) 包裹。"""
+ parsed = _parse_top_level_select_list(inner, connector)
+ if parsed is None:
+ return False
+ ncols, frag = parsed
+ if ncols != 1:
+ return False
+ if _has_top_level_group_by(inner, connector):
+ return False
+ fs = frag.strip()
+ if not fs or re.match(r"(?is)^\*\s*$", fs):
+ return False
+ return bool(_AGGREGATE_SCALAR_PREFIX.match(fs))
+
+
+def _strip_outer_trailing_order_by(sql: str, connector: ConnectorType) -> str:
+ """
+ SQL Server / MySQL / MariaDB:派生表内 ``ORDER BY`` 若无 TOP/OFFSET/LIMIT 会报语法错误。
+ 在 COUNT 子查询包装前去掉最外层(括号深度为 0)最后一次出现的 ``ORDER BY`` 及其后内容;
+ 不改变行数统计语义。注释与字符串内的括号/关键字不参与解析。
+ """
+ if connector not in ("sqlserver", "mysql", "mariadb"):
+ return sql
+ scan = _get_scanner(connector)
+ n = len(sql)
+ i = 0
+ depth = 0
+ last_order_by_start = -1
+ while i < n:
+ tok = scan(sql, i)
+ if tok["type"] != _TOKEN_PLAIN:
+ i = tok["end"]
+ continue
+ # 方言扫描器对「普通字符」常一次只前进一个字符;ORDER BY 必须从当前位置看到串尾才匹配得到
+ if depth == 0:
+ m = re.match(r"(?is)order\s+by\b", sql[i:n])
+ if m:
+ last_order_by_start = i
+ i += m.end()
+ continue
+ c = sql[i]
+ if c == "(":
+ depth += 1
+ elif c == ")":
+ depth = max(0, depth - 1)
+ i += 1
+ if last_order_by_start < 0:
+ return sql
+ return sql[:last_order_by_start].rstrip().rstrip(";").rstrip()
+
+
+def _first_sql_keyword(sql: str, connector: ConnectorType) -> str:
+ """剥离注释与字符串字面量后,取首个非空白词(小写)。"""
+ cleaned = strip_comments_and_strings(sql, connector)
+ cleaned = cleaned.strip().lower()
+ m = re.search(r"\S+", cleaned)
+ return m.group(0) if m else ""
+
+
+def _sql_wrapped_count(
+ inner_sql: str,
+ connector: ConnectorType,
+ ncols: int | None,
+) -> str:
+ inner = inner_sql.strip()
+ if not inner:
+ raise ValueError("SQL 不能为空")
+ inner = _strip_outer_trailing_order_by(inner, connector)
+ core = (
+ "SELECT COUNT(*) AS __dbhub_cnt FROM (\n"
+ f"{inner}\n"
+ ") AS __dbhub_subq"
+ )
+ if ncols is not None and connector in ("sqlserver", "mysql", "mariadb"):
+ names = ", ".join(f"__dbhub_c{k}" for k in range(ncols))
+ return f"{core} ({names})"
+ return core
+
+
+def _execute_sql_statements_count_only(
+ engine: Engine,
+ statements: list[str],
+ connector: ConnectorType,
+) -> int:
+ """
+ 顺序执行多条语句;对最后一条:优先用 COUNT 子查询只取标量(select/with/explain),
+ 否则单次执行后在 Python 侧逐行计数(不组装 rows 明细)。
+ """
+ with engine.connect() as conn:
+ for i, stmt in enumerate(statements):
+ s = stmt.strip()
+ if i < len(statements) - 1:
+ r = conn.execute(text(s))
+ if r.returns_rows:
+ r.fetchall()
+ continue
+ if not s:
+ raise ValueError("SQL 为空")
+ first = _first_sql_keyword(s, connector)
+ if first in _COUNT_WRAP_FIRST_WORDS:
+ inner = _strip_outer_trailing_order_by(s.strip(), connector)
+ if _is_aggregate_scalar_fast_path(inner, connector):
+ r = conn.execute(text(inner))
+ row = r.fetchone()
+ if row is None or row[0] is None:
+ return 0
+ cell = row[0]
+ if isinstance(cell, Decimal):
+ return int(cell)
+ return int(cell)
+ parsed = _parse_top_level_select_list(inner, connector)
+ if parsed is None:
+ r = conn.execute(text(inner))
+ if not r.returns_rows:
+ return 0
+ n = 0
+ for _ in r:
+ n += 1
+ return n
+ ncols, _frag = parsed
+ wrapped = _sql_wrapped_count(inner, connector, ncols)
+ r = conn.execute(text(wrapped))
+ row = r.fetchone()
+ if row is None:
+ return 0
+ return int(row[0])
+ r = conn.execute(text(s))
+ if not r.returns_rows:
+ return 0
+ n = 0
+ for _ in r:
+ n += 1
+ return n
+
+
+def _serialize_sql_cell(val: object) -> object:
+ if isinstance(val, Decimal):
+ return str(val)
+ if isinstance(val, (bytes, bytearray)):
+ return bytes(val).decode("utf-8", errors="replace")
+ if hasattr(val, "isoformat"):
+ try:
+ return val.isoformat()
+ except Exception: # noqa: BLE001
+ return val
+ return val
+
+
+def _execute_sql_statements(
+ engine: Engine,
+ statements: list[str],
+ *,
+ max_rows: int | None,
+) -> tuple[list[str], list[dict[str, object]], bool]:
+ """顺序执行多条语句,仅返回最后一条有结果集语句的列与行。``max_rows`` 为 None 时对最后一条 ``fetchall``,不截断。"""
+ last_cols: list[str] = []
+ last_rows: list[dict[str, object]] = []
+ truncated = False
+ with engine.connect() as conn:
+ result = None
+ for i, stmt in enumerate(statements):
+ s = stmt.strip()
+ result = conn.execute(text(s))
+ if i < len(statements) - 1:
+ if result.returns_rows:
+ result.fetchall()
+ continue
+ if not result.returns_rows:
+ last_cols = []
+ last_rows = []
+ truncated = False
+ break
+ last_cols = list(result.keys())
+ if max_rows is None:
+ rows_data = result.fetchall()
+ truncated = False
+ else:
+ mr = max(1, int(max_rows))
+ fetched = result.fetchmany(mr + 1)
+ truncated = len(fetched) > mr
+ rows_data = fetched[:mr]
+ last_rows = []
+ for row in rows_data:
+ row_map: dict[str, object] = {}
+ for j, col in enumerate(last_cols):
+ row_map[col] = _serialize_sql_cell(row[j])
+ last_rows.append(row_map)
+ return last_cols, last_rows, truncated
+
+
+def _execute_sql(
+ sql: str,
+ *,
+ readonly: bool = True,
+ max_rows: int | None = None,
+) -> dict[str, object]:
+ """
+ 在已配置的业务库上执行 SQL(可多语句,分号分隔)。
+ :param sql: 待执行 SQL 字符串。多条语句用分号分隔;
+ :param readonly: 是否启用只读校验,默认 True。为 True 时,任一条语句不符合 G3SB 允许的首关键字
+ (如 select、with、explain 等,随方言而异)或含变更类关键字则抛出 ValueError。
+ 为 False 时不做上述校验,可执行 DML/DDL(风险自负,勿用于不可信输入)。
+ :param max_rows: 对「最后一条返回结果集的语句」最多取多少行;None 时使用配置项 sql_max_rows,
+ 并在实现侧夹紧到 1~10000。前面的语句若产生结果集会被消费掉但不返回。
+ :return: 字典,含 rows、count、source_id(固定 ``default``)、columns、truncated。
+ 最后一条语句无结果集时 rows、columns 为空,count 为 0。未配置 database_url、SQL 为空或执行失败时
+ 抛出 ValueError(或包装后的底层异常信息)。
+ """
+ raw = (sql or "").strip()
+ if not raw:
+ raise ValueError("SQL 不能为空")
+
+ url = (_get_optional_str("database_url") or "").strip()
+ if not url:
+ raise ValueError("database_url 未配置")
+
+ engine = get_engine()
+ connector = sqlalchemy_dialect_to_connector(engine.dialect.name)
+ statements = split_sql_statements(raw, connector)
+
+ if not statements:
+ raise ValueError("SQL 为空")
+
+ if readonly:
+ for st in statements:
+ if not is_read_only_sql(st, connector):
+ kw = ", ".join(allowed_keywords_list(connector)) or "none"
+ raise ValueError(
+ f"Read-only mode is enabled. Only the following SQL operations are allowed: {kw}",
+ )
+
+ mr = max_rows if max_rows is not None else int(_get_optional_str("sql_max_rows") or 10000)
+ mr = max(1, min(10000, int(mr)))
+
+ log.info(f"_execute_sql 执行语句数={len(statements)} readonly={readonly} max_rows={mr}")
+
+ try:
+ cols, rows, truncated = _execute_sql_statements(engine, statements, max_rows=mr)
+ except Exception as e: # noqa: BLE001
+ log.error(f"_execute_sql 执行失败: {e}")
+ raise ValueError(str(e)) from e
+
+ count = len(rows)
+ return {
+ "rows": rows,
+ "count": count,
+ "source_id": "default",
+ "columns": cols,
+ "truncated": truncated,
+ }
+
+
+def _execute_sql_all(
+ sql: str,
+ *,
+ readonly: bool = True,
+) -> dict[str, object]:
+ """
+ 与 ``_execute_sql`` 相同校验与返回结构,但对最后一条有结果集的语句不做 ``max_rows`` 截断,全部 ``fetchall``。
+ 超大结果集会占用大量内存,仅用于可信 SQL / 已自行 LIMIT 的场景。
+ """
+ raw = (sql or "").strip()
+ if not raw:
+ raise ValueError("SQL 不能为空")
+
+ url = (_get_optional_str("database_url") or "").strip()
+ if not url:
+ raise ValueError("database_url 未配置")
+
+ engine = get_engine()
+ connector = sqlalchemy_dialect_to_connector(engine.dialect.name)
+ statements = split_sql_statements(raw, connector)
+
+ if not statements:
+ raise ValueError("SQL 为空")
+
+ if readonly:
+ for st in statements:
+ if not is_read_only_sql(st, connector):
+ kw = ", ".join(allowed_keywords_list(connector)) or "none"
+ raise ValueError(
+ f"Read-only mode is enabled. Only the following SQL operations are allowed: {kw}",
+ )
+
+ log.info(f"_execute_sql_all 执行语句数={len(statements)} readonly={readonly}")
+
+ try:
+ cols, rows, truncated = _execute_sql_statements(engine, statements, max_rows=None)
+ except Exception as e: # noqa: BLE001
+ log.error(f"_execute_sql_all 执行失败: {e}")
+ raise ValueError(str(e)) from e
+
+ count = len(rows)
+ return {
+ "rows": rows,
+ "count": count,
+ "source_id": "default",
+ "columns": cols,
+ "truncated": truncated,
+ }
+
+
+def _execute_sql_count_only(
+ sql: str,
+ *,
+ readonly: bool = True,
+) -> dict[str, int]:
+ """
+ 在已配置的业务库上执行 SQL,仅返回「最后一条有结果集语句」的行数,不返回单元格明细。
+
+ 对以 select / with / explain 开头(注释剥离后)的最后一条语句:若已为单列常见聚合(如 ``count(1)``)且无 ``GROUP BY``,
+ 则直接执行该语句并取标量,避免 ``COUNT(*)`` 外包一层只得到 1 行。否则在库侧用 ``SELECT COUNT(*) FROM ( ... )``;
+ 对 SQL Server / MySQL / MariaDB 会为派生表补上 ``(__dbhub_c0, ...)`` 列名(避免匿名列 8155 等),并在包装前去掉
+ 最外层无意义的 ``ORDER BY``。无法解析选择列表时(如 ``WITH``、``SELECT *``)改为拉全结果在 Python 侧逐行计数。
+ 其它只读语句(如部分 show/pragma)仍单次执行后逐行计数,不组装为 dict。
+
+ 多条语句用分号分隔时,前面的语句照常执行并消费结果集;仅对最后一条给出计数。
+ 最后一条无结果集时 count 为 0。进入执行阶段后若任一步失败(含语法、权限、连接中断等),返回 ``{"count": -1}``;
+ SQL 为空、未配置 ``database_url``、只读校验不通过仍抛 ``ValueError``。
+ """
+ raw = (sql or "").strip()
+ if not raw:
+ raise ValueError("SQL 不能为空")
+
+ url = (_get_optional_str("database_url") or "").strip()
+ if not url:
+ raise ValueError("database_url 未配置")
+
+ engine = get_engine()
+ connector = sqlalchemy_dialect_to_connector(engine.dialect.name)
+ statements = split_sql_statements(raw, connector)
+
+ if not statements:
+ raise ValueError("SQL 为空")
+
+ if readonly:
+ for st in statements:
+ if not is_read_only_sql(st, connector):
+ kw = ", ".join(allowed_keywords_list(connector)) or "none"
+ raise ValueError(
+ f"Read-only mode is enabled. Only the following SQL operations are allowed: {kw}",
+ )
+
+ log.info(f"_execute_sql_count_only 执行语句数={len(statements)} readonly={readonly}")
+
+ try:
+ n = _execute_sql_statements_count_only(engine, statements, connector)
+ except Exception as e: # noqa: BLE001
+ log.error(f"_execute_sql_count_only 执行失败: {e}")
+ return {"count": -1}
+
+ return {"count": n}
+
+
+def probe_sql_execution_status_ex(
+ sql: str, *, max_rows: int = 1
+) -> tuple[int | None, str | None]:
+ """
+ 在已配置的业务库上对 SQL 做只读试执行,返回状态码与失败时的错误摘要。
+
+ Returns:
+ 二元组 ``(status, error_message)``:
+
+ - ``(None, None)``:未配置 ``database_url``,跳过探针。
+ - ``(1, None)``:执行成功,且**至少返回一行数据**(与「有列无行」的 SELECT 区分)。
+ - ``(0, None)``:执行成功,但**数据行数为 0**(可有列名而无行,或无非空结果集)。
+ - ``(-1, msg)``:执行失败;``msg`` 为异常信息摘要(便于生成阶段重试)。
+ """
+ url = (_get_optional_str("database_url") or "").strip()
+ if not url:
+ return None, None
+ try:
+ out = _execute_sql(sql, readonly=True, max_rows=max_rows)
+ except Exception as e: # noqa: BLE001
+ log.warning(f"probe_sql_execution_status 执行失败: {e}")
+ detail = str(e).strip() or "unknown error"
+ return -1, detail
+ rows = list(out.get("rows") or [])
+ if len(rows) > 0:
+ return 1, None
+ return 0, None
+
+
+def probe_sql_execution_status(sql: str, *, max_rows: int = 1) -> int | None:
+ """
+ 在已配置的业务库上对 SQL 做只读试执行,返回紧凑状态码(用于 Validator 探针)。
+
+ 语义与 :func:`probe_sql_execution_status_ex` 的首个返回值一致;失败细节请用 ``_ex``。
+ """
+ status, _ = probe_sql_execution_status_ex(sql, max_rows=max_rows)
+ return status
+
+
+def _search_objects(
+ object_type: ObjectType,
+ pattern: str = "%",
+ schema: str | None = None,
+ table: str | None = None,
+ detail_level: DetailLevel = "names",
+ limit: int = 100,
+) -> dict[str, Any]:
+ """
+ 按对象类型在已配置的业务库中探索 schema、表、列、索引、存储过程/函数等元数据(对齐 DBHub search_objects)。
+ 会过滤系统 schema(如information_schema、sys、pg_catalog 等),再按 object_type 与 pattern 做 SQL LIKE 风格匹配;
+ :param object_type: 要探索的对象类别。schema-模式名;table-数据表;column-列(可与 table、schema 联用
+ 限定单表);index-索引(同上);procedure-存储过程;function-函数(标量/表值等,视库而定)。
+ :param pattern: SQL LIKE 模式,默认 "%" 表示不过滤名称。"%" 匹配任意长度子串,"_" 匹配单个字符;对表名、
+ 列名、索引名、例程名等做大小写不敏感匹配。
+ :param schema: 限定在某个 schema 内查找;None 表示在多个非系统 schema 上依次查找。若给出具体名称,
+ 必须是库中已存在的 schema,否则抛 ValueError。使用参数 table 时必须同时指定 schema。
+ :param table: 仅在 object_type 为 column 或 index 时允许传入,与 schema 共同限定「只查这一张表或视图」上的列
+ 或索引;用于其它 object_type 时会抛 ValueError。名称会在该 schema 下与反射得到的表名、视图名做
+ 不区分大小写匹配(并支持 [Name] 写法),以兼容 SQL Server 等对目录大小写不敏感但 API 需真实写法的情况。
+ :param detail_level: 返回粒度。names-仅对象名及定位字段(如 schema、table);summary-增加简要元数据
+ (如表的 column_count、row_count,列的类型、可空等);full-表级返回列列表、索引列表等完整结构。
+ :param limit: 最多返回的结果条数,默认 100,有效范围 1~1000(传入值会被夹紧到该区间)。
+ :return: 包含 object_type、pattern、schema、table、detail_level、count、results、truncated 的字典。
+ 未配置 database_url、连接失败或内省失败时抛出 ValueError(或其它底层异常)。
+ """
+ if table and not schema:
+ raise ValueError("The 'table' parameter requires 'schema' to be specified")
+ if table and object_type not in ("column", "index"):
+ raise ValueError(
+ f"The 'table' parameter only applies to object_type 'column' or 'index', not '{object_type}'",
+ )
+
+ lim = max(1, min(1000, int(limit)))
+ pat = pattern if pattern is not None else "%"
+ rx = _like_pattern_to_regex(pat)
+
+ url = (_get_optional_str("database_url") or "").strip()
+ if not url:
+ raise ValueError("database_url 未配置")
+
+ try:
+ engine = get_engine()
+ except ValueError as e:
+ raise ValueError(str(e)) from e
+
+ try:
+ insp = inspect(engine)
+ except SQLAlchemyError as e:
+ log.error(f"_search_objects inspect 失败: {e}")
+ raise ValueError(f"数据库内省失败: {e}") from e
+
+ try:
+ all_schema_names = list(insp.get_schema_names())
+ except Exception: # noqa: BLE001
+ all_schema_names = []
+ filtered_schemas = _filter_schemas(all_schema_names)
+
+ if schema:
+ if schema not in all_schema_names:
+ avail = ", ".join(all_schema_names) or "(none)"
+ raise ValueError(f"Schema '{schema}' does not exist. Available schemas: {avail}")
+ schemas_to_search = [schema]
+ else:
+ schemas_to_search = filtered_schemas if filtered_schemas else [None]
+
+ dialect = engine.dialect.name
+ results: list[Any] = []
+
+ if object_type == "schema":
+ base_schemas = filtered_schemas if filtered_schemas else all_schema_names
+ candidates = [s for s in base_schemas if rx.match(s)]
+ for schema_name in candidates:
+ if len(results) >= lim:
+ break
+ if detail_level == "names":
+ results.append({"name": schema_name})
+ else:
+ try:
+ tbls = list(insp.get_table_names(schema=schema_name))
+ except Exception: # noqa: BLE001
+ tbls = []
+ results.append({"name": schema_name, "table_count": len(tbls)})
+
+ elif object_type == "table":
+ for schema_name in schemas_to_search:
+ if len(results) >= lim:
+ break
+ try:
+ tbls = list(insp.get_table_names(schema=schema_name))
+ except Exception: # noqa: BLE001
+ continue
+ for table_name in tbls:
+ if len(results) >= lim:
+ break
+ bare = _bare_table_name(str(table_name))
+ if not rx.match(bare):
+ continue
+ if detail_level == "names":
+ results.append({"name": bare, "schema": schema_name})
+ elif detail_level == "summary":
+ cols = _column_dicts(insp, schema_name, bare)
+ rc = _table_row_count(engine, schema_name, bare)
+ comment = _table_comment(insp, schema_name, bare)
+ result_dict: dict[str, Any] = {
+ "name": bare,
+ "schema": schema_name,
+ "column_count": len(cols),
+ "row_count": rc,
+ }
+ if comment:
+ result_dict["comment"] = comment
+ results.append(result_dict)
+ else:
+ cols = _column_dicts(insp, schema_name, bare)
+ idxs = _index_dicts(insp, schema_name, bare)
+ rc = _table_row_count(engine, schema_name, bare)
+ comment = _table_comment(insp, schema_name, bare)
+ result_dict_full: dict[str, Any] = {
+ "name": bare,
+ "schema": schema_name,
+ "column_count": len(cols),
+ "row_count": rc,
+ "columns": [
+ {
+ "name": c["column_name"],
+ "type": c["data_type"],
+ "nullable": c["is_nullable"] == "YES",
+ "default": c["column_default"],
+ **({"description": c["description"]} if "description" in c else {}),
+ }
+ for c in cols
+ ],
+ "indexes": idxs,
+ }
+ if comment:
+ result_dict_full["comment"] = comment
+ results.append(result_dict_full)
+
+ elif object_type in ("column", "index"):
+ for schema_name in schemas_to_search:
+ if len(results) >= lim:
+ break
+ try:
+ if table:
+ tables_to_search = [_resolve_table_or_view_bare_name(insp, schema_name, table)]
+ else:
+ tables_to_search = [_bare_table_name(str(t)) for t in insp.get_table_names(schema=schema_name)]
+ except Exception: # noqa: BLE001
+ continue
+ for table_name in tables_to_search:
+ if len(results) >= lim:
+ break
+ bare = _bare_table_name(str(table_name))
+ if object_type == "column":
+ cols = _column_dicts(insp, schema_name, bare)
+ for c in cols:
+ if len(results) >= lim:
+ break
+ if not rx.match(c["column_name"]):
+ continue
+ if detail_level == "names":
+ results.append({"name": c["column_name"], "table": bare, "schema": schema_name})
+ else:
+ result_dict_col: dict[str, Any] = {
+ "name": c["column_name"],
+ "table": bare,
+ "schema": schema_name,
+ "type": c["data_type"],
+ "nullable": c["is_nullable"] == "YES",
+ "default": c["column_default"],
+ }
+ if "description" in c:
+ result_dict_col["description"] = c["description"]
+ results.append(result_dict_col)
+ else:
+ idxs = _index_dicts(insp, schema_name, bare)
+ for idx in idxs:
+ if len(results) >= lim:
+ break
+ iname = str(idx.get("name") or "")
+ if not rx.match(iname):
+ continue
+ if detail_level == "names":
+ results.append({"name": iname, "table": bare, "schema": schema_name})
+ else:
+ results.append(
+ {
+ "name": iname,
+ "table": bare,
+ "schema": schema_name,
+ "columns": idx.get("columns"),
+ "unique": idx.get("unique"),
+ "primary": idx.get("primary"),
+ },
+ )
+
+ elif object_type in ("procedure", "function"):
+ rt = "PROCEDURE" if object_type == "procedure" else "FUNCTION"
+ try:
+ routine_rows = _fetch_routines(engine, dialect, schema_filter=schema, routine_sql_type=rt)
+ except Exception as e: # noqa: BLE001
+ log.warning(f"_search_objects 例程列表失败 dialect={dialect}: {e}")
+ routine_rows = []
+ for sch, name, rtype, defn in routine_rows:
+ if len(results) >= lim:
+ break
+ if not rx.match(name):
+ continue
+ if detail_level == "names":
+ results.append({"name": name, "schema": sch})
+ elif detail_level == "summary":
+ results.append(
+ {
+ "name": name,
+ "schema": sch,
+ "type": rtype,
+ "language": None,
+ "return_type": None,
+ },
+ )
+ else:
+ results.append(
+ {
+ "name": name,
+ "schema": sch,
+ "type": rtype,
+ "language": None,
+ "parameters": None,
+ "return_type": None,
+ "definition": defn,
+ },
+ )
+
+ else:
+ raise ValueError(f"Unsupported object_type: {object_type}")
+
+ return {
+ "object_type": object_type,
+ "pattern": pat,
+ "schema": schema,
+ "table": table,
+ "detail_level": detail_level,
+ "count": len(results),
+ "results": results,
+ "truncated": len(results) == lim,
+ }
+
+
+class DbHubTools:
+ """
+ 与 DBHub 对齐的数据库工具封装:SQL 执行与元数据探索。
+
+ 方法 execute_sql 行为对齐 dbhub execute-sql.ts;
+ 方法 search_objects 对齐 DBHub search_objects。
+ execute_sql_all 为不截断行的全量拉取;
+ execute_sql_count_only 为仅统计行数;
+ 实现委托至模块内 ``_execute_sql`` / ``_execute_sql_all`` / ``_execute_sql_count_only`` / ``_search_objects``。
+ """
+
+ def execute_sql(
+ self,
+ sql: str,
+ *,
+ readonly: bool = True,
+ max_rows: int | None = None,
+ ) -> dict[str, object]:
+ """
+ 在已配置的业务库上执行 SQL(可多语句,分号分隔)。
+ :param sql: 待执行 SQL 字符串。多条语句用分号分隔;
+ :param readonly: 是否启用只读校验,默认 True。为 True 时,任一条语句不符合 G3SB 允许的首关键字
+ (如 select、with、explain 等,随方言而异)或含变更类关键字则抛出 ValueError
+ 为 False 时不做上述校验,可执行 DML/DDL(风险自负,勿用于不可信输入)。
+ :param max_rows: 对「最后一条返回结果集的语句」最多取多少行;None 时使用配置项 sql_max_rows,
+ 并在实现侧夹紧到 1~10000。前面的语句若产生结果集会被消费掉但不返回。
+ :return: 字典,含 rows(每行一个 dict,列名到值的映射)、count(等于 rows 长度,等同 DBHub rowCount)、
+ columns(列名列表)、truncated(是否因超过 max_rows 而截断)。
+ 最后一条语句无结果集时 rows、columns 为空,count 为 0。未配置 database_url、SQL 为空或执行失败时
+ 抛出 ValueError(或包装后的底层异常信息)。
+ """
+ return _execute_sql(sql, readonly=readonly, max_rows=max_rows)
+
+ def search_objects(
+ self,
+ object_type: ObjectType,
+ pattern: str = "%",
+ schema: str | None = None,
+ table: str | None = None,
+ detail_level: DetailLevel = "names",
+ limit: int = 100,
+ ) -> dict[str, Any]:
+ """
+ 按对象类型在已配置的业务库中探索 schema、表、列、索引、存储过程/函数等元数据(对齐 DBHub search_objects)。
+ 会过滤系统 schema(如information_schema、sys、pg_catalog 等),再按 object_type 与 pattern 做 SQL LIKE 风格匹配;
+ :param object_type: 要探索的对象类别。schema-模式名;table-数据表;column-列(可与 table、schema 联用
+ 限定单表);index-索引(同上);procedure-存储过程;function-函数(标量/表值等,视库而定)。
+ :param pattern: SQL LIKE 模式,默认 "%" 表示不过滤名称。"%" 匹配任意长度子串,"_" 匹配单个字符;对表名、
+ 列名、索引名、例程名等做大小写不敏感匹配。
+ :param schema: 限定在某个 schema 内查找;None 表示在多个非系统 schema 上依次查找。若给出具体名称,
+ 必须是库中已存在的 schema,否则抛 ValueError。使用参数 table 时必须同时指定 schema。
+ :param table: 仅在 object_type 为 column 或 index 时允许传入,与 schema 共同限定「只查这一张表或视图」上的列
+ 或索引;用于其它 object_type 时会抛 ValueError。名称会在该 schema 下与反射得到的表名、视图名做
+ 不区分大小写匹配(并支持 [Name] 写法),以兼容 SQL Server 等对目录大小写不敏感但 API 需真实写法的情况。
+ :param detail_level: 返回粒度。names-仅对象名及定位字段(如 schema、table);summary-增加简要元数据
+ (如表的 column_count、row_count,列的类型、可空等);full-表级返回列列表、索引列表等完整结构。
+ :param limit: 最多返回的结果条数,默认 100,有效范围 1~1000(传入值会被夹紧到该区间)。
+ :return: 包含 object_type、pattern、schema、table、detail_level、count、results、truncated 的字典。
+ 未配置 database_url、连接失败或内省失败时抛出 ValueError(或其它底层异常)。
+ """
+ return _search_objects(
+ object_type,
+ pattern=pattern,
+ schema=schema,
+ table=table,
+ detail_level=detail_level,
+ limit=limit,
+ )
+
+ def execute_sql_all(
+ self,
+ sql: str,
+ *,
+ readonly: bool = True,
+ ) -> dict[str, object]:
+ """
+ 在已配置的业务库上执行 SQL,对最后一条产生结果集的语句 fetchall 全量取行,不做行数上限截断。(4.SQL 执行接口)
+ :param sql: 待执行 SQL。多条语句用分号分隔;仅返回最后一条有结果集语句的 columns/rows,
+ 前面语句若产生结果集会被执行并消费掉但不返回。
+ :param readonly: 是否启用只读校验,默认 True。为 True 时任一条不符合 G3SB 允许的首关键字(如 select、with、explain 等,
+ 随方言而异)或含变更类关键字则抛出 ValueError;为 False 时不做该校验(风险自负,勿用于不可信输入)。
+ :return: 字典含 rows、count(等于 rows 长度)、columns、truncated。本路径下 truncated 恒为 False。
+ 最后一条语句无结果集时 rows、columns 为空,count 为 0。未配置 database_url、SQL 为空、校验失败或执行失败时
+ 抛出 ValueError(或带底层信息的包装异常)。
+ """
+ return _execute_sql_all(sql, readonly=readonly)
+
+ def execute_sql_count_only(self, sql: str) -> dict[str, int]:
+ """
+ 仅返回结果集语句的行数,不返回 rows/columns 明细;(3.SQL 检验接口)
+ :param sql: 待执行 SQL 字符串。多条语句用分号分隔;
+ :return: {count: int}
+ """
+ return _execute_sql_count_only(sql, readonly=True)
+
+
+# 默认工具实例;下方模块级 execute_sql / execute_sql_all / execute_sql_count_only / search_objects 均委托至此。
+dbhub_tools = DbHubTools()
+
+
+def search_objects(**kwargs: Any) -> dict[str, Any]:
+ return dbhub_tools.search_objects(**kwargs)
+
+
+def execute_sql(sql: str, **kwargs: Any) -> dict[str, object]:
+ return dbhub_tools.execute_sql(sql, **kwargs)
+
+
+def execute_sql_all(sql: str, **kwargs: Any) -> dict[str, object]:
+ return dbhub_tools.execute_sql_all(sql, **kwargs)
+
+
+def execute_sql_count_only(sql: str, **kwargs: Any) -> dict[str, int]:
+ return dbhub_tools.execute_sql_count_only(sql, **kwargs)
+
+
+def main() -> None:
+ """
+ 命令行演示:需已配置 database_url。SQL Server 示例默认 schema 为 dbo;
+ PostgreSQL 可将下面示例中的 schema 改为 public。
+ """
+ demos: list[tuple[str, dict[str, Any]]] = [
+ (
+ "1) 列出 schema(detail_level=names)",
+ {"object_type": "schema", "detail_level": "names", "limit": 5},
+ ),
+ (
+ "2) 列出某 schema 下的表名",
+ {"object_type": "table", "schema": "dbo", "detail_level": "names", "limit": 5},
+ ),
+ (
+ "3) 表名 LIKE 模糊匹配(%Account%)",
+ {
+ "object_type": "table",
+ "schema": "dbo",
+ "pattern": "%Account%",
+ "detail_level": "names",
+ "limit": 20,
+ },
+ ),
+ (
+ "4) 表级摘要(列数、行数 COUNT)",
+ {"object_type": "table", "schema": "dbo", "detail_level": "summary", "limit": 5},
+ ),
+ (
+ "6) 表级返回列列表、索引列表等完整结构",
+ {"object_type": "table", "schema": "dbo", "detail_level": "full", "limit": 5},
+ ),
+ (
+ "5) 某表或视图的列(支持大小写/视图;无则改 table 为库中真实对象名)",
+ {
+ "object_type": "column",
+ "schema": "dbo",
+ "table": "BCAccountAccruedCustodianFee",
+ "pattern": "%",
+ "detail_level": "names",
+ "limit": 5,
+ },
+ ),
+ ]
+
+ # for title, kwargs in demos:
+ # print(f"\n=== {title} ===")
+ # try:
+ # out = dbhub_tools.search_objects(**kwargs)
+ # print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
+ # except Exception as e: # noqa: BLE001
+ # log.warning(f"示例跳过或失败: {title} err={e}")
+ # print(f"(失败) {e}")
+ sql = """
+ SELECT
+ b.BrokerID,
+ m.Name AS BrokerName,
+ b.CurrencyID,
+ b.Settled AS OwedToUsAmount,
+ ABS(b.Settled) AS OutstandingAmount
+FROM BCBrokerCash b
+INNER JOIN MCBroker m ON b.BrokerID = m.BrokerID
+WHERE b.Settled < 0
+ORDER BY ABS(b.Settled) DESC;
+
+ """
+ out = dbhub_tools.execute_sql(sql)
+ print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
+ sql = """
+ SELECT
+ b.BrokerID,
+ m.Name AS BrokerName,
+ b.CurrencyID,
+ b.Settled AS OwedToUsAmount,
+ ABS(b.Settled) AS OutstandingAmount
+FROM BCBrokerCash b
+INNER JOIN MCBroker m ON b.BrokerID = m.BrokerID
+WHERE b.Settled < 0
+ORDER BY ABS(b.Settled) DESC;
+
+ """
+ out = dbhub_tools.execute_sql_count_only(sql)
+ print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
+ sql = """
+ SELECT
+ b.BrokerID,
+ m.Name AS BrokerName,
+ b.CurrencyID,
+ b.Settled AS OwedToUsAmount,
+ ABS(b.Settled) AS OutstandingAmount
+FROM BCBrokerCash b
+INNER JOIN MCBroker m ON b.BrokerID = m.BrokerID
+WHERE b.Settled < 0
+ORDER BY ABS(b.Settled) DESC;
+
+ """
+ out = dbhub_tools.execute_sql_all(sql)
+ print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/db/engine.py b/backend/db/engine.py
new file mode 100644
index 0000000..defb173
--- /dev/null
+++ b/backend/db/engine.py
@@ -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
diff --git a/llm/__init__.py b/backend/llm/__init__.py
similarity index 100%
rename from llm/__init__.py
rename to backend/llm/__init__.py
diff --git a/llm/__pycache__/__init__.cpython-312.pyc b/backend/llm/__pycache__/__init__.cpython-312.pyc
similarity index 100%
rename from llm/__pycache__/__init__.cpython-312.pyc
rename to backend/llm/__pycache__/__init__.cpython-312.pyc
diff --git a/backend/llm/__pycache__/deepseek_client.cpython-312.pyc b/backend/llm/__pycache__/deepseek_client.cpython-312.pyc
new file mode 100644
index 0000000..08a1894
Binary files /dev/null and b/backend/llm/__pycache__/deepseek_client.cpython-312.pyc differ
diff --git a/llm/deepseek_client.py b/backend/llm/deepseek_client.py
similarity index 77%
rename from llm/deepseek_client.py
rename to backend/llm/deepseek_client.py
index 5f3069c..1862f0e 100644
--- a/llm/deepseek_client.py
+++ b/backend/llm/deepseek_client.py
@@ -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:
"""
diff --git a/main.py b/backend/main.py
similarity index 82%
rename from main.py
rename to backend/main.py
index 5ae9cb7..4aca88b 100644
--- a/main.py
+++ b/backend/main.py
@@ -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()
+ interactive_mode(orchestrator, args.dialect)
if __name__ == "__main__":
diff --git a/backend/nl_lite_store.py b/backend/nl_lite_store.py
new file mode 100644
index 0000000..0e9f649
--- /dev/null
+++ b/backend/nl_lite_store.py
@@ -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()
diff --git a/schema/__init__.py b/backend/schema/__init__.py
similarity index 100%
rename from schema/__init__.py
rename to backend/schema/__init__.py
diff --git a/schema/__pycache__/__init__.cpython-312.pyc b/backend/schema/__pycache__/__init__.cpython-312.pyc
similarity index 100%
rename from schema/__pycache__/__init__.cpython-312.pyc
rename to backend/schema/__pycache__/__init__.cpython-312.pyc
diff --git a/schema/__pycache__/indexer.cpython-312.pyc b/backend/schema/__pycache__/indexer.cpython-312.pyc
similarity index 100%
rename from schema/__pycache__/indexer.cpython-312.pyc
rename to backend/schema/__pycache__/indexer.cpython-312.pyc
diff --git a/schema/__pycache__/loader.cpython-312.pyc b/backend/schema/__pycache__/loader.cpython-312.pyc
similarity index 100%
rename from schema/__pycache__/loader.cpython-312.pyc
rename to backend/schema/__pycache__/loader.cpython-312.pyc
diff --git a/schema/__pycache__/manager.cpython-312.pyc b/backend/schema/__pycache__/manager.cpython-312.pyc
similarity index 100%
rename from schema/__pycache__/manager.cpython-312.pyc
rename to backend/schema/__pycache__/manager.cpython-312.pyc
diff --git a/schema/__pycache__/models.cpython-312.pyc b/backend/schema/__pycache__/models.cpython-312.pyc
similarity index 100%
rename from schema/__pycache__/models.cpython-312.pyc
rename to backend/schema/__pycache__/models.cpython-312.pyc
diff --git a/schema/indexer.py b/backend/schema/indexer.py
similarity index 100%
rename from schema/indexer.py
rename to backend/schema/indexer.py
diff --git a/schema/loader.py b/backend/schema/loader.py
similarity index 100%
rename from schema/loader.py
rename to backend/schema/loader.py
diff --git a/schema/manager.py b/backend/schema/manager.py
similarity index 100%
rename from schema/manager.py
rename to backend/schema/manager.py
diff --git a/schema/models.py b/backend/schema/models.py
similarity index 100%
rename from schema/models.py
rename to backend/schema/models.py
diff --git a/utils/__init__.py b/backend/utils/__init__.py
similarity index 100%
rename from utils/__init__.py
rename to backend/utils/__init__.py
diff --git a/utils/__pycache__/__init__.cpython-312.pyc b/backend/utils/__pycache__/__init__.cpython-312.pyc
similarity index 100%
rename from utils/__pycache__/__init__.cpython-312.pyc
rename to backend/utils/__pycache__/__init__.cpython-312.pyc
diff --git a/backend/utils/__pycache__/dialog_classifier.cpython-312.pyc b/backend/utils/__pycache__/dialog_classifier.cpython-312.pyc
new file mode 100644
index 0000000..2941fa4
Binary files /dev/null and b/backend/utils/__pycache__/dialog_classifier.cpython-312.pyc differ
diff --git a/utils/__pycache__/embedding.cpython-312.pyc b/backend/utils/__pycache__/embedding.cpython-312.pyc
similarity index 100%
rename from utils/__pycache__/embedding.cpython-312.pyc
rename to backend/utils/__pycache__/embedding.cpython-312.pyc
diff --git a/utils/__pycache__/fewshot_selector.cpython-312.pyc b/backend/utils/__pycache__/fewshot_selector.cpython-312.pyc
similarity index 93%
rename from utils/__pycache__/fewshot_selector.cpython-312.pyc
rename to backend/utils/__pycache__/fewshot_selector.cpython-312.pyc
index 1499b00..4ce799d 100644
Binary files a/utils/__pycache__/fewshot_selector.cpython-312.pyc and b/backend/utils/__pycache__/fewshot_selector.cpython-312.pyc differ
diff --git a/backend/utils/__pycache__/question_locale.cpython-312.pyc b/backend/utils/__pycache__/question_locale.cpython-312.pyc
new file mode 100644
index 0000000..830c7fd
Binary files /dev/null and b/backend/utils/__pycache__/question_locale.cpython-312.pyc differ
diff --git a/utils/__pycache__/sql_parser.cpython-312.pyc b/backend/utils/__pycache__/sql_parser.cpython-312.pyc
similarity index 100%
rename from utils/__pycache__/sql_parser.cpython-312.pyc
rename to backend/utils/__pycache__/sql_parser.cpython-312.pyc
diff --git a/utils/__pycache__/validators.cpython-312.pyc b/backend/utils/__pycache__/validators.cpython-312.pyc
similarity index 50%
rename from utils/__pycache__/validators.cpython-312.pyc
rename to backend/utils/__pycache__/validators.cpython-312.pyc
index 1ecc9c6..f77392b 100644
Binary files a/utils/__pycache__/validators.cpython-312.pyc and b/backend/utils/__pycache__/validators.cpython-312.pyc differ
diff --git a/backend/utils/dialog_classifier.py b/backend/utils/dialog_classifier.py
new file mode 100644
index 0000000..6d333b1
--- /dev/null
+++ b/backend/utils/dialog_classifier.py
@@ -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)
diff --git a/utils/embedding.py b/backend/utils/embedding.py
similarity index 100%
rename from utils/embedding.py
rename to backend/utils/embedding.py
diff --git a/utils/fewshot_selector.py b/backend/utils/fewshot_selector.py
similarity index 98%
rename from utils/fewshot_selector.py
rename to backend/utils/fewshot_selector.py
index 229f6f1..129ba54 100644
--- a/utils/fewshot_selector.py
+++ b/backend/utils/fewshot_selector.py
@@ -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))
diff --git a/backend/utils/question_locale.py b/backend/utils/question_locale.py
new file mode 100644
index 0000000..ff5e6fc
--- /dev/null
+++ b/backend/utils/question_locale.py
@@ -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
diff --git a/utils/sql_parser.py b/backend/utils/sql_parser.py
similarity index 100%
rename from utils/sql_parser.py
rename to backend/utils/sql_parser.py
diff --git a/utils/validators.py b/backend/utils/validators.py
similarity index 77%
rename from utils/validators.py
rename to backend/utils/validators.py
index 75d93e1..f683a2b 100644
--- a/utils/validators.py
+++ b/backend/utils/validators.py
@@ -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",
diff --git a/llm/__pycache__/deepseek_client.cpython-312.pyc b/llm/__pycache__/deepseek_client.cpython-312.pyc
deleted file mode 100644
index e1b90d5..0000000
Binary files a/llm/__pycache__/deepseek_client.cpython-312.pyc and /dev/null differ
diff --git a/requirements.txt b/requirements.txt
index d507cf5..d8298c7 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -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
@@ -27,4 +29,8 @@ tenacity>=8.0.0
structlog>=23.0.0
tqdm>=4.65.0
jieba>=0.42.0
-numpy>=1.24.0,<2
\ No newline at end of file
+numpy>=1.24.0,<2
+
+# API框架
+fastapi>=0.100.0
+uvicorn>=0.20.0
\ No newline at end of file
diff --git a/scripts/integrate_fewshot.py b/scripts/integrate_fewshot.py
index aa06480..2e3352e 100644
--- a/scripts/integrate_fewshot.py
+++ b/scripts/integrate_fewshot.py
@@ -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"
```
diff --git a/scripts/patch_fewshot.py b/scripts/patch_fewshot.py
index 3939c93..aae3f81 100644
--- a/scripts/patch_fewshot.py
+++ b/scripts/patch_fewshot.py
@@ -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. 调整参数:")
diff --git a/tech_architecture.md b/tech_architecture.md
new file mode 100644
index 0000000..d6a2126
--- /dev/null
+++ b/tech_architecture.md
@@ -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 生成需求。
diff --git a/tools/__init__.py b/tools/__init__.py
new file mode 100644
index 0000000..62c2373
--- /dev/null
+++ b/tools/__init__.py
@@ -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",
+]
diff --git a/tools/__pycache__/__init__.cpython-311.pyc b/tools/__pycache__/__init__.cpython-311.pyc
new file mode 100644
index 0000000..2b48201
Binary files /dev/null and b/tools/__pycache__/__init__.cpython-311.pyc differ
diff --git a/tools/__pycache__/__init__.cpython-312.pyc b/tools/__pycache__/__init__.cpython-312.pyc
new file mode 100644
index 0000000..c7ca724
Binary files /dev/null and b/tools/__pycache__/__init__.cpython-312.pyc differ
diff --git a/tools/__pycache__/dbhub_allowed_keywords.cpython-311.pyc b/tools/__pycache__/dbhub_allowed_keywords.cpython-311.pyc
new file mode 100644
index 0000000..a42f15a
Binary files /dev/null and b/tools/__pycache__/dbhub_allowed_keywords.cpython-311.pyc differ
diff --git a/tools/__pycache__/dbhub_execute_sql.cpython-311.pyc b/tools/__pycache__/dbhub_execute_sql.cpython-311.pyc
new file mode 100644
index 0000000..36aefcc
Binary files /dev/null and b/tools/__pycache__/dbhub_execute_sql.cpython-311.pyc differ
diff --git a/tools/__pycache__/dbhub_sql_parser.cpython-311.pyc b/tools/__pycache__/dbhub_sql_parser.cpython-311.pyc
new file mode 100644
index 0000000..ebd2faa
Binary files /dev/null and b/tools/__pycache__/dbhub_sql_parser.cpython-311.pyc differ
diff --git a/tools/__pycache__/dbhub_tools.cpython-311.pyc b/tools/__pycache__/dbhub_tools.cpython-311.pyc
new file mode 100644
index 0000000..5bd1357
Binary files /dev/null and b/tools/__pycache__/dbhub_tools.cpython-311.pyc differ
diff --git a/tools/__pycache__/dbhub_tools.cpython-312.pyc b/tools/__pycache__/dbhub_tools.cpython-312.pyc
new file mode 100644
index 0000000..d668696
Binary files /dev/null and b/tools/__pycache__/dbhub_tools.cpython-312.pyc differ
diff --git a/tools/dbhub_allowed_keywords.py b/tools/dbhub_allowed_keywords.py
new file mode 100644
index 0000000..bad66dd
--- /dev/null
+++ b/tools/dbhub_allowed_keywords.py
@@ -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]
diff --git a/tools/dbhub_sql_parser.py b/tools/dbhub_sql_parser.py
new file mode 100644
index 0000000..cdf8b3a
--- /dev/null
+++ b/tools/dbhub_sql_parser.py
@@ -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
diff --git a/tools/dbhub_tools.py b/tools/dbhub_tools.py
new file mode 100644
index 0000000..d264091
--- /dev/null
+++ b/tools/dbhub_tools.py
@@ -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",
+]