464 lines
17 KiB
Python
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, redis=redis
|
|
)
|
|
|
|
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,
|
|
)
|