0.1.1 暂存
This commit is contained in:
@@ -0,0 +1,135 @@
|
||||
"""
|
||||
用户输入意图分类:区分「自然语言查数 / 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)
|
||||
Reference in New Issue
Block a user