import hashlib import json from collections.abc import Sequence from dataclasses import dataclass from datetime import UTC, datetime, timedelta from uuid import uuid4 from sqlalchemy import select 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, SessionNotAccessibleError, ) from app.model.audit import InteractionAudit from app.model.conversation import ConversationMessage from app.model.platform import AgentRun, RequestIdempotency from app.model.session import ConversationSession from app.repository.outbox_repository import OutboxRepository from app.service.advisor_rollout_service import AdvisorRolloutService 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) class RunAccepted: run_id: str trace_id: str 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], *, clarification_round: int = 0, session_context: Sequence[str] = (), ) -> dict[str, object]: """构造 Worker 使用的内部元数据,不信任外部传入的闲聊计数。""" metadata = request.metadata if request.agent_type == "customer_service": 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, 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: if request.agent_type == "advisor": await AdvisorRolloutService(self.session).ensure_allowed(context) try: self.factory.authorize(request.agent_type, context) except ForbiddenAgentError: async with self.session.begin(): self.session.add(InteractionAudit( actor_type="user", actor_id=int(context.user_id), portal=context.portal, action_type="agent.access_denied", session_id=request.session_id, detail={"agent_type": request.agent_type, "trace_id": context.trace_id}, 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(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(): session_row = await self.session.scalar( select(ConversationSession).where( ConversationSession.session_id == request.session_id, ConversationSession.user_id == user_id, ).with_for_update() ) if session_row is not None: if session_row.status != "active" or session_row.agent_type != request.agent_type: raise ForbiddenAgentError("会话不可用于当前 Agent") session_row.message_count += 1 session_row.last_active_at = now owner = await self.session.scalar( select(ConversationMessage.customer_id) .where(ConversationMessage.session_id == request.session_id) .where(ConversationMessage.customer_id.is_not(None)) .limit(1) ) if owner is not None and owner != user_id: raise SessionNotAccessibleError("会话不属于当前用户") existing = await self.session.scalar( select(RequestIdempotency).where( RequestIdempotency.user_id == user_id, RequestIdempotency.agent_type == request.agent_type, RequestIdempotency.idempotency_key == request.idempotency_key, ) ) if existing is not None: if existing.request_hash != request_hash: raise IdempotencyConflictError("同一幂等键对应不同请求") run = await self.session.scalar( select(AgentRun).where(AgentRun.idempotency_id == existing.id) ) if run is None: raise RuntimeError("idempotency record has no run") return RunAccepted(run.run_id, run.trace_id, run.status) 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( 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=stored_request.message, trace_id=trace_id, created_at=now, ) self.session.add(message) await self.session.flush() idem = RequestIdempotency( user_id=user_id, session_id=request.session_id, agent_type=request.agent_type, idempotency_key=request.idempotency_key, request_hash=request_hash, trace_id=trace_id, expire_at=now + timedelta(hours=24), created_at=now, updated_at=now, ) self.session.add(idem) try: await self.session.flush() except IntegrityError as exc: raise IdempotencyConflictError("幂等键正在被并发请求占用") from exc run_id = str(uuid4()) run = AgentRun( run_id=run_id, idempotency_id=idem.id, session_id=request.session_id, user_id=user_id, agent_type=request.agent_type, trace_id=trace_id, request_message_id=message.id, created_at=now, updated_at=now, ) self.session.add(run) await self.session.flush() await OutboxRepository(self.session).append(DomainEvent( event_id=str(uuid4()), event_type="agent.run_requested", aggregate_type="agent_run", aggregate_id=run_id, trace_id=trace_id, payload={ "run_id": run_id, "actor_type": "visitor" if "visitor" in context.roles else "authenticated", "metadata": outbox_metadata, }, occurred_at=now, )) return RunAccepted(run_id, trace_id)