Refactor api_server.py to import environment setup and schema loading from bootstrap.py, enhancing modularity. Introduce a new function in Text2SQLOrchestrator to prioritize VCUserAccessibleFunction in table selection, improving SQL generation accuracy. Update validation logic to enforce restrictions on CJK characters in SQL string literals, ensuring compliance with business rules. Enhance prompts to clarify SQL generation constraints regarding date conditions and CJK usage.

This commit is contained in:
陈辅元
2026-04-17 15:15:01 +08:00
parent bcb9a205fa
commit 3891aae7fe
19 changed files with 707 additions and 306 deletions
Binary file not shown.
Binary file not shown.
Binary file not shown.
+41 -6
View File
@@ -302,6 +302,28 @@ class Text2SQLOrchestrator:
logger.info("对手方/经纪商问题:优先纳入 %s,调整后选表:%s", present, merged)
return merged[:max_tables]
_VC_USER_ACCESSIBLE_FUNCTION = "VCUserAccessibleFunction"
def _prioritize_vc_user_accessible_function(
self, relevant_tables: List[str], max_tables: int = 5
) -> List[str]:
"""
若 Schema 中存在 VCUserAccessibleFunction,则置于选表列表最前,便于模型先根据
FunctionID / DatabaseView 等列定位业务视图,再关联其余表生成 SQL。
"""
vc = self._VC_USER_ACCESSIBLE_FUNCTION
if not self.schema_manager.get_table(vc):
return relevant_tables[:max_tables]
rest = [t for t in relevant_tables if t != vc]
merged = [vc] + rest
logger.info(
"已优先纳入目录视图 %s(置于选表前列),当前选表:%s",
vc,
merged[:max_tables],
)
return merged[:max_tables]
def _expand_relations(self, table_names: List[str]) -> List[str]:
"""
外键扩展:自动添加关联表
@@ -413,8 +435,9 @@ class Text2SQLOrchestrator:
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
"\n**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。"
"\n在 `WHERE`/`HAVING`/`JOIN ON` 及 `CASE/WHEN` 的**条件**中,**禁止**用含中日韩的 `'`/`N'…'` 字面量与代码型列比对;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN(勿写 `= '过户费'`)。`SELECT` 或 `CASE … THEN/ELSE` 的展示用中文标签允许。"
"若用户问题与对话上文均未要求按时间筛选,不得在 WHERE 中擅自添加日期列条件(与系统提示 1b 一致)。"
)
if validation_feedback:
@@ -577,8 +600,9 @@ class Text2SQLOrchestrator:
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
"\n**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。"
"\n在 `WHERE`/`HAVING`/`JOIN ON` 及 `CASE/WHEN` 的**条件**中,**禁止**用含中日韩的 `'`/`N'…'` 字面量与代码型列比对;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN(勿写 `= '过户费'`)。`SELECT` 或 `CASE … THEN/ELSE` 的展示用中文标签允许。"
"若用户问题与对话上文均未要求按时间筛选,不得在 WHERE 中擅自添加日期列条件(与系统提示 1b 一致)。"
)
if validation_feedback:
@@ -659,9 +683,9 @@ class Text2SQLOrchestrator:
if not danger_ok:
errors.extend(danger_errors)
# T-SQL:禁止中文等业务词出现在字符串字面量(如 FeeNatureID = '过户费')
# T-SQL:禁止中文出现在 WHERE/HAVING/ON/CASE 条件等比对语境(展示用 CASE THEN/ELSE 允许)
if dialect == "tsql":
cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql)
cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql, dialect=dialect)
if not cjk_ok:
errors.extend(cjk_errors)
@@ -845,6 +869,9 @@ class Text2SQLOrchestrator:
relevant_tables = self._prioritize_broker_tables(
linker_question, relevant_tables
)
relevant_tables = self._prioritize_vc_user_accessible_function(
relevant_tables
)
# 1.3 外键扩展
expanded_tables = self._expand_relations(relevant_tables)
@@ -856,6 +883,14 @@ class Text2SQLOrchestrator:
include_columns=True,
max_columns_per_table=20
)
if self._VC_USER_ACCESSIBLE_FUNCTION in expanded_tables:
filtered_schema_str += (
"\n\n【选表提示】已包含视图 "
+ self._VC_USER_ACCESSIBLE_FUNCTION
+ "(列含 UserID、FunctionID、Category、Name、DatabaseView)。"
"生成 SQL 时可先通过该视图用 DatabaseView / FunctionID 等定位目标业务视图或功能,"
"再与 Schema 中其余表做 JOIN 或子查询;若问题已明确具体表名,可直接查询该表。"
)
logger.info(f" 选中表:{relevant_tables},扩展后:{expanded_tables}")
else:
# 重试:在子 Schema 中并入「上次失败 SQL」实际引用到的表,并对齐程序校验与生成上下文
+215
View File
@@ -0,0 +1,215 @@
"""
Text2SQL 共享启动逻辑:环境检查、Schema 加载、Orchestrator 构造。
供 `main` CLI 与 `api_server` 复用,避免 API 层依赖 CLI 入口模块。
"""
from __future__ import annotations
import logging
import os
import sys
from pathlib import Path
from typing import Any, Optional
from dotenv import load_dotenv
logger = logging.getLogger(__name__)
def _repo_root() -> Path:
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
return Path(__file__).resolve().parent.parent
def load_project_env() -> None:
"""加载项目根目录 .env,供后续 os.getenv 使用。"""
load_dotenv(_repo_root() / ".env")
def resolve_sql_dialect(name: str) -> str:
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
n = (name or "sqlserver").lower().strip()
if n in ("sqlserver", "mssql"):
return "tsql"
return n
def setup_environment() -> bool:
"""环境检查"""
load_project_env()
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
oa_key_ok = oa_key and not (
oa_key.startswith("http://") or oa_key.startswith("https://")
)
if ms_key:
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
elif oa_key_ok:
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
logger.warning("未设置 OPENAI_EMBEDDING_MODEL(远程 Embedding 模型名)")
return False
else:
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
logger.warning(
"未设置 MODELSCOPE_API_KEY、"
"OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
)
return False
base = (
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
).strip()
if not base:
logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
return False
if not os.getenv("DASHSCOPE_MODEL", "").strip():
logger.warning("未设置 DASHSCOPE_MODEL")
return False
# 检查Schema文件(支持相对路径和绝对路径)
schema_path_str = os.getenv(
"SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json"
)
schema_path = Path(schema_path_str)
# 如果是相对路径,尝试从多个位置查找
if not schema_path.is_absolute():
# 尝试1: PyInstaller 临时目录(单文件模式)
if getattr(sys, "frozen", False) and hasattr(sys, "_MEIPASS"):
meipass_schema = Path(sys._MEIPASS) / schema_path_str
if meipass_schema.exists():
schema_path = meipass_schema
# 尝试2: 当前工作目录
if not schema_path.exists():
schema_path = Path.cwd() / schema_path_str
# 尝试3: 仓库根目录(本文件位于 backend/)
if not schema_path.exists():
script_dir = Path(__file__).resolve().parent.parent
schema_path = script_dir / schema_path_str
# 尝试4: 可执行文件所在目录
if not schema_path.exists() and getattr(sys, "frozen", False):
exe_dir = Path(sys.executable).parent
schema_path = exe_dir / schema_path_str
if not schema_path.exists():
logger.warning(f"Schema文件不存在: {schema_path}")
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
return False
# 检查 LLM Key(DeepSeek / OpenAI 可切换)
llm_sc = (os.getenv("LLM_SERVICE_CODE") or "").strip().lower()
if llm_sc and llm_sc not in ("deepseek", "openai"):
logger.warning("未知 LLM_SERVICE_CODE=%r(仅支持 deepseek/openai)", llm_sc)
return False
if llm_sc == "openai":
if not (os.getenv("OPENAI_API_KEY") or "").strip():
logger.warning("LLM_SERVICE_CODE=openai 但 OPENAI_API_KEY 未设置")
return False
else:
if not (os.getenv("DEEPSEEK_API_KEY") or "").strip():
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置(默认 LLM_SERVICE_CODE=deepseek)")
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
return False
logger.info(f"[OK] 环境检查通过")
logger.info(f" - Schema: {schema_path}")
logger.info(
" - LLM: %s",
(llm_sc or "deepseek(auto)"),
)
return True
def _default_g3sb_meta_path(structure_path: str) -> Optional[str]:
"""若存在与 table_structure 同名的 table_meta 文件则返回其路径。"""
p = Path(structure_path)
if "table_structure" not in p.name:
return None
cand = p.parent / p.name.replace("table_structure", "table_meta")
return str(cand) if cand.is_file() else None
def load_schema(schema_path: str, schema_meta_path: Optional[str] = None):
"""加载 Schema;G3SB structure JSON 会自动尝试配对 table_meta(可用 --schema-meta 指定)。"""
from schema.manager import SchemaManager
meta = (
schema_meta_path
if schema_meta_path is not None
else _default_g3sb_meta_path(schema_path)
)
logger.info(f"加载Schema: {schema_path}")
if meta:
logger.info(f" 表注释(meta): {meta}")
schema_mgr = SchemaManager.load_from_json(
schema_path, g3sb_meta_path=meta
)
stats = schema_mgr.get_statistics()
logger.info(
f"[OK] Schema加载完成: {stats['database']}, "
f"共{stats['total_tables']}张表, {stats['total_columns']}个字段"
)
return schema_mgr
def create_orchestrator(schema_mgr: Any, args: Any):
"""创建编排器"""
from agents.orchestrator import Text2SQLOrchestrator
from llm.router import create_llm_client, resolve_llm_service_code
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
sc = resolve_llm_service_code()
if sc == "deepseek":
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
if not api_key:
raise ValueError(
"未配置 DeepSeek API Key:请在 .env 中设置 DEEPSEEK_API_KEY,"
"或使用命令行参数 --api-key"
)
base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
cfg = DeepSeekConfig(
api_key=api_key,
base_url=base_url,
model_name=args.model or "deepseek-chat",
temperature=args.temperature,
max_tokens=args.max_tokens,
)
llm_client = create_llm_client("deepseek", **cfg.__dict__)
else:
# openai:完全由 OPENAI_* 决定;同时沿用 temperature/max_tokens 作为默认值覆盖
llm_client = create_llm_client(
"openai",
temperature=args.temperature,
max_tokens=args.max_tokens,
)
orchestrator = Text2SQLOrchestrator(
schema_manager=schema_mgr,
llm_client=llm_client,
vector_db_path=args.vector_db,
max_retry=args.max_retry,
use_vector_search=not args.no_vector_search,
# Few-shot配置
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
Binary file not shown.
+35 -22
View File
@@ -55,8 +55,9 @@ 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 维表** 用其键列或英文名列过滤,**不得**用中文当字面量。
1b. **无时间表述则不加日期条件(强制)**:若用户问题及对话上文**均未**出现任何可映射为**按时间筛选**的表述——包括但不限于:具体日历日期、年月/季度区间、「今天/昨日/本周/本月/本年/本季度」「最近N天/过去一周/过去一年」等相对时间——则 **不得**在 `WHERE`/`HAVING` 中**新增**对日期/时间类型列的过滤(例如 `TradeDate >= '...'`、`BETWEEN ... AND ...`、与 `CAST(GETDATE() AS DATE)` / `DATEADD` 结合的日期条件)。**禁止**以「防止结果集过大」「报表通常只看近期」「默认只查当年」等理由擅加日期窗。仅当用户**明确**提出时间要求、或问题语义**显式**指向某时段(如「2024年1月的销售额」「今天的成交」)时,才写对应日期条件;完全未提时间时,查询在日期维度上可为全表/全历史(仅受问题中**已出现**的非时间条件约束)。
2. **禁止虚构值与占位符**:不得使用 `'[日期]'`、`TODO`、`xxx`、空泛占位等冒充具体字面量。若用户未给出具体日期、代码或 ID:**日期类**仅当问题里**已经**出现可映射的时间表述时,才按上文与「常见时间推断指南」写出具体区间或 `GETDATE()` 条件;若全文无任何时间表述,**不得**为凑条件而编造日期过滤(与 **1b** 一致)。**非日期类**(状态码、ID 等)仍可根据问题上下文与 Schema 枚举填写合理值,**不要**留空或写占位符。
2b. **禁止在过滤/比对条件中使用中文(CJK)字面量**:在 `WHERE`/`HAVING`/`JOIN … ON` 以及 `CASE WHEN` 的**条件部分**,用户问题里的中文业务词(如「过户费」「未结算」「活跃」)**禁止**写成 `'…中文…'` 或 `N'…中文…'` 去和代码型列(如 `FeeNatureID`、`SettleStatus`、`State`)比较。必须根据 **Schema 字段注释** 写成库内真实**代码/单字母/数字**(如 `State = 'A'`、`SettleStatus = 'U'`);若业务词对应维表或码表,应 **JOIN 维表** 用其键列或英文名列过滤。**允许**在 `SELECT` 列表达式或 `CASE … THEN … ELSE …` 的**展示结果**中使用中文标签字符串(如状态说明),此类不属于「与代码列比对」。
3. **输出版式与别名风格(统一规范)**:除遵守目标方言语法外,SQL **排版与命名**须与下方「标准版式范例」一致:
- **关键字**:`SELECT`、`FROM`、`JOIN`/`LEFT JOIN`、`ON`、`WHERE`、`AND`、`GROUP BY`、`ORDER BY`、`HAVING` 等使用**大写**。
- **换行与缩进**:`SELECT` 后换行;每个输出列**独占一行**,行首 **4 个空格**,列表达式之间用**行尾逗号**分隔(最后一列无逗号)。
@@ -86,6 +87,7 @@ SQL_GENERATOR_SYSTEM = """你是一个精通SQL的数据库专家,有10年以
- 负债/负数含义:许多余额字段负值表示负债(如LoanBalance)
**常见时间/状态推断指南**(需结合Schema字段注释):
- **前提**:下列日期规则**仅当**用户问题或对话中**已出现**对应时间表述时适用;若完全未提时间,**不要**套用下列规则去加日期条件(见 **1b**)。
- "2024年1月" → `WHERE date_col >= '2024-01-01' AND date_col < '2024-02-01'`
- "今天" / "当日" → `WHERE date_col >= CAST(GETDATE() AS DATE) AND date_col < DATEADD(DAY,1,CAST(GETDATE() AS DATE))`
- "最近N天" → `WHERE date_col >= DATEADD(DAY, -N, CAST(GETDATE() AS DATE))`
@@ -103,7 +105,7 @@ SQL_GENERATOR_SYSTEM = """你是一个精通SQL的数据库专家,有10年以
用户问题(示例):按对手方列出截至 2026-04-02 的所有未结算交易。
(日期规则:用户问题里**已写出具体日期**时,在 WHERE 中写入相同字面量,例如 `CashSettleDate <= '2026-04-02'`;**未给出具体日期**但包含时间范围描述(如"2024年1月")时,应合理推断为日期区间,例如 `OrderDate >= '2024-01-01' AND OrderDate < '2024-02-01'`,禁止使用 `'[日期]'` 等占位符。)
(日期规则:用户问题里**已写出具体日期**时,在 WHERE 中写入相同字面量,例如 `CashSettleDate <= '2026-04-02'`;**未给出具体日期**但包含时间范围描述(如"2024年1月")时,应合理推断为日期区间,例如 `OrderDate >= '2024-01-01' AND OrderDate < '2024-02-01'`,禁止使用 `'[日期]'` 等占位符。**若用户完全未提及任何时间与时段**,则不要添加日期列条件,勿因本范例含日期而照抄日期过滤。)
SQL(表名、字段名须与当前 Schema 一致;**以下版式、别名、JOIN/WHERE/GROUP BY/ORDER BY 结构为强制模板**):
@@ -130,30 +132,41 @@ ORDER BY TotalUnsettledAmount DESC;
**示例(版式与范例一致)**:
示例1 - 单表查询:
问题:查询账户ID为'ACC001'的账户余额
Schema: MCAccount(AccountID, Name, AvailableBalance, MarketValue, MarginValue)
问题:查询账户ID为'M050013'的账户余额
Schema: MCAccount(AccountID, AvailableBalance, AssetBalance, LiabilityBalance, MaximumAvailableBalance)
SQL:
SELECT
AccountID,
Name AS AccountName,
RTRIM(AccountID) AS AccountID,
AvailableBalance,
MarketValue,
MarginValue
FROM MCAccount
WHERE AccountID = 'ACC001';
AssetBalance,
LiabilityBalance,
MaximumAvailableBalance
FROM dbo.MCAccount
WHERE RTRIM(AccountID) = N'M050013';
示例2 - 多表 INNER JOIN:
问题:查询账户'ACC001'持有的所有股票及数量
Schema: MCAccount(AccountID), MCAccountInstrument(AccountID, MarketID, InstrumentID, Settled)
问题:查询账户'M050013'持有的所有股票及数量
Schema: BCAccountInstrument(AccountID, MarketID, InstrumentID, DailyOpenLedgerQuantity, DailyOpenSettledQuantity), MCInstrument(MarketID, InstrumentID, Name, InstrumentTypeID), MCInstrumentType(InstrumentTypeID, Name)
SQL:
SELECT
a.AccountID,
i.MarketID,
i.InstrumentID AS InstrumentCode,
i.Settled AS HoldingQty
FROM MCAccount a
JOIN MCAccountInstrument i ON a.AccountID = i.AccountID
WHERE a.AccountID = 'ACC001';
RTRIM(bai.AccountID) AS AccountID,
RTRIM(bai.MarketID) AS MarketID,
RTRIM(bai.InstrumentID) AS InstrumentID,
RTRIM(mi.Name) AS InstrumentName,
RTRIM(mi.InstrumentTypeID) AS InstrumentTypeID,
it.Name AS InstrumentTypeName,
bai.DailyOpenLedgerQuantity AS LedgerQuantity,
bai.DailyOpenSettledQuantity AS SettledQuantity
FROM dbo.BCAccountInstrument AS bai
INNER JOIN dbo.MCInstrument AS mi
ON bai.MarketID = mi.MarketID
AND bai.InstrumentID = mi.InstrumentID
LEFT JOIN dbo.MCInstrumentType AS it
ON mi.InstrumentTypeID = it.InstrumentTypeID
WHERE RTRIM(bai.AccountID) = N'M050013'
AND bai.DailyOpenLedgerQuantity <> 0
ORDER BY bai.MarketID, bai.InstrumentID;
示例3 - 时间范围推断(关键!):
问题:查询2024年1月的总销售额
@@ -223,7 +236,7 @@ SQL_GENERATOR_USER = """Schema信息:
数据库方言:{dialect}
请生成**有用 SQL**(见系统提示定义):必须与「业务级黄金范例」**同构**——大写关键字、多行缩进版式、PascalCase 别名、该展示对手方/账户等名称时须 LEFT JOIN 维表;禁止输出挤成一行的「极简 SQL」。"""
请生成**有用 SQL**(见系统提示定义):必须与「业务级黄金范例」**同构**——大写关键字、多行缩进版式、PascalCase 别名、该展示对手方/账户等名称时须 LEFT JOIN 维表;禁止输出挤成一行的「极简 SQL」。**若当前问题未要求按时间筛选,不得在 WHERE 中擅自添加日期条件(系统提示 1b)。**"""
# ========== Few-shot 黄金 SQL 条件适配(Chroma 库内为已校验正确答案)==========
@@ -233,7 +246,7 @@ GOLDEN_SQL_ADAPT_SYSTEM = """你是精通 Microsoft SQL Server (T-SQL) 的数据
**你必须遵守**:
1. **以标准答案为主干**:优先保留其 `FROM`/`JOIN`/`ON`、主 `SELECT` 列清单与聚合/分组逻辑;**不要随意更换主表、不要拆掉必要 JOIN**,除非当前 Schema 片段中已不存在该表(此时在 Schema 内做最小替换并说明等价关系仅在脑中完成)。
2. **只改「条件类」内容**:重点调整 `WHERE`/`HAVING`/`ORDER BY`/`TOP` 中的字面量、日期区间、状态码、账户/合约/代码等过滤;将用户问题中的时间范围、业务对象、筛选口径反映到这些条件中。
2. **只改「条件类」内容**:重点调整 `WHERE`/`HAVING`/`ORDER BY`/`TOP` 中的字面量、日期区间、状态码、账户/合约/代码等过滤;将用户问题中的时间范围、业务对象、筛选口径反映到这些条件中。**若当前用户问题相较范例问题「少了」时间要求**(完全未提时间或时段),应**去掉**标准答案中仅因范例日期而存在的日期过滤,**禁止**保留与当前问题无关的日期条件(与 SQL_GENERATOR 的 **1b** 一致)。
3. **Schema 绝对优先**:表名、列名必须来自下方「当前 Schema 片段」;禁止臆造字段。若标准答案中某列在片段中不存在,按片段改写为合法列。
4. **T-SQL 与版式**:与常规生成一致——关键字大写、多行缩进、`WHERE` 续行以 `AND` 开头、需要时 PascalCase 英文别名;禁止 MySQL 反引号与 `CURDATE()` 等。
5. **禁止在字符串字面量中写中日韩文字**去匹配代码列;须用 Schema 注释中的代码或 JOIN 维表(与系统提示 SQL_GENERATOR 一致)。
+8 -200
View File
@@ -10,10 +10,6 @@ import sys
import argparse
import logging
from pathlib import Path
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:
@@ -27,201 +23,13 @@ logging.basicConfig(
)
logger = logging.getLogger(__name__)
def _repo_root() -> Path:
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
return Path(__file__).resolve().parent.parent
def _load_project_env():
"""加载项目根目录 .env,供后续 os.getenv 使用。"""
load_dotenv(_repo_root() / ".env")
def resolve_sql_dialect(name: str) -> str:
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
n = (name or "sqlserver").lower().strip()
if n in ("sqlserver", "mssql"):
return "tsql"
return n
def setup_environment():
"""环境检查"""
_load_project_env()
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
oa_key_ok = oa_key and not (
oa_key.startswith("http://") or oa_key.startswith("https://")
)
if ms_key:
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
elif oa_key_ok:
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
logger.warning("未设置 OPENAI_EMBEDDING_MODEL(远程 Embedding 模型名)")
return False
else:
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
logger.warning(
"未设置 MODELSCOPE_API_KEY、"
"OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
)
return False
base = (
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
).strip()
if not base:
logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
return False
if not os.getenv("DASHSCOPE_MODEL", "").strip():
logger.warning("未设置 DASHSCOPE_MODEL")
return False
# 检查Schema文件(支持相对路径和绝对路径)
schema_path_str = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
schema_path = Path(schema_path_str)
# 如果是相对路径,尝试从多个位置查找
if not schema_path.is_absolute():
# 尝试1: PyInstaller 临时目录(单文件模式)
import sys
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
meipass_schema = Path(sys._MEIPASS) / schema_path_str
if meipass_schema.exists():
schema_path = meipass_schema
# 尝试2: 当前工作目录
if not schema_path.exists():
schema_path = Path.cwd() / schema_path_str
# 尝试3: 脚本所在目录
if not schema_path.exists():
script_dir = Path(__file__).resolve().parent.parent
schema_path = script_dir / schema_path_str
# 尝试4: 可执行文件所在目录
if not schema_path.exists() and getattr(sys, 'frozen', False):
exe_dir = Path(sys.executable).parent
schema_path = exe_dir / schema_path_str
if not schema_path.exists():
logger.warning(f"Schema文件不存在: {schema_path}")
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
return False
# 检查 LLM Key(DeepSeek / OpenAI 可切换)
llm_sc = (os.getenv("LLM_SERVICE_CODE") or "").strip().lower()
if llm_sc and llm_sc not in ("deepseek", "openai"):
logger.warning("未知 LLM_SERVICE_CODE=%r(仅支持 deepseek/openai)", llm_sc)
return False
if llm_sc == "openai":
if not (os.getenv("OPENAI_API_KEY") or "").strip():
logger.warning("LLM_SERVICE_CODE=openai 但 OPENAI_API_KEY 未设置")
return False
else:
if not (os.getenv("DEEPSEEK_API_KEY") or "").strip():
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置(默认 LLM_SERVICE_CODE=deepseek)")
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
return False
logger.info(f"[OK] 环境检查通过")
logger.info(f" - Schema: {schema_path}")
logger.info(
" - LLM: %s",
(llm_sc or "deepseek(auto)"),
)
return True
def _default_g3sb_meta_path(structure_path: str) -> Optional[str]:
"""若存在与 table_structure 同名的 table_meta 文件则返回其路径。"""
p = Path(structure_path)
if "table_structure" not in p.name:
return None
cand = p.parent / p.name.replace("table_structure", "table_meta")
return str(cand) if cand.is_file() else None
def load_schema(schema_path: str, schema_meta_path: Optional[str] = None):
"""加载 Schema;G3SB structure JSON 会自动尝试配对 table_meta(可用 --schema-meta 指定)。"""
from schema.manager import SchemaManager
meta = (
schema_meta_path
if schema_meta_path is not None
else _default_g3sb_meta_path(schema_path)
)
logger.info(f"加载Schema: {schema_path}")
if meta:
logger.info(f" 表注释(meta): {meta}")
schema_mgr = SchemaManager.load_from_json(
schema_path, g3sb_meta_path=meta
)
stats = schema_mgr.get_statistics()
logger.info(
f"[OK] Schema加载完成: {stats['database']}, "
f"共{stats['total_tables']}张表, {stats['total_columns']}个字段"
)
return schema_mgr
def create_orchestrator(schema_mgr, args):
"""创建编排器"""
from agents.orchestrator import Text2SQLOrchestrator
from llm.router import create_llm_client, resolve_llm_service_code
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
sc = resolve_llm_service_code()
if sc == "deepseek":
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
if not api_key:
raise ValueError(
"未配置 DeepSeek API Key:请在 .env 中设置 DEEPSEEK_API_KEY,"
"或使用命令行参数 --api-key"
)
base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
cfg = DeepSeekConfig(
api_key=api_key,
base_url=base_url,
model_name=args.model or "deepseek-chat",
temperature=args.temperature,
max_tokens=args.max_tokens,
)
llm_client = create_llm_client("deepseek", **cfg.__dict__)
else:
# openai:完全由 OPENAI_* 决定;同时沿用 temperature/max_tokens 作为默认值覆盖
llm_client = create_llm_client(
"openai",
temperature=args.temperature,
max_tokens=args.max_tokens,
)
orchestrator = Text2SQLOrchestrator(
schema_manager=schema_mgr,
llm_client=llm_client,
vector_db_path=args.vector_db,
max_retry=args.max_retry,
use_vector_search=not args.no_vector_search,
# Few-shot配置
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
from bootstrap import (
create_orchestrator,
load_project_env,
load_schema,
resolve_sql_dialect,
setup_environment,
)
def single_query(orchestrator, question: str, dialect: str = "tsql"):
@@ -361,7 +169,7 @@ def main():
default=2,
help="最大重试次数(默认: 2)"
)
_load_project_env()
load_project_env()
parser.add_argument(
"--vector-db",
default=os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma").strip(),
Binary file not shown.
+87 -9
View File
@@ -4,7 +4,7 @@ SQL 验证工具集
import re
import logging
from typing import Tuple, List, Dict
from typing import Tuple, List, Dict, Optional
logger = logging.getLogger(__name__)
@@ -14,11 +14,56 @@ _CJK_IN_STRING_RE = re.compile(
)
def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
"""
扫描 SQL 中单引号字符串(含 T-SQL N'…'),若字面量内出现 CJK 则判失败。
def _cjk_text_in_string_literal(text: str) -> bool:
return bool(_CJK_IN_STRING_RE.search(text))
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本,避免误报。
def _node_in_subtree(root, target) -> bool:
if root is None:
return False
for n in root.walk():
if n is target:
return True
return False
def _cjk_string_in_forbidden_context(node, exp) -> bool:
"""
禁止含 CJK 的字面量出现在「比对/过滤」语境:WHERE、HAVING、JOIN ON、
以及 CASE 分支的 WHEN 条件(含简单 CASE 的 WHEN 值),
但允许出现在 SELECT 投影、CASE 的 THEN/ELSE 结果等纯展示位置。
"""
if node.find_ancestor(exp.Where):
return True
if node.find_ancestor(exp.Having):
return True
join = node.find_ancestor(exp.Join)
if join is not None:
on = join.args.get("on")
if on is not None and _node_in_subtree(on, node):
return True
case = node.find_ancestor(exp.Case)
while case is not None:
default = case.args.get("default")
if default is not None and _node_in_subtree(default, node):
return False
for br in case.args.get("ifs") or []:
then_expr = br.args.get("true")
cond = br.this
if then_expr is not None and _node_in_subtree(then_expr, node):
return False
if cond is not None and _node_in_subtree(cond, node):
return True
case = case.find_ancestor(exp.Case)
return False
def _check_no_cjk_legacy_text_scan(sql: str) -> Tuple[bool, List[str]]:
"""
解析失败时的回退:扫描单引号字符串(含 N'…'),字面量内出现 CJK 即失败。
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本。
"""
errors: List[str] = []
i = 0
@@ -27,7 +72,6 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
in_block_comment = False
def _read_single_quoted_string(start: int) -> Tuple[str, int]:
"""从 start 指向的 opening `'` 之后开始读,返回 (内容, 闭合引号后下标)。"""
j = start
parts: List[str] = []
while j < n:
@@ -66,10 +110,9 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
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):
if _cjk_text_in_string_literal(body):
prev = body[:48] + ("…" if len(body) > 48 else "")
errors.append(
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
@@ -79,7 +122,7 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
if sql[i] == "'":
body, i = _read_single_quoted_string(i + 1)
if _CJK_IN_STRING_RE.search(body):
if _cjk_text_in_string_literal(body):
prev = body[:48] + ("…" if len(body) > 48 else "")
errors.append(
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
@@ -92,6 +135,41 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
return len(errors) == 0, errors
def check_no_cjk_in_sql_string_literals(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
"""
禁止在「过滤/比对」语境使用含中日韩字符的字符串字面量(含 T-SQL ``N'…'``)。
允许在 SELECT 投影、CASE 的 THEN/ELSE 结果等展示用字面量中使用中文标签。
解析失败时回退为全文扫描(与旧版一致,偏严)。
"""
from sqlglot import parse_one, exp
errors: List[str] = []
try:
parsed = parse_one(sql, dialect=dialect)
except Exception as e:
logger.debug("CJK 校验回退为全文扫描(SQL 解析失败): %s", e)
return _check_no_cjk_legacy_text_scan(sql)
for node in parsed.walk():
text: Optional[str] = None
if isinstance(node, exp.Literal) and node.is_string:
text = str(node.this)
elif isinstance(node, exp.National):
text = str(node.this)
if text is None or not _cjk_text_in_string_literal(text):
continue
if _cjk_string_in_forbidden_context(node, exp):
prev = text[:48] + ("…" if len(text) > 48 else "")
errors.append(
"SQL 在 WHERE/HAVING/JOIN ON 或 CASE/WHEN 条件中出现含中文的字符串字面量(禁止)。"
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
)
return len(errors) == 0, errors
# 危险操作关键词(除非明确允许)
DANGEROUS_KEYWORDS = [
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE",