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", +]