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
+158 -6
View File
@@ -1,17 +1,27 @@
"""
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
支持环境变量 ``DIALOG_INTENT_CLASSIFIER``:
- ``rules``:仅用关键词与短语规则(无 LLM 调用)。
- ``hybrid``(默认):明显查数词/寒暄/元问题走规则;其余交 LLM 判断(需调用方传入 ``llm_client``)。
"""
from __future__ import annotations
import logging
import os
import re
import unicodedata
from enum import Enum
from typing import NamedTuple, Optional
from typing import TYPE_CHECKING, NamedTuple, Optional
if TYPE_CHECKING:
from llm.deepseek_client import DeepSeekClient
logger = logging.getLogger(__name__)
_INTENT_CLASSIFIER_ENV = "DIALOG_INTENT_CLASSIFIER"
class DialogIntent(str, Enum):
TEXT2SQL = "text2sql"
@@ -100,12 +110,24 @@ def _strip_trailing_punct(t: str) -> str:
return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip()
def classify_dialog(user_text: str) -> DialogClassifyResult:
"""
对用户一轮输入做粗分类。
def _intent_classifier_mode() -> str:
m = (os.getenv(_INTENT_CLASSIFIER_ENV) or "hybrid").strip().lower()
if m not in ("rules", "hybrid"):
logger.warning(
"[dialog] unknown DIALOG_INTENT_CLASSIFIER=%r, use hybrid", m
)
return "hybrid"
return m
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 CONVERSATION;
其余默认 TEXT2SQL,避免漏判真实查询。
def _classify_dialog_rules(
user_text: str,
*,
last_turn_was_data_query: bool = False,
) -> DialogClassifyResult:
"""
规则分类(与历史行为一致):优先查数关键词;寒暄/元问题为 conversation;
若上一轮为数据查询且非寒暄/元问题,则倾向 text2sql。
"""
t = _normalize(user_text)
if not t:
@@ -131,5 +153,135 @@ def classify_dialog(user_text: str) -> DialogClassifyResult:
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
)
if last_turn_was_data_query:
logger.debug("[dialog] intent=text2sql (follow-up after data query)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
logger.debug("[dialog] intent=text2sql (default)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
def _classify_dialog_llm(
client: DeepSeekClient,
user_text: str,
*,
last_turn_was_data_query: bool,
dialog_context: Optional[str],
) -> DialogClassifyResult:
from config.prompts import (
DIALOG_INTENT_CLASSIFIER_SYSTEM,
DIALOG_INTENT_CLASSIFIER_USER,
)
last_label = "是" if last_turn_was_data_query else "否"
ctx = (dialog_context or "").strip()
if len(ctx) > 2400:
ctx = ctx[:2400].rstrip() + "\n…(已截断)"
if not ctx:
ctx = "(无)"
raw = client.chat_with_json(
[
{"role": "system", "content": DIALOG_INTENT_CLASSIFIER_SYSTEM},
{
"role": "user",
"content": DIALOG_INTENT_CLASSIFIER_USER.format(
last_turn_label=last_label,
context_snip=ctx,
user_message=_normalize(user_text),
),
},
],
temperature=0.0,
max_tokens=256,
top_p=1.0,
)
if raw.get("_json_decode_failed"):
raise ValueError("intent JSON decode failed")
intent_s = (raw.get("intent") or "").strip().lower()
reply = (raw.get("reply_zh") or "").strip()
if intent_s == "conversation":
if not reply:
reply = (
"若上一版 SQL 或结果不符合预期,请具体说明:希望增加/修改哪些条件、"
"时间范围或统计维度,以便重新生成。"
)
logger.info("[dialog] intent=conversation (LLM)")
return DialogClassifyResult(DialogIntent.CONVERSATION, reply_suggestion=reply)
if intent_s == "text2sql":
logger.info("[dialog] intent=text2sql (LLM)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
raise ValueError(f"unexpected intent field: {intent_s!r}")
def classify_dialog(
user_text: str,
*,
last_turn_was_data_query: bool = False,
dialog_context: Optional[str] = None,
llm_client: Optional[DeepSeekClient] = None,
) -> DialogClassifyResult:
"""
对用户一轮输入做分类。
``DIALOG_INTENT_CLASSIFIER``:
- ``rules``:仅规则。
- ``hybrid``:规则快速命中查数词/寒暄/元问题后返回;否则在有 ``llm_client`` 时用 LLM,
失败或无客户端时回退规则。
Args:
user_text: 用户输入。
last_turn_was_data_query: 上一轮助手是否为数据查询(供规则与 LLM 参考)。
dialog_context: 会话摘要,供 LLM 参考(可选)。
llm_client: DeepSeek 客户端;hybrid 下 LLM 分支需要。
"""
mode = _intent_classifier_mode()
if mode == "rules":
return _classify_dialog_rules(
user_text, last_turn_was_data_query=last_turn_was_data_query
)
# hybrid
t = _normalize(user_text)
if not t:
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=_EMPTY_INPUT_REPLY
)
if _SQL_OR_QUERY_HINT_RE.search(t):
logger.debug("[dialog] intent=text2sql (query/business hint, hybrid fast)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
core = _strip_trailing_punct(t)
if core.casefold() in _CHITCHAT_KEYS:
logger.debug("[dialog] intent=conversation (chitchat phrase, hybrid fast)")
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
)
if _META_QUESTION_RE.search(t):
logger.debug("[dialog] intent=conversation (meta question, hybrid fast)")
return DialogClassifyResult(
DialogIntent.CONVERSATION,
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
)
if llm_client is not None:
try:
return _classify_dialog_llm(
llm_client,
user_text,
last_turn_was_data_query=last_turn_was_data_query,
dialog_context=dialog_context,
)
except Exception as e:
logger.warning("[dialog] LLM intent failed, fallback rules: %s", e)
return _classify_dialog_rules(
user_text, last_turn_was_data_query=last_turn_was_data_query
)