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
+85 -6
View File
@@ -32,6 +32,10 @@ from main import setup_environment, load_schema, create_orchestrator, resolve_sq
from agents.orchestrator import GenerationResult
from nl_lite_store import lite_nl_store
from utils.dialog_classifier import DialogIntent, classify_dialog
from utils.dialog_context import (
last_assistant_was_data_query,
messages_to_text2sql_context,
)
logging.basicConfig(
level=logging.INFO,
@@ -54,7 +58,12 @@ def get_orchestrator():
return orchestrator
if not setup_environment():
raise RuntimeError("环境检查失败,请检查配置")
raise RuntimeError(
"环境检查失败(请查看上方 WARNING/INFO,常见原因:"
"USE_LOCAL_EMBEDDING=true 但未配置或找不到 EMBEDDING_MODEL_PATH;"
"或 USE_LOCAL_EMBEDDING=false 但未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_*;"
"或 Schema 文件路径不对、DEEPSEEK_API_KEY 未设置)"
)
schema_path = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
schema_meta_path = os.getenv("SCHEMA_META_PATH", None)
@@ -111,6 +120,20 @@ class DataQueryResult(BaseModel):
truncated: bool = Field(False, description="是否截断")
sql_explain: Optional[str] = Field(None, description="SQL自然语言说明")
can_export: Optional[bool] = Field(True, description="是否允许导出")
db_execution_status: Optional[int] = Field(
None,
description="库探针:1=有数据行,0=无行需追问补充,-1=执行失败(重试),None=未探针",
)
db_empty_feedback: Optional[str] = Field(
None, description="探针0时的无行说明与追问(含问题分析)"
)
follow_up_required: bool = Field(
False,
description="True 表示探针为0:需用户补充条件后重新提问以重新生成SQL",
)
sql_delivery_message: Optional[str] = Field(
None, description="探针为1时LLM生成的面向用户SQL交付说明"
)
class NLChatSuccessData(BaseModel):
@@ -185,6 +208,32 @@ def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
}
async def _load_session_text2sql_context(
request: NLChatRequest,
) -> tuple[str, bool]:
"""
在写入本轮之前读取会话历史,构造 Text2SQL 上文,并判断上一轮助手是否为数据查询。
"""
sid = (request.session_id or "").strip()
if not sid:
return "", False
data = await lite_nl_store.get_messages(
request.user_id,
request.visitor_biz_id,
sid,
limit=200,
offset=0,
)
if not data:
return "", False
items = data.get("items") or []
if not items:
return "", False
last_data = last_assistant_was_data_query(items)
block, _n = messages_to_text2sql_context(items)
return block, last_data
async def _append_session_if_needed(request: NLChatRequest, user_text: str, data: Dict[str, Any]) -> None:
if not (request.session_id and request.session_id.strip()):
return
@@ -226,6 +275,10 @@ def _nl_dict_from_generation(result: GenerationResult) -> Dict[str, Any]:
dbe = result.metadata.get("db_empty_feedback")
if dbe:
branch_result["db_empty_feedback"] = dbe
branch_result["follow_up_required"] = bool(result.valid and dbs == 0)
sdm = result.metadata.get("sql_delivery_message")
if sdm:
branch_result["sql_delivery_message"] = str(sdm)
conf = 1.0 if result.valid else 0.0
if result.valid:
reason = f"使用了 {len(result.tables_used)} 张表"
@@ -247,6 +300,10 @@ def _sql_gen_stream_html(result: GenerationResult) -> str:
inner = json.dumps({"sql": sql}, ensure_ascii=False)
parts = [f"<data>{inner}</data>"]
explain = ""
if result.metadata.get("sql_delivery_message"):
parts.append(
f'<div class="sql-delivery-body">{html_lib.escape(str(result.metadata["sql_delivery_message"]))}</div>'
)
if result.metadata.get("db_empty_feedback"):
explain = str(result.metadata["db_empty_feedback"])
elif result.warnings:
@@ -258,15 +315,21 @@ def _sql_gen_stream_html(result: GenerationResult) -> str:
return "".join(parts)
async def _run_generate(question: str, top_k: int = 20) -> GenerationResult:
async def _run_generate(
question: str,
top_k: int = 20,
dialog_context: Optional[str] = None,
) -> GenerationResult:
orch = get_orchestrator()
dialect = resolve_sql_dialect(os.getenv("TEXT2SQL_DIALECT", "sqlserver"))
dc = (dialog_context or "").strip() or None
def _call() -> GenerationResult:
return orch.generate(
question=question.strip(),
dialect=dialect,
top_k_candidates=top_k,
dialog_context=dc,
)
return await asyncio.to_thread(_call)
@@ -277,7 +340,15 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
yield _sse_data({"code": 400, "msg": "消息内容不能为空", "data": None})
return
text = request.message.strip()
classified = classify_dialog(text)
orch = get_orchestrator()
dialog_block, last_data = await _load_session_text2sql_context(request)
classified = await asyncio.to_thread(
classify_dialog,
text,
last_turn_was_data_query=last_data,
dialog_context=dialog_block or None,
llm_client=orch.deepseek,
)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
logger.info("[API/stream] 对话意图: conversation(跳过 Text2SQL,与 CLI single_query 一致)")
@@ -290,7 +361,7 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
yield _sse_data({"stage": "orchestrator", "stream_kind": "content", "content": "DATA_QUERY"})
try:
result = await _run_generate(text)
result = await _run_generate(text, dialog_context=dialog_block or None)
except Exception as e:
logger.error(f"[API/stream] 生成异常: {e}")
yield _sse_data({"code": 500, "msg": str(e), "data": None})
@@ -353,7 +424,15 @@ async def nl_chat(request: NLChatRequest):
raise HTTPException(status_code=400, detail="消息内容不能为空")
text = request.message.strip()
logger.info(f"[API] 问题: {text[:100]}...")
classified = classify_dialog(text)
orch = get_orchestrator()
dialog_block, last_data = await _load_session_text2sql_context(request)
classified = await asyncio.to_thread(
classify_dialog,
text,
last_turn_was_data_query=last_data,
dialog_context=dialog_block or None,
llm_client=orch.deepseek,
)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
logger.info("[API/chat] 对话意图: conversation(跳过 Text2SQL)")
@@ -361,7 +440,7 @@ async def nl_chat(request: NLChatRequest):
await _append_session_if_needed(request, text, data_dict)
return {"code": 200, "msg": "success", "data": data_dict}
result = await _run_generate(text)
result = await _run_generate(text, dialog_context=dialog_block or None)
data_dict = _nl_dict_from_generation(result)
response_data = NLChatSuccessData.model_validate(data_dict)
msg = "success" if result.valid else "partial"