diff --git a/rag/intent.py b/rag/intent.py index 2637a49..acd8786 100644 --- a/rag/intent.py +++ b/rag/intent.py @@ -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)) diff --git a/service/client_agent/runtime.py b/service/client_agent/runtime.py index 2cc2804..04b2cc0 100644 --- a/service/client_agent/runtime.py +++ b/service/client_agent/runtime.py @@ -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() diff --git a/service/customer_agent/chat.py b/service/customer_agent/chat.py index 0138830..393776d 100644 --- a/service/customer_agent/chat.py +++ b/service/customer_agent/chat.py @@ -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, } diff --git a/service/customer_agent/runtime.py b/service/customer_agent/runtime.py index 45b5b0a..038c2fe 100644 --- a/service/customer_agent/runtime.py +++ b/service/customer_agent/runtime.py @@ -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(