""" 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] 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**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;" "业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。" ) 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**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;" "业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。" ) 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:禁止中文等业务词出现在字符串字面量(如 FeeNatureID = '过户费') if dialect == "tsql": cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql) if not cjk_ok: errors.extend(cjk_errors) # === 阶段1.5:数据库试执行(仅程序校验全部通过时;需配置 database_url) === if len(errors) == 0: from db.dbhub_tools import probe_sql_execution_status_ex db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql) if db_execution_status == -1: msg = ( "【库探针结果:-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 ) # 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 ) 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, }