"""Redis 短期会话记忆。""" from __future__ import annotations import asyncio import json import time import uuid from datetime import datetime, timezone from inspect import isawaitable from typing import Any, Callable from config.database.redis import client as redis_client from .schemas import ShortTermMessage class ShortTermMemoryError(RuntimeError): """短期记忆操作失败。""" class SessionExpiredError(ShortTermMemoryError): """会话已超过最长生命周期。""" async def _config(config_getter, key: str, default): """读取项目配置,并将数据库配置值转换为默认值类型。""" value = config_getter(key, str(default)) if isawaitable(value): value = await value return type(default)(value) class ShortTermMemory: """管理当前客服会话的 Redis 消息、Token 预算和生命周期。""" MESSAGE_KEY = "session:{session_id}:messages" META_KEY = "session:{session_id}:meta" ALLOWED_ROLES = frozenset({"user", "assistant", "system"}) def __init__( self, redis=None, *, config_getter=None, token_counter: Callable[[str], int] | None = None, clock=time.time, fail_soft: bool = True, ): """创建短期记忆服务,默认复用项目级 Redis 客户端。""" self.redis = redis or redis_client() self.config_getter = config_getter or (lambda _key, default: default) self.token_counter = token_counter or (lambda text: max(1, len(text) // 4)) self.clock = clock self.fail_soft = fail_soft self.last_warnings: list[str] = [] self._locks: dict[str, asyncio.Lock] = {} @classmethod def message_key(cls, session_id: str) -> str: """生成会话消息列表 Key。""" return cls.MESSAGE_KEY.format(session_id=session_id) @classmethod def meta_key(cls, session_id: str) -> str: """生成会话元数据 Hash Key。""" return cls.META_KEY.format(session_id=session_id) def _lock(self, session_id: str) -> asyncio.Lock: """获取进程内会话锁,避免并发截断互相覆盖。""" return self._locks.setdefault(session_id, asyncio.Lock()) async def append_message( self, session_id: str, role: str, content: str, *, message_id: str | None = None, agent_run_id: str | None = None, tool_calls: list[dict[str, Any]] | None = None, ) -> ShortTermMessage | None: """追加消息、刷新 TTL,并按 Token 预算从旧到新截断。""" self.last_warnings = [] if role not in self.ALLOWED_ROLES: raise ValueError(f"非法消息角色: {role}") if not isinstance(content, str) or not content.strip(): raise ValueError("消息内容不能为空") message = ShortTermMessage( message_id=message_id or uuid.uuid4().hex, session_id=session_id, role=role, content=content, token_count=self.token_counter(content), agent_run_id=agent_run_id, tool_calls=tool_calls or [], create_time=datetime.fromtimestamp(self.clock(), tz=timezone.utc), ) try: async with self._lock(session_id): await self._ensure_meta(session_id) await self.redis.rpush( self.message_key(session_id), self._serialize(message) ) await self._touch(session_id) await self._trim(session_id) return message except SessionExpiredError: raise except Exception as exc: return self._degrade("append_message", exc) async def load_messages(self, session_id: str) -> list[ShortTermMessage]: """按写入顺序读取当前会话消息,并刷新空闲 TTL。""" self.last_warnings = [] try: if not await self._session_is_active(session_id): return [] raw_messages = await self.redis.lrange(self.message_key(session_id), 0, -1) await self._touch(session_id) return [self._deserialize(raw) for raw in raw_messages if raw] except Exception as exc: return self._degrade("load_messages", exc) or [] async def get_message_count(self, session_id: str) -> int: """返回当前会话消息数量。""" return len(await self.load_messages(session_id)) async def get_token_count(self, session_id: str) -> int: """返回当前会话消息的 Token 估算总数。""" return sum(message.token_count for message in await self.load_messages(session_id)) async def clear_session(self, session_id: str) -> None: """清理会话消息和元数据。""" self.last_warnings = [] try: await self.redis.delete(self.message_key(session_id), self.meta_key(session_id)) except Exception as exc: self._degrade("clear_session", exc) async def _ensure_meta(self, session_id: str) -> None: """初始化会话元数据,并固定 24 小时绝对过期时间。""" meta_key = self.meta_key(session_id) meta = await self.redis.hgetall(meta_key) now = self.clock() if meta: absolute_expire_at = float(meta.get("absolute_expire_at", now)) if absolute_expire_at <= now: await self.clear_session(session_id) raise SessionExpiredError("会话已超过最长生命周期") return max_lifetime = await _config( self.config_getter, "agent.customer.session.max_lifetime", 86400 ) await self.redis.hset( meta_key, mapping={ "created_at": str(now), "absolute_expire_at": str(now + max_lifetime), }, ) await self.redis.expire(meta_key, max_lifetime) async def _session_is_active(self, session_id: str) -> bool: """检查会话元数据是否存在且未达到绝对过期时间。""" meta = await self.redis.hgetall(self.meta_key(session_id)) if not meta: return False if float(meta.get("absolute_expire_at", 0)) <= self.clock(): await self.clear_session(session_id) return False return True async def _touch(self, session_id: str) -> None: """刷新空闲 TTL,但不超过绝对过期时间。""" meta = await self.redis.hgetall(self.meta_key(session_id)) if not meta: return remaining = int(float(meta["absolute_expire_at"]) - self.clock()) if remaining <= 0: await self.clear_session(session_id) raise SessionExpiredError("会话已超过最长生命周期") idle_ttl = await _config( self.config_getter, "agent.customer.session.ttl", 1800 ) archive_grace = await _config( self.config_getter, "agent.customer.session.archive_grace", 60 ) ttl = min(idle_ttl + archive_grace, remaining) await self.redis.expire(self.message_key(session_id), ttl) await self.redis.expire(self.meta_key(session_id), remaining) async def _trim(self, session_id: str) -> None: """保留最新消息,确保最新一条消息不会因超预算被删除。""" limit = await _config( self.config_getter, "agent.customer.session_max_token", 4096 ) key = self.message_key(session_id) raw_messages = [raw for raw in await self.redis.lrange(key, 0, -1) if raw] total = 0 keep_from = len(raw_messages) for index in range(len(raw_messages) - 1, -1, -1): message = self._deserialize(raw_messages[index]) total += message.token_count keep_from = index if total > limit: keep_from = index + 1 break if raw_messages and keep_from == len(raw_messages): keep_from = len(raw_messages) - 1 await self.redis.ltrim(key, keep_from, -1) @staticmethod def _serialize(message: ShortTermMessage) -> str: """将消息转换为 Redis List 中的 JSON 字符串。""" return json.dumps(message.model_dump(mode="json"), ensure_ascii=False) @staticmethod def _deserialize(raw: str | bytes) -> ShortTermMessage: """将 Redis JSON 字符串恢复为消息 DTO。""" if isinstance(raw, bytes): raw = raw.decode("utf-8") return ShortTermMessage.model_validate(json.loads(raw)) def _degrade(self, operation: str, exc: Exception): """记录降级原因;fail_soft 模式下不让 Redis 故障击穿客服请求。""" warning = f"short_term_{operation}_degraded:{type(exc).__name__}" self.last_warnings = [warning] if not self.fail_soft: raise ShortTermMemoryError(warning) from exc return None __all__ = ["SessionExpiredError", "ShortTermMemory", "ShortTermMemoryError"]