Files
group_fqcd_jr/app/service/agent_run_application_service.py
T

245 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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.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:
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)