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:
Binary file not shown.
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user