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

292 lines
9.3 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 的对话。
支持环境变量 ``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.info("[dialog] intent=text2sql (query/business hint) preview=%r", user_text[:120])
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
core = _strip_trailing_punct(t)
if core.casefold() in _CHITCHAT_KEYS:
logger.info("[dialog] intent=conversation (chitchat phrase) core=%r", core)
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
)
if _META_QUESTION_RE.search(t):
logger.info("[dialog] intent=conversation (meta question) preview=%r", user_text[:120])
return DialogClassifyResult(
DialogIntent.CONVERSATION,
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
)
if last_turn_was_data_query:
logger.info("[dialog] intent=text2sql (follow-up after data query) preview=%r", user_text[:120])
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
logger.info("[dialog] intent=text2sql (default rules) preview=%r", user_text[:120])
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) reply_chars=%s reply_preview=%r",
len(reply),
reply[:300] + ("…" if len(reply) > 300 else ""),
)
return DialogClassifyResult(DialogIntent.CONVERSATION, reply_suggestion=reply)
if intent_s == "text2sql":
logger.info("[dialog] intent=text2sql (LLM) user_preview=%r", user_text[:200])
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.info("[dialog] intent=text2sql (hybrid fast: query hint) preview=%r", user_text[:120])
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
core = _strip_trailing_punct(t)
if core.casefold() in _CHITCHAT_KEYS:
logger.info("[dialog] intent=conversation (hybrid fast: chitchat) core=%r", core)
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
)
if _META_QUESTION_RE.search(t):
logger.info("[dialog] intent=conversation (hybrid fast: meta) preview=%r", user_text[:120])
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
)