""" 用户输入意图分类:区分「自然语言查数 / 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 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" CONVERSATION = "conversation" class DialogClassifyResult(NamedTuple): intent: DialogIntent """若为 CONVERSATION,可展示给用户的引导文案;TEXT2SQL 时为 None。""" reply_suggestion: Optional[str] = None DEFAULT_CONVERSATION_REPLY = ( "您好,我是业务库 Text2SQL 助手。\n" "请用自然语言描述要查询或统计的内容(例如:查询某账户可用余额、按经纪商汇总未结算交易笔数)。\n" "输入 quit 或 exit 可退出。" ) _EMPTY_INPUT_REPLY = "请输入具体的业务查询问题,或输入 quit 退出。" # 一旦出现,倾向于按「要查数据」处理(含常见业务词,避免误判) _SQL_OR_QUERY_HINT_RE = re.compile( r"(查|查询|查出|检索|统计|列出|汇总|求和|平均|分组|排序|排名|显示|导出|筛选|过滤|" r"多少|几个|几张|哪些|占比|同比|环比|" r"余额|交易|账户|持仓|报表|结算|合约|订单|流水|经纪商|对手方|证券|资金|" r"query|select|list|show|count|sum|avg|how\s+many|statistics|\bfrom\b|\bwhere\b|\btable\b)", re.IGNORECASE, ) _CHITCHAT_PHRASES = frozenset( { "你好", "您好", "嗨", "哈喽", "hello", "hi", "hey", "早上好", "下午好", "晚上好", "在吗", "在不在", "谢谢", "多谢", "感谢", "thanks", "thank you", "thx", "再见", "拜拜", "bye", "goodbye", "哈哈", "哈哈哈", "嗯", "嗯嗯", "好的", "好", "ok", "okay", "行", "收到", "👋", "😀", "哈哈谢谢", } ) _CHITCHAT_KEYS = frozenset(p.casefold() for p in _CHITCHAT_PHRASES) _META_QUESTION_RE = re.compile( r"(你是谁|你是什么|你能(做|干)什么|你会什么|怎么用|如何使用|使用说明|帮助|help\b|" r"什么功能|干啥的)", re.IGNORECASE, ) def _normalize(text: str) -> str: t = unicodedata.normalize("NFKC", text or "").strip() t = re.sub(r"\s+", " ", t) return t def _strip_trailing_punct(t: str) -> str: return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip() 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 def _classify_dialog_rules( user_text: str, *, last_turn_was_data_query: bool = False, ) -> DialogClassifyResult: """ 规则分类(与历史行为一致):优先查数关键词;寒暄/元问题为 conversation; 若上一轮为数据查询且非寒暄/元问题,则倾向 text2sql。 """ 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)") return DialogClassifyResult(DialogIntent.TEXT2SQL, None) core = _strip_trailing_punct(t) if core.casefold() in _CHITCHAT_KEYS: logger.debug("[dialog] intent=conversation (chitchat phrase)") return DialogClassifyResult( DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY ) if _META_QUESTION_RE.search(t): logger.debug("[dialog] intent=conversation (meta question)") return DialogClassifyResult( DialogIntent.CONVERSATION, 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 )