Enhance dialog context handling for Text2SQL queries by integrating session history. Introduce new methods for summarizing previous assistant messages and determining if the last interaction was a data query. Update environment configuration for embedding options and improve error handling in SQL generation. Add user-facing delivery messages for successful SQL execution. This update supports more coherent follow-up questions and improves user experience in conversational interactions.

This commit is contained in:
陈辅元
2026-04-14 18:02:12 +08:00
parent 026d3bf88c
commit ca8bc5e7de
87 changed files with 5808 additions and 6351 deletions
+102 -19
View File
@@ -64,7 +64,7 @@ class Text2SQLOrchestrator:
schema_manager: Schema管理器实例
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
deepseek_config: DeepSeek配置对象(优先于api_key)
embedding_model_path: Qwen3-Embedding模型路径
embedding_model_path: 本地 Embedding 模型目录(仅 USE_LOCAL_EMBEDDING=true)
vector_db_path: 向量数据库路径
max_retry: 最大重试次数(包含首次生成)
use_vector_search: 是否使用向量检索粗筛
@@ -96,12 +96,22 @@ class Text2SQLOrchestrator:
if self.fewshot_enabled:
try:
path = fewshot_samples_path or os.getenv(
"FEWSHOT_DATA_PATH",
"./data/experiences/all_samples.jsonl"
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,
path or None,
embedding_model_path=self._embedding_model_path,
)
logger.info(
@@ -119,6 +129,23 @@ class Text2SQLOrchestrator:
+ (", en→zh=on" if self.translate_english_to_zh else ", en→zh=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:
@@ -291,6 +318,7 @@ class Text2SQLOrchestrator:
schema_str: str,
dialect: str = "tsql",
validation_feedback: Optional[str] = None,
dialog_context: Optional[str] = None,
) -> str:
"""
SQL生成(SQL Generator Agent)
@@ -301,17 +329,24 @@ class Text2SQLOrchestrator:
dialect: SQL方言
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
Returns:
SQL语句
"""
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
from utils.sql_parser import normalize_sql_for_dialect
dc = (dialog_context or "").strip()
fewshot_question = question
if dc:
fewshot_question = f"{dc}\n\n【当前问】{question}"
# Few-shot 增强
if self.fewshot_enabled and self.fewshot_selector:
try:
examples = self.fewshot_selector.select(
question=question,
question=fewshot_question,
top_k=self.fewshot_top_k,
min_rating=self.fewshot_min_rating
)
@@ -329,7 +364,15 @@ class Text2SQLOrchestrator:
if dialect == "tsql":
dialect_label = "Microsoft SQL Server (T-SQL)"
user_content = SQL_GENERATOR_USER.format(
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,
@@ -380,6 +423,7 @@ class Text2SQLOrchestrator:
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)
@@ -389,6 +433,7 @@ class Text2SQLOrchestrator:
schema_str: Schema描述
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
question: 用户自然语言(探针为 0 时用于生成补充说明)
dialog_context: 会话上文;探针 0 时与 question 一并传入说明模型
Returns:
(是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
@@ -433,11 +478,11 @@ class Text2SQLOrchestrator:
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
if db_execution_status == -1:
msg = (
"数据库执行验证失败:SQL 在目标库执行报错(探针状态 -1),"
"将据此重新生成 SQL。"
"【库探针结果:-1 执行失败】SQL 在目标库执行报错,"
"系统将依据下列错误**自动重新生成** SQL(请等待重试结果)。"
)
if db_probe_err:
msg += f" 数据库返回:{db_probe_err}"
msg += f"\n数据库返回:{db_probe_err}"
errors.append(msg)
elif db_execution_status is None:
warnings.append(
@@ -482,23 +527,35 @@ class Text2SQLOrchestrator:
# === 阶段2b:探针 0 时生成用户可读补充说明(仍返回 SQL,由 API/CLI 一并展示) ===
if db_execution_status == 0 and len(errors) == 0:
prefix = (
"该 SQL 已在数据库成功执行,但返回的数据行数为 0(未查到匹配记录)。"
"请将下方 SQL 与说明一并核对;若不符合预期,请补充或调整条件后再次提问。"
"【库探针结果: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=question,
question=fb_q,
sql=sql,
schema=schema_str,
)
empty_feedback = f"{prefix}\n\n【分析与建议】\n{llm_fb}"
empty_feedback = f"{prefix}\n\n【问题分析】\n{llm_fb}"
except Exception as e:
logger.warning(f"无数据说明生成失败: {e}")
empty_feedback = (
f"{prefix}\n\n【分析与建议】\n"
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
return is_valid, errors, warnings, db_execution_status, empty_feedback
@@ -507,7 +564,8 @@ class Text2SQLOrchestrator:
question: str,
dialect: str = "tsql",
top_k_candidates: int = 20,
include_schema_in_result: bool = False
include_schema_in_result: bool = False,
dialog_context: Optional[str] = None,
) -> GenerationResult:
"""
主生成流程
@@ -517,6 +575,7 @@ class Text2SQLOrchestrator:
dialect: SQL方言(默认 tsql;与 sqlglot 一致)
top_k_candidates: 粗筛候选表数量
include_schema_in_result: 结果中是否包含使用的Schema字符串
dialog_context: 前几轮对话可读摘要;选表、向量粗筛、SQL 生成与无数据说明会参考
Returns:
GenerationResult对象
@@ -543,6 +602,9 @@ class Text2SQLOrchestrator:
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(f"[GEN] 开始生成SQL:{question[:50]}...")
attempt = 0
@@ -558,14 +620,18 @@ class Text2SQLOrchestrator:
# === Step 1: Schema筛选(仅首次) ===
if attempt == 0:
# 1.1 粗筛
candidate_tables = self._coarse_filter(question, top_k=top_k_candidates)
candidate_tables = self._coarse_filter(
retrieval_question, top_k=top_k_candidates
)
# 1.2 LLM精筛
relevant_tables, reasoning = self._llm_select_tables(
question,
linker_question,
candidate_tables,
)
relevant_tables = self._prioritize_broker_tables(question, relevant_tables)
relevant_tables = self._prioritize_broker_tables(
linker_question, relevant_tables
)
# 1.3 外键扩展
expanded_tables = self._expand_relations(relevant_tables)
@@ -613,6 +679,7 @@ class Text2SQLOrchestrator:
filtered_schema_str,
dialect,
validation_feedback=feedback,
dialog_context=dc_raw or None,
)
last_sql = sql
except Exception as e:
@@ -626,6 +693,7 @@ class Text2SQLOrchestrator:
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
@@ -643,6 +711,19 @@ class Text2SQLOrchestrator:
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)
result = GenerationResult(
sql=sql,
@@ -663,6 +744,8 @@ class Text2SQLOrchestrator:
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)
return GenerationResult(
sql=last_sql or "",
valid=False,