253 lines
9.3 KiB
Python
253 lines
9.3 KiB
Python
"""Runtime assembly for the client Agent."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import contextvars
|
|
import logging
|
|
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)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
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,
|
|
trace_id=trace_id,
|
|
)
|
|
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,
|
|
*,
|
|
trace_id: str | None = None,
|
|
):
|
|
"""提取并保存候选客户记忆,单个阶段或候选失败不影响客服回答。"""
|
|
try:
|
|
candidates = await self.extractor.extract(
|
|
query,
|
|
context={
|
|
"profile": context.customer_profile,
|
|
"memories": [self._memory_to_dict(item) for item in context.long_term_memories],
|
|
},
|
|
)
|
|
except Exception as exc:
|
|
logger.exception(
|
|
"client memory extraction failed: trace_id=%s customer_id=%s session_id=%s",
|
|
trace_id,
|
|
customer_id,
|
|
session_id,
|
|
)
|
|
warnings.append(f"memory_extraction_failed:{type(exc).__name__}")
|
|
return
|
|
|
|
for index, candidate in enumerate(candidates):
|
|
try:
|
|
evidence_count = 1
|
|
if candidate.get("signal_type") == "interest_query":
|
|
# 兴趣主题首次出现即保存为候选,后续由长期记忆按精确内容合并证据。
|
|
candidate["memory_type"] = "CUSTOMER_PREFERENCE"
|
|
candidate["source"] = "dialogue_inferred"
|
|
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}],
|
|
evidence_count=evidence_count,
|
|
)
|
|
_, save_warnings = await self.memory_service.save_memory(
|
|
customer_id=customer_id,
|
|
memory=memory,
|
|
)
|
|
warnings.extend(save_warnings)
|
|
except Exception as exc:
|
|
logger.exception(
|
|
"client memory candidate save failed: trace_id=%s customer_id=%s session_id=%s index=%s",
|
|
trace_id,
|
|
customer_id,
|
|
session_id,
|
|
index,
|
|
)
|
|
warnings.append(f"memory_save_failed:{type(exc).__name__}")
|
|
|
|
@staticmethod
|
|
def _memory_to_dict(item) -> dict:
|
|
"""将召回记忆兼容转换为可序列化字典。"""
|
|
if hasattr(item, "model_dump"):
|
|
return item.model_dump(mode="json")
|
|
if isinstance(item, dict):
|
|
return dict(item)
|
|
raise TypeError(f"unsupported memory context item: {type(item).__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,
|
|
)
|