288 lines
8.9 KiB
Python
288 lines
8.9 KiB
Python
"""
|
||
用户输入意图分类:区分「自然语言查数 / 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
|
||
)
|