Enhance SQL generation and streaming capabilities in the API server. Introduce optional parameters for streaming throttle and SQL stream granularity in NLChatRequest. Implement new functions for iterating SQL generation content pieces and adjusting streaming behavior based on user-defined settings. Update prompts for few-shot SQL adaptation and improve logging for SQL generation processes. Refactor orchestrator methods to support streaming responses and integrate few-shot SQL conditions. Update impact analysis documentation to reflect these changes.

This commit is contained in:
陈辅元
2026-04-16 13:48:44 +08:00
parent 695356a496
commit bff5f85d60
16 changed files with 586 additions and 103 deletions
+182 -4
View File
@@ -5,13 +5,13 @@ Text2SQL 多智能体编排器
import logging
import os # 新增
from typing import Dict, List, Optional, Tuple
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 FewShotSelector # 新增
from utils.fewshot_selector import ExperienceSample, FewShotSelector # 新增
logger = logging.getLogger(__name__)
@@ -121,6 +121,9 @@ class Text2SQLOrchestrator:
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}"
@@ -335,6 +338,118 @@ class Text2SQLOrchestrator:
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,
@@ -342,6 +457,7 @@ class Text2SQLOrchestrator:
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)
@@ -353,6 +469,7 @@ class Text2SQLOrchestrator:
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
sql_stream_callback: 若提供且 LLM 支持 chat_stream,则在 SQL 主生成/黄金适配时流式回传原始文本分片。
Returns:
SQL语句
@@ -360,11 +477,57 @@ class Text2SQLOrchestrator:
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:
@@ -431,8 +594,9 @@ class Text2SQLOrchestrator:
]
# 与选表一致:生成阶段默认贪心解码,减少同一中文问题多次 SQL 不一致
response = self.deepseek.chat(messages, temperature=0.0, top_p=1.0)
sql = response.content.strip()
sql = self._sql_chat_completion_text(
messages, sql_stream_callback=sql_stream_callback
)
# 清理可能的markdown代码块
if "```sql" in sql:
@@ -609,6 +773,7 @@ class Text2SQLOrchestrator:
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:
"""
主生成流程
@@ -619,10 +784,12 @@ class Text2SQLOrchestrator:
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
@@ -726,6 +893,7 @@ class Text2SQLOrchestrator:
dialect,
validation_feedback=feedback,
dialog_context=dc_raw or None,
sql_stream_callback=sql_stream_callback,
)
last_sql = sql
except Exception as e:
@@ -770,6 +938,11 @@ class Text2SQLOrchestrator:
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,
@@ -792,6 +965,11 @@ class Text2SQLOrchestrator:
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,