Files
Video-Tree-TRM5/adapters/redis_cache.py

129 lines
3.9 KiB
Python
Raw Permalink 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.
"""Redis 响应缓存 — content-addressed SHA256 键,Redis 不可用时静默降级。"""
from __future__ import annotations
import dataclasses
import hashlib
import json
from typing import Any
from loguru import logger
from core.types import LLMResponse
def _resolve_cache_ttl(ttl: int) -> int:
"""校验 Redis 缓存 TTL:必须为正整数(消灭 0=永不过期 的隐式语义)。
Args:
ttl: 待校验的 TTL 秒数。
Returns:
校验通过的正整数 TTL。
Raises:
ValueError: ttl <= 0。
"""
if ttl <= 0:
raise ValueError(
f"REDIS_CACHE_TTL 必须为正整数秒,实际 {ttl}。"
"训练场景建议 >= 单次训练时长(如 86400)。"
)
return ttl
class RedisResponseCache:
"""基于 Redis 的 LLM 响应缓存。
使用 content-addressed 策略:key = sha256(model + json(messages))
值为 JSON 序列化的 LLMResponse。
当 Redis 不可用时静默降级:get 返回 None,set 吞异常并记录 warning。
Args:
redis: 异步 Redis 客户端实例(duck-typed,需支持 get/set 方法)。
ttl_s: 缓存过期时间(秒)。None 表示永不过期。
"""
def __init__(self, redis: Any, ttl_s: int | None) -> None:
self._redis = redis
self._ttl_s = ttl_s
def _build_key(
self,
model: str,
messages: list[dict[str, str]],
cache_salt: str | None = None,
) -> str:
"""构造 content-addressed 缓存键。
Args:
model: 模型名称。
messages: 消息列表。
cache_salt: 可选缓存盐(如跨 epoch 强制重采样)。仅当非 None 时才加入
键 payload,保证默认 None 时键结构与旧缓存一字节不差、旧键不失效。
Returns:
sha256 哈希字符串作为 Redis 键。
"""
key_obj: dict[str, Any] = {"model": model, "messages": messages}
if cache_salt is not None:
key_obj["salt"] = cache_salt
payload = json.dumps(key_obj, sort_keys=True, ensure_ascii=False)
digest = hashlib.sha256(payload.encode("utf-8")).hexdigest()
return f"llm_cache:{digest}"
async def get(
self,
model: str,
messages: list[dict[str, str]],
cache_salt: str | None = None,
) -> LLMResponse | None:
"""从缓存读取 LLM 响应。
Args:
model: 模型名称。
messages: 消息列表。
cache_salt: 可选缓存盐,透传到键构造。
Returns:
缓存命中时返回 LLMResponse,未命中或 Redis 异常时返回 None。
"""
try:
key = self._build_key(model, messages, cache_salt)
raw = await self._redis.get(key)
except Exception:
logger.warning("Redis 缓存读取失败,降级为未命中")
return None
if raw is None:
return None
data = json.loads(raw)
return LLMResponse(**data)
async def set(
self,
model: str,
messages: list[dict[str, str]],
response: LLMResponse,
cache_salt: str | None = None,
) -> None:
"""将 LLM 响应写入缓存。
Args:
model: 模型名称。
messages: 消息列表。
response: 待缓存的 LLMResponse。
cache_salt: 可选缓存盐,透传到键构造。
"""
try:
key = self._build_key(model, messages, cache_salt)
value = json.dumps(dataclasses.asdict(response), ensure_ascii=False)
if self._ttl_s:
await self._redis.set(key, value, ex=self._ttl_s)
else:
await self._redis.set(key, value)
except Exception:
logger.warning("Redis 缓存写入失败,跳过缓存")