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,
Binary file not shown.
+32
View File
@@ -225,6 +225,38 @@ SQL_GENERATOR_USER = """Schema信息:
请生成**有用 SQL**(见系统提示定义):必须与「业务级黄金范例」**同构**——大写关键字、多行缩进版式、PascalCase 别名、该展示对手方/账户等名称时须 LEFT JOIN 维表;禁止输出挤成一行的「极简 SQL」。"""
# ========== Few-shot 黄金 SQL 条件适配(Chroma 库内为已校验正确答案)==========
GOLDEN_SQL_ADAPT_SYSTEM = """你是精通 Microsoft SQL Server (T-SQL) 的数据库专家。
**任务背景**:下方「标准答案 SQL」来自向量库中**已校验通过**的业务范例,结构与写法正确。当前用户问题与范例问题**语义高度相似**,仅时间、账户、状态、代码、筛选口径等「特殊条件」可能不同。
**你必须遵守**:
1. **以标准答案为主干**:优先保留其 `FROM`/`JOIN`/`ON`、主 `SELECT` 列清单与聚合/分组逻辑;**不要随意更换主表、不要拆掉必要 JOIN**,除非当前 Schema 片段中已不存在该表(此时在 Schema 内做最小替换并说明等价关系仅在脑中完成)。
2. **只改「条件类」内容**:重点调整 `WHERE`/`HAVING`/`ORDER BY`/`TOP` 中的字面量、日期区间、状态码、账户/合约/代码等过滤;将用户问题中的时间范围、业务对象、筛选口径反映到这些条件中。
3. **Schema 绝对优先**:表名、列名必须来自下方「当前 Schema 片段」;禁止臆造字段。若标准答案中某列在片段中不存在,按片段改写为合法列。
4. **T-SQL 与版式**:与常规生成一致——关键字大写、多行缩进、`WHERE` 续行以 `AND` 开头、需要时 PascalCase 英文别名;禁止 MySQL 反引号与 `CURDATE()` 等。
5. **禁止在字符串字面量中写中日韩文字**去匹配代码列;须用 Schema 注释中的代码或 JOIN 维表(与系统提示 SQL_GENERATOR 一致)。
**输出**:仅输出一条完整可执行 SQL,不要解释。"""
GOLDEN_SQL_ADAPT_USER = """【标准答案 SQL】(向量库中的已校验正确答案;相似度 {golden_score:.4f},越高越应保留结构)
对应范例问题:{golden_question}
```sql
{golden_sql}
```
【当前 Schema 片段】(标识符必须与此一致)
{schema}
【当前用户问题】
{question}
【数据库方言】{dialect}
请输出一条完整 T-SQL:**在标准答案基础上仅调整条件/字面量/排序等与当前问题相关的部分**,保留正确的 JOIN 与整体查询意图;版式与 SQL_GENERATOR 黄金范例同构。"""
# ========== Validator Agent Prompt ==========
VALIDATOR_SYSTEM = """你是一个严谨的SQL审核员,负责验证SQL语句的正确性和安全性。
+38 -1
View File
@@ -6,7 +6,7 @@ DeepSeek API 客户端封装
import os
import json
import logging
from typing import Dict, List, Optional, Any, Union
from typing import Any, Dict, Iterator, List, Optional, Union
from dataclasses import dataclass, field
from openai import OpenAI, AsyncOpenAI
from openai.types.chat import ChatCompletion, ChatCompletionMessage
@@ -104,6 +104,43 @@ class DeepSeekClient:
logger.error(f"DeepSeek API调用失败: {e}")
raise
def chat_stream(
self,
messages: List[Dict[str, str]],
**kwargs: Any,
) -> Iterator[str]:
"""
流式聊天:按 completion 增量产出文本片段(与 chat 同参,固定 stream=True)。
"""
params: Dict[str, Any] = {
"model": self.config.model_name,
"messages": messages,
"temperature": kwargs.get("temperature", self.config.temperature),
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
"top_p": kwargs.get("top_p", self.config.top_p),
"frequency_penalty": kwargs.get(
"frequency_penalty", self.config.frequency_penalty
),
"presence_penalty": kwargs.get(
"presence_penalty", self.config.presence_penalty
),
"stream": True,
"timeout": kwargs.get("timeout", self.config.timeout),
}
if self.config.extra_headers:
params["extra_headers"] = self.config.extra_headers
try:
stream = self.client.chat.completions.create(**params)
for chunk in stream:
if not chunk.choices:
continue
delta = chunk.choices[0].delta
if delta and getattr(delta, "content", None):
yield delta.content
except Exception as e:
logger.error(f"DeepSeek API流式调用失败: {e}")
raise
def chat_with_json(
self,
messages: List[Dict[str, str]],
+64 -1
View File
@@ -10,7 +10,7 @@ import json
import logging
import os
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any, Dict, Iterator, List, Optional
from openai import AsyncOpenAI, OpenAI # type: ignore[import-not-found]
from openai.types.chat import ( # type: ignore[import-not-found]
@@ -119,6 +119,69 @@ class OpenAIClient:
logger.error("OpenAI API调用失败: %s", e)
raise
def chat_stream(
self,
messages: List[Dict[str, str]],
**kwargs: Any,
) -> Iterator[str]:
"""流式聊天:按 completion 增量产出文本片段(与 chat 同参,固定 stream=True)。"""
max_tokens = kwargs.get("max_tokens", self.config.max_tokens)
max_completion_tokens = kwargs.get("max_completion_tokens", None)
params: Dict[str, Any] = {
"model": self.config.model_name,
"messages": messages,
"temperature": kwargs.get("temperature", self.config.temperature),
"top_p": kwargs.get("top_p", self.config.top_p),
"frequency_penalty": kwargs.get(
"frequency_penalty", self.config.frequency_penalty
),
"presence_penalty": kwargs.get(
"presence_penalty", self.config.presence_penalty
),
"stream": True,
"timeout": kwargs.get("timeout", self.config.timeout),
}
if max_completion_tokens is not None:
params["max_completion_tokens"] = max_completion_tokens
else:
params["max_tokens"] = max_tokens
if self.config.extra_headers:
params["extra_headers"] = self.config.extra_headers
try:
stream = self.client.chat.completions.create(**params)
for chunk in stream:
if not chunk.choices:
continue
delta = chunk.choices[0].delta
if delta and getattr(delta, "content", None):
yield delta.content
except Exception as e:
msg = str(e)
if (
"Unsupported parameter" in msg
and "max_tokens" in msg
and "max_completion_tokens" in msg
and "max_completion_tokens" not in params
):
params.pop("max_tokens", None)
params["max_completion_tokens"] = max_tokens
try:
stream = self.client.chat.completions.create(**params)
for chunk in stream:
if not chunk.choices:
continue
delta = chunk.choices[0].delta
if delta and getattr(delta, "content", None):
yield delta.content
return
except Exception:
pass
logger.error("OpenAI API流式调用失败: %s", e)
raise
def chat_with_json(
self, messages: List[Dict[str, str]], **kwargs
) -> Dict[str, Any]:
+87 -1
View File
@@ -19,7 +19,7 @@ Few-shot示例选择器 - 基于经验数据集动态选择相关示例
import json
import os
from pathlib import Path
from typing import List, Dict, Optional
from typing import List, Dict, Optional, Tuple
from dataclasses import dataclass
import numpy as np
import logging
@@ -158,6 +158,92 @@ class FewShotSelector:
self.cache_path = None
logger.info("[OK] Few-shot 使用 Chroma(%s 条)", self._chroma_store.count())
@property
def is_chroma_backend(self) -> bool:
"""是否使用 Chroma 持久化库(`data/embeddings/chroma_fewshot` 等)。"""
return self._chroma_store is not None
def _passes_filters(
self,
sample: "ExperienceSample",
min_rating: Optional[int],
required_tags: Optional[List[str]],
max_difficulty: str,
exclude_qids: Optional[List[str]],
) -> bool:
if exclude_qids and sample.qid in exclude_qids:
return False
if min_rating and sample.rating is not None and sample.rating < min_rating:
return False
if max_difficulty == "easy" and sample.difficulty != "easy":
return False
if max_difficulty == "medium" and sample.difficulty == "hard":
return False
if required_tags and not all(tag in sample.tags for tag in required_tags):
return False
return True
def select_best_with_score(
self,
question: str,
min_rating: Optional[int] = None,
required_tags: Optional[List[str]] = None,
max_difficulty: str = "hard",
exclude_qids: Optional[List[str]] = None,
) -> Optional[Tuple[ExperienceSample, float]]:
"""
返回通过筛选的**相似度最高**一条样本及分数 ``[0,1]``(与向量余弦一致:1 - distance)。
无命中时返回 ``None``。
"""
if self._chroma_store is not None:
from utils.fewshot_chroma_store import sample_from_chroma_metadata
over_fetch = 96
rows = self._chroma_store.search_raw(question, top_k=over_fetch)
for score, meta, doc in rows:
s = sample_from_chroma_metadata(meta, doc)
if not self._passes_filters(
s, min_rating, required_tags, max_difficulty, exclude_qids
):
continue
logger.info(
"[Few-shot] best_with_score: qid=%s score=%.4f preview=%r",
s.qid,
float(score),
(s.question_zh or "")[:100],
)
return (s, float(score))
return None
if self._embedder is None or self.embeddings is None or len(self.samples) == 0:
return None
q_emb = self._embedder.encode(
[question],
batch_size=1,
normalize=True,
show_progress=False,
)[0]
scores = np.dot(self.embeddings, q_emb)
best: Optional[Tuple[float, ExperienceSample]] = None
for idx, (score, sample) in enumerate(zip(scores, self.samples)):
if not self._passes_filters(
sample, min_rating, required_tags, max_difficulty, exclude_qids
):
continue
s = float(score)
if best is None or s > best[0]:
best = (s, sample)
if best is None:
return None
logger.info(
"[Few-shot] best_with_score(numpy): qid=%s score=%.4f preview=%r",
best[1].qid,
best[0],
(best[1].question_zh or "")[:100],
)
return (best[1], best[0])
def _load_samples(self):
"""从 JSONL 加载样本到内存(非 Chroma 模式必需;Chroma 空库时用于首次灌库)。"""
if self.samples_path is None: