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