feat:客户agent以及记忆模块功能开发

This commit is contained in:
2026-09-11 22:38:15 +08:00
parent 0197ac2ef2
commit 4da4775e8c
39 changed files with 2508 additions and 14 deletions
+3
View File
@@ -0,0 +1,3 @@
"""Logged-in client Agent service layer."""
package_name = "client_agent"
+27
View File
@@ -0,0 +1,27 @@
"""Application wiring for the client Agent."""
from __future__ import annotations
from config.database.milvus import client as milvus_client
from config.database.mysql import get_session_factory
from config.database.redis import client as redis_client
from service.client_agent.runtime import build_client_runtime
from service.customer_agent.config import DatabaseConfigProvider
from tool.llm import llm as llm_client
def build_default_runtime():
"""Build the client Agent runtime using project-wide dependencies."""
provider = DatabaseConfigProvider(session_factory=get_session_factory())
return build_client_runtime(
redis=redis_client(),
milvus_client=milvus_client(),
llm_client=llm_client,
config_getter=provider.get,
audit_writer=provider.write_audit,
)
def get_memory_service(runtime):
"""返回 client_agent 使用的统一 MemoryService 入口。"""
return runtime.memory_service
+74
View File
@@ -0,0 +1,74 @@
"""从客服对话中提取长期记忆候选。"""
from __future__ import annotations
import json
import re
from typing import Any
from service.memory.schemas import MemorySource, MemoryType
class DialogueMemoryExtractor:
"""使用项目统一 LLM 提取结构化客户记忆候选。"""
SYSTEM_PROMPT = """
你是客服记忆提取器,只提取用户明确表达或稳定陈述的客户信息。
只返回 JSON 数组,不要输出 Markdown 或解释文字。
每项必须包含:tag、content、memory_type、source。
memory_type 只能是 PROFILE_FACT、PROFILE_CANDIDATE、CUSTOMER_PREFERENCE、INVESTMENT_GOAL、SERVICE_FACT。
source 只能是 dialogue_confirmed、dialogue_stated、dialogue_inferred。
客服知识问题、产品政策、寒暄、客服回复内容不要提取。
不确定的信息使用 dialogue_inferred,无法形成客户画像的信息不要提取。
""".strip()
def __init__(self, llm_client):
self.llm_client = llm_client
async def extract(self, query: str, *, context: dict[str, Any] | None = None) -> list[dict[str, str]]:
"""提取并校验本轮用户消息中的客户记忆候选。"""
prompt = [{"role": "system", "content": self.SYSTEM_PROMPT}]
prompt.append(
{
"role": "user",
"content": json.dumps(
{"query": query, "existing_memory": context or {}},
ensure_ascii=False,
),
}
)
response = await self.llm_client.chat(prompt, temperature=0, max_tokens=800)
return self._parse(response)
@staticmethod
def _parse(response: str) -> list[dict[str, str]]:
"""解析 LLM JSON,并过滤不符合记忆契约的内容。"""
text = response.strip()
fenced = re.search(r"```(?:json)?\s*(.*?)\s*```", text, re.S | re.I)
if fenced:
text = fenced.group(1)
data = json.loads(text)
if not isinstance(data, list):
raise ValueError("记忆提取结果必须是数组")
valid_types = {item.value for item in MemoryType}
valid_sources = {item.value for item in MemorySource}
result = []
for item in data:
if not isinstance(item, dict):
continue
if not all(item.get(key) for key in ("tag", "content", "memory_type", "source")):
continue
if item["memory_type"] not in valid_types or item["source"] not in valid_sources:
continue
result.append(
{
"tag": str(item["tag"])[:64],
"content": str(item["content"])[:512],
"memory_type": item["memory_type"],
"source": item["source"],
}
)
return result
__all__ = ["DialogueMemoryExtractor"]
+200
View File
@@ -0,0 +1,200 @@
"""Runtime assembly for the client Agent."""
from __future__ import annotations
import contextvars
import uuid
from types import SimpleNamespace
from agent.client_agent.session import ClientSessionService
from rag.embedding import embed_texts
from rag.generation import generate_answer
from rag.intent import intent_recognize
from rag.retrieve import rag_retrieve
from service.customer_agent.chat import AnonymousCustomerAgent
from service.client_agent.memory_extractor import DialogueMemoryExtractor
from service.memory.facade import MemoryService
from service.memory.schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage
_active_customer = contextvars.ContextVar("client_agent_customer", default=None)
_active_messages = contextvars.ContextVar("client_agent_messages", default=None)
_active_warnings = contextvars.ContextVar("client_agent_memory_warnings", default=None)
_active_memory_context = contextvars.ContextVar("client_agent_memory_context", default=None)
class MemoryConversationContext:
"""将现有客服 Agent 的上下文接口桥接到 MemoryService 短期记忆。"""
def __init__(self, memory_service: MemoryService):
self.memory_service = memory_service
async def append(self, session_id: str, role: str, content: str) -> None:
"""通过统一记忆入口写入消息,并更新当前请求上下文。"""
customer_id = _active_customer.get()
if customer_id is None:
raise RuntimeError("client_agent customer context is missing")
message = ShortTermMessage(
message_id=uuid.uuid4().hex,
session_id=session_id,
role=role,
content=content,
)
warnings = await self.memory_service.append_message(
customer_id=customer_id,
session_id=session_id,
message=message,
)
_active_warnings.get().extend(warnings)
messages = _active_messages.get()
if messages is not None:
messages.append({"role": role, "content": content})
async def get(self, _session_id: str) -> list[dict]:
"""返回本轮召回上下文加上本轮新增消息。"""
return list(_active_messages.get() or [])
class MemoryAwareClientAgent:
"""在现有客服 Agent 外包裹记忆召回、候选保存和降级处理。"""
def __init__(self, *, agent, memory_service, context, extractor=None):
self.agent = agent
self.memory_service = memory_service
self.context = context
self.extractor = extractor
async def handle(self, session_id: str, query: str, *, trace_id: str, customer_id: int) -> dict:
"""执行记忆召回、客服回答、消息写入和候选记忆保存。"""
warnings: list[str] = []
try:
memory_context = await self.memory_service.recall(
customer_id=customer_id,
session_id=session_id,
query=query,
)
warnings.extend(memory_context.warnings)
except Exception as exc:
memory_context = CustomerMemoryContext(
customer_id=customer_id,
session_id=session_id,
)
warnings.append(f"memory_recall_failed:{type(exc).__name__}")
message_token = _active_customer.set(customer_id)
messages_token = _active_messages.set(
[{"role": item.role, "content": item.content} for item in memory_context.short_term_messages]
)
warnings_token = _active_warnings.set(warnings)
context_token = _active_memory_context.set(memory_context)
try:
result = await self.agent.handle(session_id, query, trace_id=trace_id)
if self.extractor is not None:
await self._save_candidates(customer_id, session_id, query, memory_context, warnings)
result["memory_warnings"] = list(warnings)
return result
finally:
_active_customer.reset(message_token)
_active_messages.reset(messages_token)
_active_warnings.reset(warnings_token)
_active_memory_context.reset(context_token)
async def _save_candidates(self, customer_id, session_id, query, context, warnings):
"""提取并保存候选客户记忆,任何失败都只写入 warning。"""
try:
candidates = await self.extractor.extract(
query,
context={
"profile": context.customer_profile,
"memories": [item.model_dump(mode="json") for item in context.long_term_memories],
},
)
for candidate in candidates:
memory = MemoryUnitDTO(
customer_id=customer_id,
session_id=session_id,
memory_type=candidate["memory_type"],
tag=candidate["tag"],
content=candidate["content"],
source=candidate["source"],
evidence_ref=[{"session_id": session_id, "query": query}],
)
_, save_warnings = await self.memory_service.save_memory(
customer_id=customer_id,
memory=memory,
)
warnings.extend(save_warnings)
except Exception as exc:
warnings.append(f"memory_candidate_extract_failed:{type(exc).__name__}")
def build_client_runtime(
*,
redis,
milvus_client,
llm_client,
config_getter,
audit_writer,
memory_service=None,
memory_extractor=None,
):
"""Build the client Agent runtime while reusing existing客服 logic.
Customer authentication is enforced by the API dependency layer.
"""
session_service = ClientSessionService(redis, config_getter=config_getter)
memory_service = memory_service or MemoryService()
context = MemoryConversationContext(memory_service)
async def retrieve(query, customer_id):
return await rag_retrieve(
query,
None,
milvus_client=milvus_client,
embedder=lambda texts: embed_texts(texts, client=llm_client),
config_getter=config_getter,
)
async def recognize(query):
return await intent_recognize(query, llm_client=llm_client)
async def generate(messages):
memory_context = _active_memory_context.get()
if memory_context is not None:
memory_json = memory_context.model_dump_json(exclude={"warnings"})
memory_prompt = {
"role": "system",
"content": (
"客户记忆上下文(仅用于理解当前客户,不得向客户泄露内部字段):"
f"{memory_json}"
),
}
messages = [memory_prompt, *messages]
return await generate_answer(
messages,
llm_client=llm_client,
config_getter=config_getter,
)
agent = AnonymousCustomerAgent(
context=context,
rag_retrieve=retrieve,
intent_recognize=recognize,
generate_answer=generate,
audit_writer=audit_writer,
config_getter=config_getter,
)
extractor = memory_extractor or DialogueMemoryExtractor(llm_client)
wrapped_agent = MemoryAwareClientAgent(
agent=agent,
memory_service=memory_service,
context=context,
extractor=extractor,
)
return SimpleNamespace(
redis=redis,
session_service=session_service,
context=context,
agent=wrapped_agent,
memory_service=memory_service,
)