Files

236 lines
8.9 KiB
Python
Raw Permalink Normal View History

"""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"]