Files
Mutual_Fund/service/client_agent/runtime.py
T

258 lines
9.4 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, 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,
)
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,
)