"""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, )