Files

1046 lines
42 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Text2SQL 多智能体编排器
协调 Schema Linker、SQL Generator、Validator 三个Agent
"""
import logging
import os # 新增
from typing import Callable, Dict, List, Optional, Tuple
from dataclasses import dataclass, field
from schema.manager import SchemaManager
from schema.indexer import SchemaIndexer
from llm.deepseek_client import DeepSeekClient, DeepSeekConfig
from utils.fewshot_selector import ExperienceSample, FewShotSelector # 新增
logger = logging.getLogger(__name__)
@dataclass
class GenerationResult:
"""SQL生成结果"""
sql: str
valid: bool
errors: List[str] = field(default_factory=list)
warnings: List[str] = field(default_factory=list)
tables_used: List[str] = field(default_factory=list)
attempts: int = 1
reasoning: Optional[str] = None
metadata: Dict[str, any] = field(default_factory=dict)
class Text2SQLOrchestrator:
"""
Text2SQL 多智能体编排器
工作流程:
1. 粗筛候选表
2. Schema Linker:LLM 精筛表
3. 外键扩展 → 拼 Schema 子集
4. SQL Generator:生成 SQL
5. Validator:验证 SQL
"""
def __init__(
self,
schema_manager: SchemaManager,
llm_client: Optional[object] = None,
deepseek_api_key: Optional[str] = None,
deepseek_config: Optional[DeepSeekConfig] = None,
vector_db_path: str = "./data/embeddings/chroma",
max_retry: int = 2,
use_vector_search: bool = True,
# Few-shot配置
fewshot_enabled: bool = True,
fewshot_samples_path: Optional[str] = None,
fewshot_top_k: int = 3,
fewshot_min_rating: int = 7,
translate_english_to_zh: bool = True,
):
"""
初始化编排器
Args:
schema_manager: Schema管理器实例
llm_client: 可选:外部传入的 LLM Client(需具备 chat/chat_with_json 等方法)。
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
deepseek_config: DeepSeek配置对象(优先于api_key)
vector_db_path: 向量数据库路径
max_retry: 最大重试次数(包含首次生成)
use_vector_search: 是否使用向量检索粗筛
translate_english_to_zh: 为 True 时(默认)对**所有**非空问句做一次 LLM 归一(temperature=0),
输出一句标准中文供检索与生成;使同一语义的中英文表述对齐,从而 SQL 一致。为 False 时
不做归一(原样英文/中文)。环境变量 ``TRANSLATE_EN_TO_ZH=false`` 可关闭。
"""
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
# 初始化 LLM 客户端(历史属性名保留为 deepseek,避免大范围改动)
if llm_client is not None:
self.deepseek = llm_client
else:
if deepseek_config:
self.deepseek = DeepSeekClient(deepseek_config)
else:
self.deepseek = DeepSeekClient(DeepSeekConfig(api_key=deepseek_api_key))
# 初始化向量索引(延迟加载)
self._vector_index: Optional[SchemaIndexer] = None
self._vector_db_path = vector_db_path
# Few-shot 初始化
self.fewshot_enabled = fewshot_enabled
self.fewshot_top_k = fewshot_top_k
self.fewshot_min_rating = fewshot_min_rating
self.fewshot_selector = None
if self.fewshot_enabled:
try:
use_chroma = os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in (
"1",
"true",
"yes",
)
if fewshot_samples_path is not None:
path = str(fewshot_samples_path).strip()
else:
env_p = os.getenv("FEWSHOT_DATA_PATH")
if env_p is not None:
path = env_p.strip()
else:
# Chroma 优先时默认不再依赖 JSONL;否则保留原默认路径
path = "" if use_chroma else "./data/experiences/all_samples.jsonl"
self.fewshot_selector = FewShotSelector(path or None)
logger.info(
f"Few-shot已启用: top_k={fewshot_top_k}, "
f"min_rating={fewshot_min_rating}"
)
except Exception as e:
logger.warning(f"Few-shot加载失败: {e},将使用标准生成")
self.fewshot_enabled = False
# 若本轮走「Chroma 黄金 SQL 条件适配」,在 metadata 中回传 qid/分数
self._last_fewshot_golden: Optional[Tuple[str, float]] = None
logger.info(
f"[OK] Text2SQLOrchestrator初始化完成: "
f"max_retry={max_retry}, use_vector_search={use_vector_search}"
+ (f", fewshot=on" if self.fewshot_enabled else "")
+ (", nl→zh_norm=on" if self.translate_english_to_zh else ", nl→zh_norm=off")
)
@staticmethod
def _merge_dialog_for_model(
dialog_context: str, question: str, max_len: int
) -> str:
"""拼接上文与当前问句,控制总长,优先保留当前问句完整。"""
dc = (dialog_context or "").strip()
q = (question or "").strip()
if not dc:
return q
tail = "\n\n【当前用户问题】\n" + q
if len(dc) + len(tail) <= max_len:
return dc + tail
room = max_len - len(tail)
if room < 80:
return tail[-max_len:]
return dc[:room].rstrip() + tail
def _get_vector_index(self) -> SchemaIndexer:
"""获取或创建向量索引(懒加载)"""
if self._vector_index is None:
from utils.embedding import get_embedder
embedder = get_embedder()
self._vector_index = SchemaIndexer(
embedder=embedder,
persist_dir=self._vector_db_path
)
return self._vector_index
def _coarse_filter(
self,
question: str,
top_k: int = 20
) -> List[str]:
"""
阶段1:粗筛(向量检索)
Args:
question: 用户问题
top_k: 返回前K个候选表
Returns:
候选表名列表
"""
if not self.use_vector_search:
# 不使用向量检索时,返回所有表
logger.info("[Orchestrator] 向量搜索已禁用,使用所有表")
return self.schema_manager.list_tables()
try:
indexer = self._get_vector_index()
indexer.ensure_index_for_schema(self.schema_manager)
# 检索
logger.info(
"[Orchestrator] 开始向量检索: query_chars=%s query_preview=%r",
len(question or ""),
(question or "")[:200] + ("…" if len(question or "") > 200 else ""),
)
results = indexer.search(
query=question,
top_k=top_k,
score_threshold=0.1 # 降低阈值以提高召回率
)
candidate_tables = [r["table_name"] for r in results]
scored = [
(r["table_name"], round(float(r.get("score", 0.0)), 4))
for r in results[: min(25, len(results))]
]
logger.info(
"[Orchestrator] 向量粗筛: 命中=%s 张(阈值内),表名+分: %s",
len(candidate_tables),
scored,
)
return candidate_tables
except Exception as e:
logger.error(f"[Orchestrator] 向量检索失败: {e},降级为使用所有表", exc_info=True)
import traceback
logger.error(traceback.format_exc())
# 降级:返回所有表
return self.schema_manager.list_tables()
def _llm_select_tables(
self,
question: str,
candidate_tables: List[str],
max_tables: int = 5,
) -> Tuple[List[str], str]:
"""
阶段2:LLM精筛(Schema Linker Agent)
Args:
question: 用户问题
candidate_tables: 候选表列表
max_tables: 最多选择的表数
Returns:
(相关表列表, 推理理由)
"""
# 构造候选表信息(只显示表名和注释)
table_infos = []
for tbl_name in candidate_tables:
table = self.schema_manager.get_table(tbl_name)
if table:
comment = table.comment or "无描述"
table_infos.append(f"- {tbl_name}: {comment}")
table_list_str = "\n".join(table_infos)
# 使用DeepSeek客户端调用(而非CAMEL Agent,更直接可控)
response = self.deepseek.select_tables(
question=question,
table_list=table_list_str,
)
relevant_tables = response.get("relevant_tables", [])
reasoning = response.get("reasoning", "")
# 限制数量
relevant_tables = relevant_tables[:max_tables]
rs = (reasoning or "").strip()
logger.info(
"LLM精筛选中表:%s | reasoning_chars=%s reasoning_preview=%r",
relevant_tables,
len(rs),
rs[:600] + ("…" if len(rs) > 600 else ""),
)
return relevant_tables, reasoning
_BROKER_KEYWORDS_CN = ("对手方", "经纪商", "券商", "對手方")
def _question_implies_broker_dimension(self, question: str) -> bool:
if not question:
return False
if any(k in question for k in self._BROKER_KEYWORDS_CN):
return True
return "broker" in question.lower()
def _prioritize_broker_tables(
self, question: str, relevant_tables: List[str], max_tables: int = 5
) -> List[str]:
"""
问题涉及对手方/经纪商时,优先纳入 TSBBrokerContract 与 MCBroker(若 Schema 中存在),
避免仅选中 VSBHK 报表视图却无 BrokerID,模型又照抄黄金范例列名导致校验失败。
"""
if not self._question_implies_broker_dimension(question):
return relevant_tables[:max_tables]
priority = ["TSBBrokerContract", "MCBroker"]
present = [t for t in priority if self.schema_manager.get_table(t)]
if not present:
return relevant_tables[:max_tables]
seen = set()
merged: List[str] = []
for t in present:
if t not in seen:
merged.append(t)
seen.add(t)
for t in relevant_tables:
if len(merged) >= max_tables:
break
if t not in seen and self.schema_manager.get_table(t):
merged.append(t)
seen.add(t)
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]:
"""
外键扩展:自动添加关联表
Args:
table_names: 已选中的表名列表
Returns:
扩展后的表名列表
"""
result = set(table_names)
for tbl_name in table_names:
table = self.schema_manager.get_table(tbl_name)
if not table:
continue
# 添加被引用的表(外键指向的表)
for fk in table.foreign_keys:
if fk.ref_table not in result:
result.add(fk.ref_table)
logger.debug(f"外键扩展:添加关联表 {fk.ref_table}")
# 添加引用当前表的表(反向外键)
for other in self.schema_manager.get_tables():
for fk in other.foreign_keys:
if fk.ref_table == tbl_name and other.name not in result:
result.add(other.name)
logger.debug(f"外键扩展:添加引用表 {other.name}")
expanded = list(result)
if len(expanded) > len(table_names):
logger.info(f"外键扩展:{table_names} → {expanded}")
return expanded
def _sql_chat_completion_text(
self,
messages: List[Dict[str, str]],
*,
sql_stream_callback: Optional[Callable[[str], None]] = None,
) -> str:
"""
SQL 生成相关的一次 LLM 调用:可选流式回调(将原始 completion 文本分片回传)。
无回调或非流式失败时与非流式 chat 行为一致。
"""
ds = self.deepseek
if sql_stream_callback is not None and hasattr(ds, "chat_stream"):
try:
parts: List[str] = []
for piece in ds.chat_stream(messages, temperature=0.0, top_p=1.0):
if piece:
parts.append(piece)
sql_stream_callback(piece)
return "".join(parts).strip()
except Exception as e:
logger.warning("[GEN] chat_stream 失败,回退非流式: %s", e)
msg = ds.chat(messages, temperature=0.0, top_p=1.0)
return (msg.content or "").strip()
def _generate_sql_golden_adapt(
self,
question: str,
schema_str: str,
dialect: str,
golden: ExperienceSample,
golden_score: float,
validation_feedback: Optional[str],
dialog_context: Optional[str],
sql_stream_callback: Optional[Callable[[str], None]] = None,
) -> str:
"""
Chroma few-shot 库内 SQL 视为正确答案:在高分相似命中下,仅让模型调整条件/字面量以匹配当前问题。
"""
from config.prompts import GOLDEN_SQL_ADAPT_SYSTEM, GOLDEN_SQL_ADAPT_USER
from utils.sql_parser import normalize_sql_for_dialect
gsql = (golden.sql or "").strip()
max_sql = int(os.getenv("FEWSHOT_GOLDEN_SQL_PROMPT_MAX", "16000"))
if len(gsql) > max_sql:
gsql = gsql[:max_sql] + "\n-- …(标准答案过长,已截断)"
dialect_label = dialect
if dialect == "tsql":
dialect_label = "Microsoft SQL Server (T-SQL)"
dc = (dialog_context or "").strip()
prefix = ""
if dc:
prefix = (
"【对话上文】(用于理解指代与续问条件;请结合「当前用户问题」调整 WHERE 等。)\n"
f"{dc}\n\n"
)
user_content = prefix + GOLDEN_SQL_ADAPT_USER.format(
golden_score=golden_score,
golden_question=(golden.question_zh or "").strip(),
golden_sql=gsql,
schema=schema_str,
question=question,
dialect=dialect_label,
)
if dialect == "tsql":
user_content += (
"\n\n【硬性要求】目标库为 SQL Server(T-SQL):禁止使用 MySQL 反引号 `;"
"标识符如需引用请使用方括号,例如 [TableName]、[ColumnName]。"
"字符串连接使用 `+`(与系统提示中的标准版式范例一致)。"
"「今日」「当天」等与日期列比较时,使用 `CAST(GETDATE() AS DATE)`,"
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
"\n在 `WHERE`/`HAVING`/`JOIN ON` 及 `CASE/WHEN` 的**条件**中,**禁止**用含中日韩的 `'`/`N'…'` 字面量与代码型列比对;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN(勿写 `= '过户费'`)。`SELECT` 或 `CASE … THEN/ELSE` 的展示用中文标签允许。"
"若用户问题与对话上文均未要求按时间筛选,不得在 WHERE 中擅自添加日期列条件(与系统提示 1b 一致)。"
)
if validation_feedback:
user_content += (
"\n\n【上次校验未通过】请在保留标准答案主干的前提下修正 SQL;"
"表名、列名必须与「当前 Schema 片段」中完全一致。\n"
f"{validation_feedback}"
)
messages = [
{"role": "system", "content": GOLDEN_SQL_ADAPT_SYSTEM},
{"role": "user", "content": user_content},
]
logger.info(
"[GEN] 黄金 few-shot 条件适配: qid=%s score=%.4f",
golden.qid,
golden_score,
)
sql = self._sql_chat_completion_text(
messages, sql_stream_callback=sql_stream_callback
)
if "```sql" in sql:
sql = sql[sql.find("```sql") + 6 : sql.find("```", sql.find("```sql") + 6)].strip()
elif "```" in sql:
sql = sql[sql.find("```") + 3 : sql.find("```", sql.find("```") + 3)].strip()
sql = normalize_sql_for_dialect(sql, dialect)
lim = 12000
body = sql if len(sql) <= lim else sql[:lim] + "\n…(日志已截断)"
logger.info("生成的SQL(黄金适配,chars=%s):\n%s", len(sql), body)
return sql
def _generate_sql(
self,
question: str,
schema_str: str,
dialect: str = "tsql",
validation_feedback: Optional[str] = None,
dialog_context: Optional[str] = None,
sql_stream_callback: Optional[Callable[[str], None]] = None,
) -> str:
"""
SQL生成(SQL Generator Agent)
Args:
question: 用户问题
schema_str: Schema描述字符串
dialect: SQL方言
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
sql_stream_callback: 若提供且 LLM 支持 chat_stream,则在 SQL 主生成/黄金适配时流式回传原始文本分片。
Returns:
SQL语句
"""
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
from utils.sql_parser import normalize_sql_for_dialect
self._last_fewshot_golden = None
dc = (dialog_context or "").strip()
fewshot_question = question
if dc:
fewshot_question = f"{dc}\n\n【当前问】{question}"
golden_reuse = os.getenv("FEWSHOT_GOLDEN_REUSE", "true").lower() in (
"1",
"true",
"yes",
)
only_chroma = os.getenv("FEWSHOT_GOLDEN_ONLY_CHROMA", "true").lower() in (
"1",
"true",
"yes",
)
golden_min = float(os.getenv("FEWSHOT_GOLDEN_MIN_SCORE", "0.88"))
chroma_ok = bool(
self.fewshot_selector and self.fewshot_selector.is_chroma_backend
)
if only_chroma and not chroma_ok:
golden_reuse = False
if (
golden_reuse
and self.fewshot_enabled
and self.fewshot_selector
):
best = self.fewshot_selector.select_best_with_score(
fewshot_question,
min_rating=self.fewshot_min_rating,
)
if (
best
and best[1] >= golden_min
and (best[0].sql or "").strip()
):
ex, sc = best
sql_out = self._generate_sql_golden_adapt(
question=question,
schema_str=schema_str,
dialect=dialect,
golden=ex,
golden_score=sc,
validation_feedback=validation_feedback,
dialog_context=dialog_context,
sql_stream_callback=sql_stream_callback,
)
self._last_fewshot_golden = (ex.qid, sc)
return sql_out
# Few-shot 增强
if self.fewshot_enabled and self.fewshot_selector:
try:
examples = self.fewshot_selector.select(
question=fewshot_question,
top_k=self.fewshot_top_k,
min_rating=self.fewshot_min_rating
)
if examples:
examples_prompt = "\n\n".join([
f"示例 {i+1}:\n问题:{ex.question_zh}\nSQL:\n{ex.sql}"
for i, ex in enumerate(examples)
])
schema_str = f"参考以下相似示例的SQL编写风格:\n\n{examples_prompt}\n\n【当前Schema】\n{schema_str}"
logger.info(
"已注入 %s 个 few-shot 示例: qid=%s question_zh_preview=%r",
len(examples),
[ex.qid for ex in examples],
[((ex.question_zh or "")[:80] + "…") if len(ex.question_zh or "") > 80 else (ex.question_zh or "") for ex in examples],
)
except Exception as e:
logger.warning(f"Few-shot检索失败: {e}")
dialect_label = dialect
if dialect == "tsql":
dialect_label = "Microsoft SQL Server (T-SQL)"
prefix = ""
if dc:
prefix = (
"【对话上文】(用于理解「这/那/同样/上面/刚才」等指代及续问条件;"
"请结合下文「当前用户问题」生成 SQL。)\n"
f"{dc}\n\n"
)
user_content = prefix + SQL_GENERATOR_USER.format(
schema=schema_str,
question=question,
dialect=dialect_label,
)
if dialect == "tsql":
user_content += (
"\n\n【硬性要求】目标库为 SQL Server(T-SQL):禁止使用 MySQL 反引号 `;"
"标识符如需引用请使用方括号,例如 [TableName]、[ColumnName]。"
"字符串连接使用 `+`(与系统提示中的标准版式范例一致)。"
"「今日」「当天」等与日期列比较时,使用 `CAST(GETDATE() AS DATE)`,"
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
"\n在 `WHERE`/`HAVING`/`JOIN ON` 及 `CASE/WHEN` 的**条件**中,**禁止**用含中日韩的 `'`/`N'…'` 字面量与代码型列比对;"
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN(勿写 `= '过户费'`)。`SELECT` 或 `CASE … THEN/ELSE` 的展示用中文标签允许。"
"若用户问题与对话上文均未要求按时间筛选,不得在 WHERE 中擅自添加日期列条件(与系统提示 1b 一致)。"
)
if validation_feedback:
user_content += (
"\n\n【上次校验未通过】请根据下列错误修正 SQL,并输出完整可执行查询;"
"表名、列名必须与「当前Schema」中完全一致,不要臆造字段名。\n"
f"{validation_feedback}"
)
messages = [
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
{"role": "user", "content": user_content},
]
# 与选表一致:生成阶段默认贪心解码,减少同一中文问题多次 SQL 不一致
sql = self._sql_chat_completion_text(
messages, sql_stream_callback=sql_stream_callback
)
# 清理可能的markdown代码块
if "```sql" in sql:
sql = sql[sql.find("```sql") + 6:sql.find("```", sql.find("```sql") + 6)].strip()
elif "```" in sql:
sql = sql[sql.find("```") + 3:sql.find("```", sql.find("```") + 3)].strip()
sql = normalize_sql_for_dialect(sql, dialect)
lim = 12000
body = sql if len(sql) <= lim else sql[:lim] + "\n…(日志已截断)"
logger.info("生成的SQL(chars=%s):\n%s", len(sql), body)
return sql
def _validate_sql(
self,
sql: str,
schema_str: str,
dialect: str = "tsql",
question: str = "",
dialog_context: Optional[str] = None,
) -> Tuple[bool, List[str], List[str], Optional[int], Optional[str]]:
"""
SQL验证(程序验证 + 库探针 + 按探针分支的 LLM)
Args:
sql: SQL语句
schema_str: Schema描述
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
question: 用户自然语言(探针为 0 时用于生成补充说明)
dialog_context: 会话上文;探针 0 时与 question 一并传入说明模型
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
# 语法验证
syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect=dialect)
if not syntax_ok:
errors.extend(syntax_errors)
# Schema一致性验证
schema_ok, schema_errors = validate_schema_consistency(
sql, self.schema_manager, dialect=dialect
)
if not schema_ok:
errors.extend(schema_errors)
# 危险操作检查
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)
# T-SQL:禁止中文出现在 WHERE/HAVING/ON/CASE 条件等比对语境(展示用 CASE THEN/ELSE 允许)
if dialect == "tsql":
cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql, dialect=dialect)
if not cjk_ok:
errors.extend(cjk_errors)
# === 阶段1.5:数据库试执行(仅程序校验全部通过时;需配置 database_url) ===
if len(errors) == 0:
from db.dbhub_tools import probe_sql_execution_status_ex
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
if db_execution_status == -1:
msg = (
"【库探针结果:-1 执行失败】SQL 在目标库执行报错,"
"系统将依据下列错误**自动重新生成** SQL(请等待重试结果)。"
)
if db_probe_err:
msg += f"\n数据库返回:{db_probe_err}"
errors.append(msg)
elif db_execution_status is None:
warnings.append(
"未配置 database_url,已跳过数据库执行探针"
)
# 探针 1:库上至少有一行数据,跳过 Validator LLM,直接将 SQL 视为可交付
# 探针 0:执行成功但行数为 0,跳过 Validator LLM,另调 LLM 生成说明并引导用户补充条件
# 探针 -1:执行失败,跳过 Validator LLM,走重试
# 探针 None:走完整 Validator LLM
skip_validator_llm = db_execution_status in (-1, 0, 1)
# === 阶段2:LLM 语义验证(仅未命中库探针 0/1/-1 时) ===
if not skip_validator_llm:
try:
llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str)
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 = (
"【库探针结果:0 行】该 SQL 已在数据库成功执行,但**返回数据行数为 0**(未查到匹配记录)。"
"下方已附带完整 SQL 与原因分析,请一并阅读。"
)
fb_q = question
dc = (dialog_context or "").strip()
if dc:
fb_q = f"{dc}\n\n【当前用户问题】\n{question}"
try:
llm_fb = self.deepseek.empty_result_user_feedback(
question=fb_q,
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"
"未能自动生成详细分析。请补充时间范围、筛选条件或业务对象后重新提问。"
)
follow = (
"\n\n【追问 — 请补充后再次提问以重新生成 SQL】\n"
"1. 请根据上述分析,尽量具体地补充或修正:**时间范围**、**业务对象**(账户/合约/代码等)、"
"**筛选口径** 或 **您认为 SQL 中不合理的条件**。\n"
"2. 补充说明后请**重新发起一次自然语言提问**(无需粘贴 SQL),系统会结合您的新描述**重新生成**查询。"
)
empty_feedback = (empty_feedback or prefix) + follow
is_valid = len(errors) == 0
logger.info(
"[validate] 程序+探针+LLM 汇总: valid=%s err_count=%s warn_count=%s "
"db_execution_status=%s sql_chars=%s",
is_valid,
len(errors),
len(warnings),
db_execution_status,
len(sql or ""),
)
if errors:
logger.info("[validate] errors 预览: %s", errors[:5])
if warnings:
logger.info("[validate] warnings: %s", warnings[:5])
return is_valid, errors, warnings, db_execution_status, empty_feedback
def generate(
self,
question: str,
dialect: str = "tsql",
top_k_candidates: int = 20,
include_schema_in_result: bool = False,
dialog_context: Optional[str] = None,
sql_stream_callback: Optional[Callable[[str], None]] = None,
) -> GenerationResult:
"""
主生成流程
Args:
question: 用户自然语言问题
dialect: SQL方言(默认 tsql;与 sqlglot 一致)
top_k_candidates: 粗筛候选表数量
include_schema_in_result: 结果中是否包含使用的Schema字符串
dialog_context: 前几轮对话可读摘要;选表、向量粗筛、SQL 生成与无数据说明会参考
sql_stream_callback: 可选;SQL 主生成 LLM 输出分片回调(用于 API SSE)
Returns:
GenerationResult对象
"""
self._last_fewshot_golden = None
original_question = (question or "").strip()
translation_meta: Dict = {}
work_question = original_question
if self.translate_english_to_zh and original_question:
try:
zh = self.deepseek.normalize_nl_question_for_text2sql(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
dc_raw = (dialog_context or "").strip()
retrieval_question = self._merge_dialog_for_model(dc_raw, question, 4000)
linker_question = self._merge_dialog_for_model(dc_raw, question, 6000)
logger.info(
"[GEN] 开始生成SQL: question_chars=%s preview=%r dialog_context_chars=%s",
len(question or ""),
(question or "")[:300] + ("…" if len(question or "") > 300 else ""),
len(dc_raw) if dc_raw else 0,
)
attempt = 0
last_sql = None
last_errors = []
last_db_execution_status: Optional[int] = None
filtered_schema_str = ""
tables_used = []
while attempt < self.max_retry:
logger.info(f" 尝试 #{attempt + 1}")
# === Step 1: Schema筛选(仅首次) ===
if attempt == 0:
# 1.1 粗筛
candidate_tables = self._coarse_filter(
retrieval_question, top_k=top_k_candidates
)
# 1.2 LLM精筛
relevant_tables, reasoning = self._llm_select_tables(
linker_question,
candidate_tables,
)
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)
tables_used = expanded_tables
# 1.4 生成Schema字符串
filtered_schema_str = self.schema_manager.to_compact_string(
table_names=expanded_tables,
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」实际引用到的表,并对齐程序校验与生成上下文
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:
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,
dialog_context=dc_raw or None,
sql_stream_callback=sql_stream_callback,
)
last_sql = sql
except Exception as e:
last_errors = [f"SQL生成失败: {str(e)}"]
attempt += 1
continue
# === Step 3: 验证 ===
is_valid, errors, warnings, db_probe, empty_feedback = self._validate_sql(
sql,
filtered_schema_str,
dialect=dialect,
question=question,
dialog_context=dc_raw or None,
)
if db_probe is not None:
last_db_execution_status = db_probe
if not is_valid:
last_errors = errors
logger.warning(f" [FAIL] 验证失败:{errors}")
attempt += 1
continue
logger.info(f"[OK] SQL生成与验证通过({attempt + 1}次尝试)")
meta: Dict = dict(translation_meta)
if db_probe is not None:
meta["db_execution_status"] = db_probe
if empty_feedback:
meta["db_empty_feedback"] = empty_feedback
if db_probe == 1:
try:
meta["sql_delivery_message"] = (
self.deepseek.sql_probe_success_delivery_message(
question=question,
sql=sql,
)
)
except Exception as e:
logger.warning("[GEN] 探针1交付说明生成失败: %s", e)
meta["sql_delivery_message"] = None
if dc_raw:
meta["dialog_context_chars"] = len(dc_raw)
if self._last_fewshot_golden:
gq, gsc = self._last_fewshot_golden
meta["fewshot_golden_reuse"] = True
meta["fewshot_golden_qid"] = gq
meta["fewshot_golden_score"] = gsc
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
if dc_raw:
fail_meta["dialog_context_chars"] = len(dc_raw)
if self._last_fewshot_golden:
gq, gsc = self._last_fewshot_golden
fail_meta["fewshot_golden_reuse"] = True
fail_meta["fewshot_golden_qid"] = gq
fail_meta["fewshot_golden_score"] = gsc
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:
"""
构建向量索引(可选,提前构建可加速首次查询)
Args:
force_rebuild: 是否强制重建
Returns:
是否成功构建
"""
indexer = self._get_vector_index()
return indexer.build_index(
self.schema_manager,
force_rebuild=force_rebuild
)
def get_statistics(self) -> Dict:
"""获取统计信息"""
schema_stats = self.schema_manager.get_statistics()
indexer = self._get_vector_index()
index_stats = indexer.get_statistics()
return {
"schema": schema_stats,
"vector_index": index_stats,
"max_retry": self.max_retry,
"use_vector_search": self.use_vector_search,
}