- Implemented `_merged_items` and `_merged_memory_text` functions to consolidate consult and chitchat memories, improving context awareness in intent classification and response generation. - Updated intent prompts to include recent dialogue history, aiding in the resolution of ambiguous user queries. - Enhanced `search_knowledge` tool to utilize context window for better query understanding, addressing issues with omitted references in user inputs. - Fixed existing test cases to reflect changes in intent constants and ensure accurate context handling during tests. This update significantly improves the handling of multi-turn dialogues, ensuring a more coherent and contextually aware interaction for users.
429 lines
14 KiB
Python
429 lines
14 KiB
Python
"""游客 Agent 编排:LangGraph 9 节点 + DeepSeek LLM。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import TypedDict
|
||
|
||
from langgraph.graph import END, StateGraph
|
||
from langchain_openai import ChatOpenAI
|
||
|
||
from app.config.settings import settings
|
||
from app.model.schemas import AuthContext, IntentType
|
||
from app.service.memory_service import VisitorMemoryService
|
||
from app.service.rag_service import VisitorRagService
|
||
from app.service.visitor_prompts import (
|
||
CHITCHAT_SYSTEM,
|
||
CHITCHAT_USER_TEMPLATE,
|
||
FALLBACK_TEXT,
|
||
GENERATE_SYSTEM,
|
||
GENERATE_USER_TEMPLATE,
|
||
INTENT_SYSTEM,
|
||
INTENT_USER_TEMPLATE,
|
||
TRANSFER_TEXT,
|
||
)
|
||
from app.utils.compliance_guard import (
|
||
REJECT_ACCOUNT,
|
||
REJECT_ADVICE,
|
||
REJECT_COMPARE,
|
||
REJECT_GENERAL,
|
||
REJECT_PREDICT,
|
||
REJECT_REALTIME,
|
||
RISK_DISCLAIMER,
|
||
should_add_disclaimer,
|
||
)
|
||
from app.utils.sanitize_postprocess import finalize_sanitized_reply
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# DeepSeek LLM 客户端
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def _build_llm() -> ChatOpenAI:
|
||
"""构建 DeepSeek LLM 客户端(OpenAI 兼容)。"""
|
||
return ChatOpenAI(
|
||
model=settings.deepseek_model,
|
||
api_key=settings.deepseek_api_key,
|
||
base_url=settings.deepseek_base_url,
|
||
temperature=settings.deepseek_temperature,
|
||
max_tokens=settings.deepseek_max_tokens,
|
||
)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# LangGraph State
|
||
# ---------------------------------------------------------------------------
|
||
|
||
class VisitorState(TypedDict):
|
||
session_id: str
|
||
trace_id: str
|
||
message: str
|
||
intent: str
|
||
chitchat_memory: list[dict]
|
||
consult_memory: list[dict]
|
||
rag_context: str
|
||
rag_sources: list[dict]
|
||
reply: str
|
||
has_disclaimer: bool
|
||
transfer_to_human: bool
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 节点函数
|
||
# ---------------------------------------------------------------------------
|
||
|
||
# 意图关键词快速路由(不调 LLM)
|
||
|
||
# 转人工关键词:用户主动要求 / 投诉 / 安全问题
|
||
_TRANSFER_KEYWORDS = (
|
||
"转人工", "人工客服", "人工服务", "联系客服", "找客服",
|
||
"投诉", "不满", "纠纷", "账户异常", "被盗", "安全",
|
||
)
|
||
|
||
# 拒绝关键词:游客无法处理,直接拒绝,不转人工
|
||
_REJECT_ADVICE_KW = (
|
||
"推荐", "建议买", "建议卖", "包赚", "保本", "稳赚",
|
||
"零风险", "无风险", "躺赚", "保本保息", "刚性兑付",
|
||
)
|
||
_REJECT_ACCOUNT_KW = (
|
||
"我的持仓", "我的账户", "我的余额", "我的风评",
|
||
"我的资产", "我的交易", "我的收益", "我的理财",
|
||
)
|
||
_REJECT_PREDICT_KW = (
|
||
"会涨", "会跌", "收益预测", "走势预测",
|
||
"能赚多少", "收益率多少", "涨跌", "明天行情",
|
||
)
|
||
_REJECT_COMPARE_KW = (
|
||
"和其他", "和XX", "哪个平台好", "哪个好",
|
||
"对比一下", "比一比", "哪家好",
|
||
)
|
||
_REJECT_REALTIME_KW = (
|
||
"实时行情", "实时价格", "现在价格",
|
||
"当前价格", "最新行情",
|
||
)
|
||
|
||
# FAQ 关键词:LLM 分类失败时仍可走 RAG(开户/账户操作类)
|
||
_FAQ_KEYWORDS = (
|
||
"开户", "开账户", "办理开户", "开户资料", "开户材料",
|
||
"需要什么资料", "需要哪些资料", "需要什么材料", "需要哪些材料",
|
||
"身份认证", "实名认证", "风险测评", "风评问卷", "忘记密码", "重置密码",
|
||
)
|
||
|
||
|
||
def _merged_items(state: VisitorState) -> list[dict]:
|
||
"""consult + chitchat 记忆按 ts 合并(recall_memory 已载入 state)。"""
|
||
items: list[dict] = []
|
||
items.extend(state.get("consult_memory") or [])
|
||
items.extend(state.get("chitchat_memory") or [])
|
||
items.sort(key=lambda m: m.get("ts", 0))
|
||
return items
|
||
|
||
|
||
def _merged_memory_text(state: VisitorState, max_items: int = 16) -> str:
|
||
"""合并两类近期对话为「用户/客服」文本(供意图分类与生成 prompt 使用)。
|
||
|
||
数据源为 recall_memory 已载入的 state 记忆,不重复回源 Redis。
|
||
"""
|
||
lines: list[str] = []
|
||
for msg in _merged_items(state)[-max_items:]:
|
||
role = "用户" if msg.get("role") == "user" else "客服"
|
||
content = (msg.get("content") or "").strip()
|
||
if content:
|
||
lines.append(f"{role}: {content}")
|
||
return "\n".join(lines)
|
||
|
||
|
||
def _rag_query(state: VisitorState, max_items: int = 6) -> str:
|
||
"""历史感知检索 query:拼最近几轮原始内容,解决「那申购呢」类省略指代。"""
|
||
msg = state["message"]
|
||
recent = [(m.get("content") or "").strip() for m in _merged_items(state)[-max_items:]]
|
||
recent = [c for c in recent if c]
|
||
if not recent:
|
||
return msg
|
||
return "\n".join(recent) + "\n" + msg
|
||
|
||
|
||
def _degraded_reply_from_rag(rag_context: str) -> str | None:
|
||
"""LLM 不可用时,从 RAG 首片段提取可读答案。"""
|
||
text = (rag_context or "").strip()
|
||
if not text:
|
||
return None
|
||
parts = text.split("\n", 1)
|
||
if len(parts) < 2:
|
||
return None
|
||
body = parts[1].strip()
|
||
if len(body) < 8:
|
||
return None
|
||
return body[:400]
|
||
|
||
|
||
def recall_memory(state: VisitorState) -> VisitorState:
|
||
"""节点 1:从 Redis 读取两类短期记忆(Redis 不可用则降级为空)。"""
|
||
mem = VisitorMemoryService()
|
||
sid = state["session_id"]
|
||
try:
|
||
chitchat_memory = mem.recall(sid, "chitchat")
|
||
consult_memory = mem.recall(sid, "consult")
|
||
except Exception:
|
||
chitchat_memory = []
|
||
consult_memory = []
|
||
return {
|
||
"chitchat_memory": chitchat_memory,
|
||
"consult_memory": consult_memory,
|
||
}
|
||
|
||
|
||
def intent_classify(state: VisitorState) -> VisitorState:
|
||
"""节点 2:意图分类(先关键词快速路由,再 DeepSeek)。"""
|
||
msg = state["message"]
|
||
|
||
# 转人工关键词(最高优先级)
|
||
if any(k in msg for k in _TRANSFER_KEYWORDS):
|
||
return {"intent": "transfer_human"}
|
||
|
||
# 拒绝关键词(直接拒绝,不转人工)
|
||
if any(k in msg for k in _REJECT_ADVICE_KW):
|
||
return {"intent": "reject", "reply": REJECT_ADVICE}
|
||
if any(k in msg for k in _REJECT_ACCOUNT_KW):
|
||
return {"intent": "reject", "reply": REJECT_ACCOUNT}
|
||
if any(k in msg for k in _REJECT_PREDICT_KW):
|
||
return {"intent": "reject", "reply": REJECT_PREDICT}
|
||
if any(k in msg for k in _REJECT_COMPARE_KW):
|
||
return {"intent": "reject", "reply": REJECT_COMPARE}
|
||
if any(k in msg for k in _REJECT_REALTIME_KW):
|
||
return {"intent": "reject", "reply": REJECT_REALTIME}
|
||
|
||
if any(k in msg for k in _FAQ_KEYWORDS):
|
||
return {"intent": "faq"}
|
||
|
||
# DeepSeek 分类
|
||
try:
|
||
llm = _build_llm()
|
||
user_prompt = INTENT_USER_TEMPLATE.format(memory=_merged_memory_text(state), message=msg)
|
||
resp = llm.invoke([
|
||
{"role": "system", "content": INTENT_SYSTEM},
|
||
{"role": "user", "content": user_prompt},
|
||
])
|
||
intent = resp.content.strip().lower()
|
||
# 合法性校验
|
||
valid = {"product_consult", "policy_interpret", "faq", "chit_chat", "reject", "transfer_human", "fallback"}
|
||
if intent not in valid:
|
||
intent = "fallback"
|
||
except Exception:
|
||
intent = "fallback"
|
||
|
||
return {"intent": intent}
|
||
|
||
|
||
def rag_search(state: VisitorState) -> VisitorState:
|
||
"""节点 3:RAG 检索(product_consult/policy_interpret/faq 分支)。"""
|
||
intent = state["intent"]
|
||
rag = VisitorRagService()
|
||
context, sources = rag.retrieve(intent, _rag_query(state))
|
||
|
||
if not context:
|
||
# 空结果 → 走兜底
|
||
return {
|
||
"rag_context": "",
|
||
"rag_sources": [],
|
||
"reply": FALLBACK_TEXT,
|
||
"intent": "fallback",
|
||
}
|
||
|
||
return {"rag_context": context, "rag_sources": sources}
|
||
|
||
|
||
def generate(state: VisitorState) -> VisitorState:
|
||
"""节点 4:基于 RAG 上下文 + 短期记忆生成回复。"""
|
||
# 空上下文时跳过 LLM(兜底已在 rag_search 中处理)
|
||
if not state.get("rag_context"):
|
||
return {"reply": state.get("reply", FALLBACK_TEXT)}
|
||
|
||
# 如果已经设置了兜底回复,直接返回
|
||
if state.get("reply") and state.get("intent") == "fallback":
|
||
has_disclaimer = should_add_disclaimer("faq")
|
||
reply = state["reply"]
|
||
if has_disclaimer:
|
||
reply = f"{reply}\n\n{RISK_DISCLAIMER}"
|
||
return {"reply": reply, "has_disclaimer": has_disclaimer}
|
||
|
||
try:
|
||
llm = _build_llm()
|
||
mem_text = _merged_memory_text(state)
|
||
user_prompt = GENERATE_USER_TEMPLATE.format(
|
||
rag_context=state["rag_context"],
|
||
memory=mem_text,
|
||
message=state["message"],
|
||
)
|
||
resp = llm.invoke([
|
||
{"role": "system", "content": GENERATE_SYSTEM},
|
||
{"role": "user", "content": user_prompt},
|
||
])
|
||
reply = resp.content.strip()
|
||
except Exception:
|
||
degraded = _degraded_reply_from_rag(state.get("rag_context", ""))
|
||
if not degraded:
|
||
return {"reply": FALLBACK_TEXT, "intent": "fallback"}
|
||
reply = degraded
|
||
|
||
# 合规护栏(1B:命中不自动 transfer)
|
||
reply, need_transfer = finalize_sanitized_reply(reply, intent=state.get("intent"))
|
||
if need_transfer:
|
||
return {"reply": reply, "transfer_to_human": True, "has_disclaimer": False}
|
||
|
||
# 咨询类附带风险提示
|
||
has_disclaimer = should_add_disclaimer(state["intent"])
|
||
if has_disclaimer:
|
||
reply = f"{reply}\n\n{RISK_DISCLAIMER}"
|
||
|
||
return {"reply": reply, "has_disclaimer": has_disclaimer}
|
||
|
||
|
||
def chitchat(state: VisitorState) -> VisitorState:
|
||
"""节点 5:闲聊生成。"""
|
||
try:
|
||
llm = _build_llm()
|
||
mem_text = _merged_memory_text(state)
|
||
user_prompt = CHITCHAT_USER_TEMPLATE.format(
|
||
memory=mem_text,
|
||
message=state["message"],
|
||
)
|
||
resp = llm.invoke([
|
||
{"role": "system", "content": CHITCHAT_SYSTEM},
|
||
{"role": "user", "content": user_prompt},
|
||
])
|
||
reply = resp.content.strip()
|
||
except Exception:
|
||
return {"reply": FALLBACK_TEXT, "intent": "fallback"}
|
||
|
||
# 合规护栏
|
||
reply, need_transfer = finalize_sanitized_reply(reply, intent=state.get("intent"))
|
||
if need_transfer:
|
||
return {"reply": reply, "transfer_to_human": True, "has_disclaimer": False}
|
||
|
||
return {"reply": reply, "has_disclaimer": False}
|
||
|
||
|
||
def reject(state: VisitorState) -> VisitorState:
|
||
"""节点:直接拒绝(不转人工)。"""
|
||
# 关键词路由已设置 reply 时直接用
|
||
if state.get("reply"):
|
||
return {}
|
||
# DeepSeek 返回 reject 但未设置 reply 时用通用拒绝
|
||
return {"reply": REJECT_GENERAL}
|
||
|
||
|
||
def transfer_human(state: VisitorState) -> VisitorState:
|
||
"""节点 6:转人工。"""
|
||
return {"reply": TRANSFER_TEXT, "transfer_to_human": True}
|
||
|
||
|
||
def fallback(state: VisitorState) -> VisitorState:
|
||
"""节点 7:兜底。"""
|
||
return {"reply": FALLBACK_TEXT}
|
||
|
||
|
||
def save_memory(state: VisitorState) -> VisitorState:
|
||
"""节点 8:保存短期记忆(先写 user,再写 assistant;Redis 失败不阻断)。"""
|
||
mem = VisitorMemoryService()
|
||
sid = state["session_id"]
|
||
intent = state.get("intent", "fallback")
|
||
|
||
# 咨询类 → consult 记忆;闲聊/转人工/兜底 → chitchat 记忆
|
||
if intent in ("product_consult", "policy_interpret", "faq"):
|
||
kind = "consult"
|
||
else:
|
||
kind = "chitchat"
|
||
|
||
try:
|
||
mem.append(sid, kind, "user", state["message"])
|
||
mem.append(sid, kind, "assistant", state["reply"])
|
||
except Exception:
|
||
pass
|
||
|
||
return {}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 图构建
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def _build_graph():
|
||
g = StateGraph(VisitorState)
|
||
|
||
g.add_node("recall_memory", recall_memory)
|
||
g.add_node("intent_classify", intent_classify)
|
||
g.add_node("rag_search", rag_search)
|
||
g.add_node("generate", generate)
|
||
g.add_node("chitchat", chitchat)
|
||
g.add_node("reject", reject)
|
||
g.add_node("transfer_human", transfer_human)
|
||
g.add_node("fallback", fallback)
|
||
g.add_node("save_memory", save_memory)
|
||
|
||
g.set_entry_point("recall_memory")
|
||
g.add_edge("recall_memory", "intent_classify")
|
||
|
||
def _route(state: VisitorState) -> str:
|
||
intent = state["intent"]
|
||
if intent in ("product_consult", "policy_interpret", "faq"):
|
||
return "rag_search"
|
||
if intent == "chit_chat":
|
||
return "chitchat"
|
||
if intent == "reject":
|
||
return "reject"
|
||
if intent == "transfer_human":
|
||
return "transfer_human"
|
||
return "fallback"
|
||
|
||
g.add_conditional_edges("intent_classify", _route, {
|
||
"rag_search": "rag_search",
|
||
"chitchat": "chitchat",
|
||
"reject": "reject",
|
||
"transfer_human": "transfer_human",
|
||
"fallback": "fallback",
|
||
})
|
||
|
||
g.add_edge("rag_search", "generate")
|
||
g.add_edge("generate", "save_memory")
|
||
g.add_edge("chitchat", "save_memory")
|
||
g.add_edge("reject", "save_memory")
|
||
g.add_edge("transfer_human", "save_memory")
|
||
g.add_edge("fallback", "save_memory")
|
||
g.add_edge("save_memory", END)
|
||
return g.compile()
|
||
|
||
|
||
_GRAPH = _build_graph()
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 编排入口(对齐 agent_service.run_chat 模式)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def run_visitor_chat(
|
||
ctx: AuthContext,
|
||
message: str,
|
||
session_id: str,
|
||
) -> tuple[str, bool, str, bool]:
|
||
"""游客对话入口,返回 (reply, has_disclaimer, intent, transfer_to_human)。"""
|
||
state: VisitorState = {
|
||
"session_id": session_id,
|
||
"trace_id": ctx.trace_id,
|
||
"message": message,
|
||
"intent": "",
|
||
"chitchat_memory": [],
|
||
"consult_memory": [],
|
||
"rag_context": "",
|
||
"rag_sources": [],
|
||
"reply": "",
|
||
"has_disclaimer": False,
|
||
"transfer_to_human": False,
|
||
}
|
||
result = _GRAPH.invoke(state)
|
||
return (
|
||
result["reply"],
|
||
result.get("has_disclaimer", False),
|
||
result.get("intent", "fallback"),
|
||
result.get("transfer_to_human", False),
|
||
)
|