feat: complete customer service safety and handover flow

This commit is contained in:
张胜宇
2026-09-11 16:11:30 +08:00
parent 0059701509
commit ef098e6a4b
30 changed files with 1868 additions and 50 deletions
+102 -16
View File
@@ -10,6 +10,7 @@ from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.contracts import AgentRequest, DomainEvent, RequestContext
from app.core.conversation_privacy import sanitize_customer_service_message
from app.core.errors import (
ForbiddenAgentError,
IdempotencyConflictError,
@@ -23,6 +24,10 @@ from app.repository.outbox_repository import OutboxRepository
from app.service.agent.bootstrap import get_agent_factory
from app.service.agent.customer_service_routing import CustomerServiceIntentRouter
from app.service.agent.factory import AgentFactory
from app.service.customer_service_session_memory_service import (
CustomerServiceSessionMemory,
build_customer_service_session_memory,
)
@dataclass(frozen=True)
@@ -32,8 +37,17 @@ class RunAccepted:
status: str = "queued"
@dataclass(frozen=True)
class CustomerServicePriorContext:
"""受理阶段取得的短期上下文,仅包含脱敏内容和当前会话需要的最小轮次。"""
user_messages: tuple[str, ...] = ()
session_context: tuple[str, ...] = ()
def build_outbox_metadata(
request: AgentRequest, prior_user_messages: Sequence[str]
request: AgentRequest, prior_user_messages: Sequence[str], *,
clarification_round: int = 0, session_context: Sequence[str] = (),
) -> dict[str, object]:
"""构造 Worker 使用的内部元数据,不信任外部传入的闲聊计数。"""
metadata = request.metadata
@@ -41,15 +55,77 @@ def build_outbox_metadata(
metadata = metadata.model_copy(update={
"chitchat_streak": CustomerServiceIntentRouter.chitchat_streak(
prior_user_messages, request.message
)
),
"clarification_round": clarification_round,
"session_context": tuple(session_context[-6:]),
})
return metadata.model_dump(mode="json")
class AgentRunApplicationService:
def __init__(self, session: AsyncSession, factory: AgentFactory | None = None) -> None:
def __init__(
self, session: AsyncSession, factory: AgentFactory | None = None,
session_memory: CustomerServiceSessionMemory | None = None,
) -> None:
self.session = session
self.factory = factory if factory is not None else get_agent_factory()
# 客服 Redis List 只保存当前会话的脱敏上下文,绝不接入长期客户记忆。
self.session_memory = (
session_memory
if session_memory is not None
else build_customer_service_session_memory()
)
async def _read_customer_service_short_context(
self, *, request: AgentRequest, user_id: int
) -> CustomerServicePriorContext | None:
"""读取 Redis 短期上下文;返回 ``None`` 时才表示需要在事务内回退 MySQL。"""
memory_read = await self.session_memory.read(
actor_id=str(user_id), session_id=request.session_id
)
if memory_read.degraded:
return None
context = tuple(
f"{'用户' if turn.role == 'user' else '助手'}:"
f"{sanitize_customer_service_message(turn.content)}"
for turn in memory_read.turns[-6:]
)
return CustomerServicePriorContext(
user_messages=memory_read.recent_user_messages[-3:], session_context=context,
)
async def _load_mysql_prior_context(self, request: AgentRequest) -> CustomerServicePriorContext:
"""Redis 故障时读取 MySQL 已脱敏历史;调用方必须已处于受理事务。"""
mysql_messages = list(await self.session.scalars(
select(ConversationMessage.content)
.where(
ConversationMessage.session_id == request.session_id,
ConversationMessage.role == "user",
)
.order_by(ConversationMessage.id.desc())
.limit(3)
))
user_messages = tuple(
sanitize_customer_service_message(str(message))
for message in reversed(mysql_messages)
)
return CustomerServicePriorContext(
user_messages=user_messages,
session_context=tuple(f"用户:{message}" for message in user_messages),
)
async def _load_customer_service_prior_context(
self, *, request: AgentRequest, user_id: int
) -> CustomerServicePriorContext:
"""为测试和非事务调用提供完整降级路径。
Redis 空 List 表示短期会话已过期或尚未写入,不能回退历史 MySQL;只有 Redis
故障才允许降级回退,防止 30 分钟会话边界在不知情的情况下被历史消息绕过。
"""
short_context = await self._read_customer_service_short_context(
request=request, user_id=user_id
)
return short_context or await self._load_mysql_prior_context(request)
async def accept(self, request: AgentRequest, context: RequestContext) -> RunAccepted:
try:
@@ -63,9 +139,21 @@ class AgentRunApplicationService:
created_at=datetime.now(UTC).replace(tzinfo=None),
))
raise
# 客服原文不会进入会话或异步链路;安全关键词保留给后续路由生成风险提示。
stored_request = request
if request.agent_type == "customer_service":
stored_request = request.model_copy(update={
"message": sanitize_customer_service_message(request.message)
})
user_id = int(context.user_id)
# Redis 是会话体验的可选依赖,不能在下方的 MySQL 行锁事务里等待网络 I/O。
short_context = (
await self._read_customer_service_short_context(request=request, user_id=user_id)
if request.agent_type == "customer_service"
else CustomerServicePriorContext()
)
request_hash = hashlib.sha256(
json.dumps(request.model_dump(mode="json"), sort_keys=True).encode("utf-8")
json.dumps(stored_request.model_dump(mode="json"), sort_keys=True).encode("utf-8")
).hexdigest()
now = datetime.now(UTC).replace(tzinfo=None)
async with self.session.begin():
@@ -105,24 +193,22 @@ class AgentRunApplicationService:
raise RuntimeError("idempotency record has no run")
return RunAccepted(run.run_id, run.trace_id, run.status)
# 仅读取当前会话最近三条用户消息;第四条闲聊即触发一次自然业务引导。
prior_user_messages = list(await self.session.scalars(
select(ConversationMessage.content)
.where(
ConversationMessage.session_id == request.session_id,
ConversationMessage.role == "user",
)
.order_by(ConversationMessage.id.desc())
.limit(3)
))
clarification_round = session_row.clarification_round if session_row is not None else 0
# Redis 短期会话与 MySQL 降级统一返回时间正序,连续闲聊计数不依赖存储实现。
if short_context is None:
prior_context = await self._load_mysql_prior_context(request)
else:
prior_context = short_context
outbox_metadata = build_outbox_metadata(
request, tuple(reversed(prior_user_messages))
stored_request, prior_context.user_messages,
clarification_round=clarification_round,
session_context=prior_context.session_context,
)
trace_id = context.trace_id
message = ConversationMessage(
session_id=request.session_id, customer_id=user_id, portal="api",
role="user", content=request.message, trace_id=trace_id, created_at=now,
role="user", content=stored_request.message, trace_id=trace_id, created_at=now,
)
self.session.add(message)
await self.session.flush()