236 lines
8.9 KiB
Python
236 lines
8.9 KiB
Python
"""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"]
|