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:
+85
-6
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user