Implement follow-up detection in dialog context to prevent context pollution in multi-topic conversations. Introduce is_likely_follow_up function to determine when to inject session history based on user input. Adjust _load_session_text2sql_context to conditionally include context based on follow-up status, enhancing SQL generation accuracy. Update relevant API endpoints to pass user text for improved context handling.

This commit is contained in:
陈辅元
2026-04-15 17:33:07 +08:00
parent a55fd18915
commit 85ad31348e
8 changed files with 103 additions and 3 deletions
+7 -3
View File
@@ -35,6 +35,7 @@ from utils.dialog_classifier import DialogIntent, classify_dialog
from utils.dialog_context import (
last_assistant_was_data_query,
messages_to_text2sql_context,
is_likely_follow_up,
)
logging.basicConfig(
@@ -246,6 +247,7 @@ def _conversation_nl_dict(reply: str) -> Dict[str, Any]:
async def _load_session_text2sql_context(
request: NLChatRequest,
user_text: str,
) -> tuple[str, bool]:
"""
在写入本轮之前读取会话历史,构造 Text2SQL 上文,并判断上一轮助手是否为数据查询。
@@ -266,7 +268,9 @@ async def _load_session_text2sql_context(
if not items:
return "", False
last_data = last_assistant_was_data_query(items)
block, _n = messages_to_text2sql_context(items)
# 方案A:仅在“续问/沿用口径”时注入少量上文;新话题直接清空,避免上下文污染 SQL。
max_pairs = 2 if is_likely_follow_up(user_text) else 0
block, _n = messages_to_text2sql_context(items, max_pairs=max_pairs)
return block, last_data
@@ -377,7 +381,7 @@ async def _chat_stream_events(request: NLChatRequest) -> AsyncIterator[bytes]:
return
text = request.message.strip()
orch = get_orchestrator()
dialog_block, last_data = await _load_session_text2sql_context(request)
dialog_block, last_data = await _load_session_text2sql_context(request, text)
classified = await asyncio.to_thread(
classify_dialog,
text,
@@ -461,7 +465,7 @@ async def nl_chat(request: NLChatRequest):
text = request.message.strip()
logger.info(f"[API] 问题: {text[:100]}...")
orch = get_orchestrator()
dialog_block, last_data = await _load_session_text2sql_context(request)
dialog_block, last_data = await _load_session_text2sql_context(request, text)
classified = await asyncio.to_thread(
classify_dialog,
text,