Files
ai-g3sb-backman2.0/backend/utils/dialog_classifier.py
T
2026-04-14 10:28:22 +08:00

136 lines
3.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
"""
from __future__ import annotations
import logging
import re
import unicodedata
from enum import Enum
from typing import NamedTuple, Optional
logger = logging.getLogger(__name__)
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 classify_dialog(user_text: str) -> DialogClassifyResult:
"""
对用户一轮输入做粗分类。
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 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,
)
logger.debug("[dialog] intent=text2sql (default)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)