feat:客服意图识别增加记忆功能
This commit is contained in:
+103
-9
@@ -1,9 +1,12 @@
|
||||
"""Customer-service intent recognition contract."""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from inspect import isawaitable
|
||||
|
||||
|
||||
logger = logging.getLogger("rag.intent")
|
||||
@@ -26,10 +29,24 @@ INTENT_VALUES = frozenset(item.value for item in Intent)
|
||||
_INTENT_PATTERN = re.compile(
|
||||
"|".join(re.escape(value) for value in sorted(INTENT_VALUES, key=len, reverse=True))
|
||||
)
|
||||
_JSON_PATTERN = re.compile(r"\{.*\}", re.S)
|
||||
|
||||
# 改写句允许的最大长度:相对原句的倍数与绝对下限取大者,防止模型把历史大段塞进改写句
|
||||
_REWRITE_MAX_RATIO = 4
|
||||
_REWRITE_MIN_LIMIT = 200
|
||||
|
||||
INTENT_SYSTEM_PROMPT = (
|
||||
"你是华夏科技(一家基金代销金融机构)智能客服的意图分类器。"
|
||||
"请根据用户输入,从以下选项中选择最匹配的意图,并仅输出对应的英文标签(不输出任何其他内容):\n"
|
||||
"你会看到最近几轮对话(可能为空)和用户的【当前输入】。\n"
|
||||
"任务:\n"
|
||||
"1. 只对【当前输入】判定意图;【最近对话】仅用于理解当前输入中省略的主语或指代"
|
||||
"(如“它”“那个”“呢”“还有吗”)。\n"
|
||||
"2. 若当前输入依赖上文才能理解,请把它改写为一句不依赖上文的完整问题,"
|
||||
"只能补全上文已明确出现的对象,不得添加新信息、不得改变原意;"
|
||||
"若当前输入本身已完整,query 原样返回。\n"
|
||||
"3. 当前输入若明显切换到新话题,以当前输入为准,不要延续上一轮的意图。\n"
|
||||
"4. 致谢、告别、寒暄即使出现在知识问答之后也归为 chitchat。\n"
|
||||
"意图选项:\n"
|
||||
"- guide_purchase: 用户询问如何购买基金、开户、注册等引导类问题\n"
|
||||
"- want_advisor: 用户希望获得个性化基金推荐或投资顾问服务\n"
|
||||
"- knowledge_qa: 用户询问基金相关的知识性问题,如净值、费率、风险、申赎规则等\n"
|
||||
@@ -41,10 +58,26 @@ INTENT_SYSTEM_PROMPT = (
|
||||
"如写代码、讲笑话、写作文、问天气、聊政治、情感咨询、做数学题等\n"
|
||||
"- no_match: 无法归入以上任何一类\n"
|
||||
"注意:寒暄性质的一两句话归为 chitchat;一旦用户提出金融之外的实质性请求,归为 off_topic。\n"
|
||||
"仅输出一个小写英文标签,不要输出解释、标点或换行。"
|
||||
"只输出一行 JSON,格式:{\"intent\": \"<小写英文标签>\", \"query\": \"<完整问题>\"},"
|
||||
"不要输出解释、Markdown 或其他内容。"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class IntentResult:
|
||||
intent: Intent
|
||||
# 结合历史补全后的独立问题;无需补全或解析失败时等于原句
|
||||
query: str
|
||||
used_history: bool = False
|
||||
|
||||
|
||||
async def _config_value(config_getter, key: str, default):
|
||||
value = config_getter(key, str(default))
|
||||
if isawaitable(value):
|
||||
value = await value
|
||||
return type(default)(value)
|
||||
|
||||
|
||||
def parse_intent(raw: str | None) -> Intent:
|
||||
"""从模型原始输出中提取第一个合法标签,提取不到则返回 NO_MATCH。"""
|
||||
if not raw:
|
||||
@@ -53,16 +86,77 @@ def parse_intent(raw: str | None) -> Intent:
|
||||
return Intent(match.group(0)) if match else Intent.NO_MATCH
|
||||
|
||||
|
||||
async def intent_recognize(query: str, *, llm_client) -> Intent:
|
||||
if not query or not query.strip():
|
||||
return Intent.NO_MATCH
|
||||
messages = [
|
||||
def parse_intent_result(raw: str | None, *, fallback_query: str, used_history: bool = False) -> IntentResult:
|
||||
"""优先按 JSON 解析意图与改写句;JSON 不可用时退回正则抽标签、原句作为 query。"""
|
||||
match = _JSON_PATTERN.search(raw) if raw else None
|
||||
data = None
|
||||
if match:
|
||||
try:
|
||||
data = json.loads(match.group(0))
|
||||
except ValueError:
|
||||
data = None
|
||||
if not isinstance(data, dict):
|
||||
return IntentResult(parse_intent(raw), fallback_query, used_history)
|
||||
intent = parse_intent(str(data.get("intent") or ""))
|
||||
rewritten = str(data.get("query") or "").strip()
|
||||
limit = max(len(fallback_query) * _REWRITE_MAX_RATIO, _REWRITE_MIN_LIMIT)
|
||||
query = rewritten if rewritten and len(rewritten) <= limit else fallback_query
|
||||
return IntentResult(intent, query, used_history)
|
||||
|
||||
|
||||
def render_history(history, *, max_turns: int, max_chars: int) -> str:
|
||||
"""把最近 max_turns 轮 user/assistant 消息折叠成一段文本,每条截断到 max_chars。"""
|
||||
if max_turns <= 0 or not history:
|
||||
return ""
|
||||
recent = [
|
||||
message for message in history
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") in ("user", "assistant")
|
||||
and message.get("content")
|
||||
][-max_turns * 2:]
|
||||
lines = []
|
||||
for message in recent:
|
||||
role = "用户" if message["role"] == "user" else "客服"
|
||||
content = str(message["content"]).replace("\n", " ").strip()
|
||||
if len(content) > max_chars:
|
||||
content = content[:max_chars] + "…"
|
||||
lines.append(f"{role}:{content}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def build_intent_messages(query: str, rendered_history: str) -> list[dict]:
|
||||
if rendered_history:
|
||||
user_content = f"【最近对话】\n{rendered_history}\n\n【当前输入】\n{query}"
|
||||
else:
|
||||
user_content = f"【当前输入】\n{query}"
|
||||
return [
|
||||
{"role": "system", "content": INTENT_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": query},
|
||||
{"role": "user", "content": user_content},
|
||||
]
|
||||
|
||||
|
||||
async def intent_recognize(
|
||||
query: str,
|
||||
*,
|
||||
llm_client,
|
||||
history: list[dict] | None = None,
|
||||
config_getter=None,
|
||||
) -> IntentResult:
|
||||
"""结合最近几轮对话识别当前输入的意图,并给出补全指代后的独立问题。
|
||||
|
||||
history 为空或 agent.customer.intent.history_turns 配置为 0 时退化为仅看当前句。
|
||||
"""
|
||||
if not query or not query.strip():
|
||||
return IntentResult(Intent.NO_MATCH, query or "")
|
||||
max_turns, max_chars = 3, 200
|
||||
if config_getter is not None:
|
||||
max_turns = await _config_value(config_getter, "agent.customer.intent.history_turns", max_turns)
|
||||
max_chars = await _config_value(config_getter, "agent.customer.intent.history_max_chars", max_chars)
|
||||
rendered = render_history(history, max_turns=max_turns, max_chars=max_chars)
|
||||
messages = build_intent_messages(query, rendered)
|
||||
try:
|
||||
raw = await llm_client.chat(messages)
|
||||
except Exception:
|
||||
logger.exception("intent recognition failed")
|
||||
return Intent.NO_MATCH
|
||||
return parse_intent(raw)
|
||||
return IntentResult(Intent.NO_MATCH, query, bool(rendered))
|
||||
return parse_intent_result(raw, fallback_query=query, used_history=bool(rendered))
|
||||
|
||||
@@ -207,8 +207,13 @@ def build_client_runtime(
|
||||
config_getter=config_getter,
|
||||
)
|
||||
|
||||
async def recognize(query):
|
||||
return await intent_recognize(query, llm_client=llm_client)
|
||||
async def recognize(query, history=None):
|
||||
return await intent_recognize(
|
||||
query,
|
||||
llm_client=llm_client,
|
||||
history=history,
|
||||
config_getter=config_getter,
|
||||
)
|
||||
|
||||
async def generate(messages):
|
||||
memory_context = _active_memory_context.get()
|
||||
|
||||
@@ -2,10 +2,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from inspect import isawaitable
|
||||
|
||||
from rag.intent import Intent
|
||||
from rag.intent import Intent, IntentResult
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class QueryTooLongError(ValueError):
|
||||
@@ -44,6 +48,12 @@ class AnonymousCustomerAgent:
|
||||
async def handle(self, session_id: str, query: str, *, trace_id: str) -> dict:
|
||||
if len(query) > 2000:
|
||||
raise QueryTooLongError("query长度不能超过2000字符")
|
||||
# 先取历史再写入当前问题,保证意图识别拿到的历史不含本轮输入;取不到历史不阻断请求
|
||||
try:
|
||||
history = await self.context.get(session_id)
|
||||
except Exception:
|
||||
logger.exception("load conversation history failed: session_id=%s", session_id)
|
||||
history = []
|
||||
await self.context.append(session_id, "user", query)
|
||||
if self._contains_sensitive_input(query):
|
||||
await _maybe_await(self.audit_writer(
|
||||
@@ -52,7 +62,11 @@ class AnonymousCustomerAgent:
|
||||
session_id=session_id,
|
||||
))
|
||||
|
||||
intent = await _maybe_await(self.intent_recognize(query))
|
||||
recognized = await _maybe_await(self.intent_recognize(query, history))
|
||||
if isinstance(recognized, IntentResult):
|
||||
intent, search_query = recognized.intent, recognized.query
|
||||
else:
|
||||
intent, search_query = recognized, query
|
||||
sources = []
|
||||
if intent == Intent.GUIDE_PURCHASE:
|
||||
answer = await _config(
|
||||
@@ -77,7 +91,8 @@ class AnonymousCustomerAgent:
|
||||
answer = await self._chitchat(session_id)
|
||||
elif intent in (Intent.KNOWLEDGE_QA, Intent.COMPANY_INFO):
|
||||
try:
|
||||
sources = await _maybe_await(self.rag_retrieve(query, None))
|
||||
# 用补全指代后的问题检索,省略主语的追问才能命中
|
||||
sources = await _maybe_await(self.rag_retrieve(search_query, None))
|
||||
except Exception:
|
||||
sources = []
|
||||
if not sources:
|
||||
@@ -118,6 +133,7 @@ class AnonymousCustomerAgent:
|
||||
"answer": answer,
|
||||
"sources": sources,
|
||||
"intent": intent.value,
|
||||
"rewritten_query": search_query,
|
||||
"trace_id": trace_id,
|
||||
}
|
||||
|
||||
|
||||
@@ -36,8 +36,13 @@ def build_anonymous_runtime(
|
||||
config_getter=config_getter,
|
||||
)
|
||||
|
||||
async def recognize(query):
|
||||
return await intent_recognize(query, llm_client=llm_client)
|
||||
async def recognize(query, history=None):
|
||||
return await intent_recognize(
|
||||
query,
|
||||
llm_client=llm_client,
|
||||
history=history,
|
||||
config_getter=config_getter,
|
||||
)
|
||||
|
||||
async def generate(messages):
|
||||
return await generate_answer(
|
||||
|
||||
Reference in New Issue
Block a user