"""Runtime assembly for the client Agent.""" from __future__ import annotations import contextvars import logging import uuid from types import SimpleNamespace from nl2sql.contracts import DataQueryRequest from nl2sql.history import archive_query_safely from nl2sql.limits import QueryLimiter from nl2sql.retrieval import retrieve_metadata from nl2sql.runtime_config import runtime_config from nl2sql.schema import load_authoritative_schema 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, DataQueryRejected, ) from service.client_agent.memory_extractor import DialogueMemoryExtractor from service.memory.facade import MemoryService from service.memory.schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage from service.nl2sql.answer_render import render_query_answer from service.nl2sql.customer_permission import load_customer_query_permission from service.nl2sql.query_service import QueryServiceError, execute_query _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, customer_id=customer_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_data_query(*, db_session_factory, milvus_client, llm_client, config_getter, redis, schema_database): """构造登录客户的 NL2SQL 数据查询依赖。 身份与行级范围由服务端强制注入:data_scope 只带当前登录客户自己的 customer_id,配合权限快照的 row_scopes 在 SQL 层兜底,杜绝水平越权。 """ async def data_query(*, question, customer_id, session_id, trace_id): async with db_session_factory() as db: permission = await load_customer_query_permission( db, customer_id, config_getter=config_getter ) if not permission.get("can_query"): raise DataQueryRejected("数据查询功能暂未开放,您可以先咨询基金知识或开户流程~") limiter = QueryLimiter(redis) if not await limiter.acquire( customer_id, daily_quota=permission.get("daily_quota", 0) or 20, max_concurrent=1, rate_limit=10, ): raise DataQueryRejected("您今天的数据查询次数已达上限,请明天再来吧~") try: return await _execute_customer_query( db=db, permission=permission, question=question, customer_id=customer_id, session_id=session_id, trace_id=trace_id, ) finally: await limiter.release(customer_id) async def _execute_customer_query(*, db, permission, question, customer_id, session_id, trace_id): query_id = uuid.uuid4().hex async def permission_loader(_user_id: int): return permission async def metadata_retriever(retrieval_question: str): return await retrieve_metadata( retrieval_question, milvus_client, top_k=runtime_config.retrieval_top_k ) async def schema_loader(table_names: set[str], _permission: dict): return await load_authoritative_schema( db, database=schema_database, candidate_tables=table_names ) request = DataQueryRequest( question=question, user_id=customer_id, trace_id=trace_id, session_id=session_id, caller_agent="client_agent", # 行级范围只允许是登录客户本人,不接受任何外部输入 data_scope={"customer_ids": [customer_id]}, include_sql=False, max_rows=min(permission.get("max_rows") or 200, runtime_config.max_rows), ) try: result = await execute_query( request, session=db, query_id=query_id, permission_loader=permission_loader, metadata_retriever=metadata_retriever, schema_loader=schema_loader, llm_client=llm_client, masks=permission.get("masks") or {}, summary_llm=llm_client, ) except QueryServiceError as exc: await archive_query_safely( db, query_id=query_id, user_id=customer_id, question=question, status="blocked", error_message=str(exc), trace_id=trace_id, session_id=session_id, caller_agent="client_agent", ) # 不向用户暴露 SQL 和内部异常细节 raise DataQueryRejected( "这个问题我暂时查不了,您可以换个问法,或者联系人工客服帮您处理~" ) from exc await archive_query_safely( db, query_id=query_id, user_id=customer_id, question=question, status="success", row_count=result.row_count, truncated=result.truncated, elapsed_ms=result.elapsed_ms, trace_id=trace_id, session_id=session_id, caller_agent="client_agent", ) return { "answer": render_query_answer(result), "sources": [], "query_id": result.query_id, "row_count": result.row_count, "truncated": result.truncated, "chart": result.chart, } return data_query def build_client_runtime( *, redis, milvus_client, llm_client, config_getter, audit_writer, db_session_factory=None, schema_database=None, 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, 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() 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, ) data_query = None if db_session_factory is not None and schema_database: data_query = _build_data_query( db_session_factory=db_session_factory, milvus_client=milvus_client, llm_client=llm_client, config_getter=config_getter, redis=redis, schema_database=schema_database, ) agent = AnonymousCustomerAgent( context=context, rag_retrieve=retrieve, intent_recognize=recognize, generate_answer=generate, audit_writer=audit_writer, config_getter=config_getter, data_query=data_query, ) 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, )