Files
Mutual_Fund/service/client_agent/runtime.py
T

464 lines
17 KiB
Python

"""Runtime assembly for the client Agent."""
from __future__ import annotations
import asyncio
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 外包裹记忆召回、候选保存和降级处理。
候选记忆保存默认放入后台任务执行(background_saves=True),不阻塞
客服响应;测试或需要确定性顺序的场景可设为 False 改回同步执行。
"""
def __init__(self, *, agent, memory_service, context, extractor=None,
background_saves: bool = True):
self.agent = agent
self.memory_service = memory_service
self.context = context
self.extractor = extractor
self.background_saves = background_saves
self._pending_saves: set[asyncio.Task] = set()
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 and not self.background_saves:
await self._save_candidates(
customer_id,
session_id,
query,
memory_context,
warnings,
trace_id=trace_id,
)
result["memory_warnings"] = list(warnings)
if self.extractor is not None and self.background_saves:
self._spawn_save(
customer_id,
session_id,
query,
memory_context,
warnings,
trace_id=trace_id,
)
return result
finally:
_active_customer.reset(message_token)
_active_messages.reset(messages_token)
_active_warnings.reset(warnings_token)
_active_memory_context.reset(context_token)
def _spawn_save(
self,
customer_id,
session_id,
query,
context,
warnings,
*,
trace_id: str | None = None,
) -> None:
"""把候选记忆保存放入后台任务;任务异常已自捕获,不击穿响应。"""
task = asyncio.create_task(
self._save_candidates(
customer_id,
session_id,
query,
context,
warnings,
trace_id=trace_id,
)
)
self._pending_saves.add(task)
task.add_done_callback(self._pending_saves.discard)
async def wait_for_pending_saves(self) -> None:
"""等待全部后台保存完成,供测试与优雅退出使用。"""
if self._pending_saves:
await asyncio.gather(*list(self._pending_saves), return_exceptions=True)
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"
try:
_, reached = await self.memory_service.record_interest_signal(
customer_id=customer_id, tag=candidate["tag"]
)
except Exception as exc:
logger.exception(
"client interest signal failed: trace_id=%s customer_id=%s tag=%s",
trace_id,
customer_id,
candidate.get("tag"),
)
warnings.append(f"interest_signal_failed:{type(exc).__name__}")
continue
if not reached:
continue
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,
)