feat: 预测缓存可选 Redis 后端(PREDICT_CACHE_URL)
PREDICT_CACHE_URL 为空=进程内 LRU+TTL dict(默认,行为不变);
填 redis:// 启用 Redis,失败自动降级内存并 warning,不中断预测。
- _CacheBackend 接口 + _MemoryCache/_RedisCache 两个后端
- 同键格式(predict:{match}:{provider}:{model}:{version}:{tpl_hash[:12}])
- 同 TTL(300s);Redis 用 pickle 序列化
- redis 包未安装/连接失败 → 降级内存;不强制依赖 redis 启动
- clear_prompt_cache 同步清空内存预测缓存
全量测试 270 通过;缓存后端单测 4/4。
This commit is contained in:
@@ -32,6 +32,12 @@ class Settings(BaseSettings):
|
||||
LLM_SPECIALIST_MODEL: str = ""
|
||||
LLM_AGGREGATOR_MODEL: str = ""
|
||||
|
||||
# ── 预测缓存 ──
|
||||
# 预测响应缓存后端:空(默认)=进程内 LRU+TTL 字典;填 redis://host:port/db 启用 Redis。
|
||||
# Redis 失败自动降级内存缓存并 warning,不中断预测;不强制依赖 redis 包。
|
||||
# TTL 固定 300s(5 分钟),键格式与内存后端一致(含 prompt 模板 hash)。
|
||||
PREDICT_CACHE_URL: str = ""
|
||||
|
||||
# --- data sources ---
|
||||
BZZOIRO_KEY: str = ""
|
||||
BZZOIRO_BASE: str = "https://sports.bzzoiro.com/api/v2"
|
||||
|
||||
+120
-19
@@ -25,9 +25,7 @@ _PROMPT_DIR = Path(__file__).resolve().parent / "prompts"
|
||||
# ── LLM 响应缓存(match+provider+model+version → 结果) ──
|
||||
_CACHE_TTL_SEC = 300 # 5 分钟
|
||||
_CACHE_MAX_SIZE = 200 # P3-1: 有上限,避免长期运行内存无限增长
|
||||
# P1-5: 缓存仅在 asyncio 协程内同步访问(dict 操作 GIL 原子),无需 threading.Lock。
|
||||
# 删除 _cache_lock,避免同步锁阻塞事件循环;dict 的 get/set 在 CPython 下原子。
|
||||
_cache: dict[str, tuple[float, PredictResult]] = {}
|
||||
_CACHE_PREFIX = "predict:" # Redis key 前缀
|
||||
|
||||
|
||||
def _cache_key(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> str:
|
||||
@@ -36,30 +34,130 @@ def _cache_key(match_id: int, provider: str, model: str, version: str, tpl_hash:
|
||||
仅用 version 做键不够 —— 编辑器里改动 `match_prediction_v1.md` 而版本号
|
||||
不变时,进程内缓存仍会返回旧模板产生的旧结果(见审查报告 P2-6)。
|
||||
把模板内容 hash 纳入键,模板一改缓存自动失效。
|
||||
内存与 Redis 共用同一键格式,TTL 一致。
|
||||
"""
|
||||
return f"{match_id}:{provider}:{model}:{version}:{tpl_hash[:12]}"
|
||||
return f"{_CACHE_PREFIX}{match_id}:{provider}:{model}:{version}:{tpl_hash[:12]}"
|
||||
|
||||
|
||||
def _get_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> PredictResult | None:
|
||||
# P1-5: 无锁访问。dict get/del 在 CPython GIL 下原子,且无 await 穿插。
|
||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||
entry = _cache.get(key)
|
||||
# ── 缓存后端:内存(LRU+TTL),可选 Redis ──────────────────────────
|
||||
class _CacheBackend:
|
||||
"""缓存后端统一接口:_get 同步返回(命中时),_set 异步(Redis 为 async,内存同步)。"""
|
||||
|
||||
def _raw_key(self, key: str) -> str:
|
||||
return key
|
||||
|
||||
def get(self, key: str) -> PredictResult | None:
|
||||
raise NotImplementedError
|
||||
|
||||
async def set(self, key: str, result: PredictResult, ttl: int) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _MemoryCache(_CacheBackend):
|
||||
"""进程内 LRU+TTL 缓存(默认后端)。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._store: dict[str, tuple[float, PredictResult]] = {}
|
||||
|
||||
def get(self, key: str) -> PredictResult | None:
|
||||
entry = self._store.get(key)
|
||||
if entry is not None:
|
||||
ts, result = entry
|
||||
if time.time() - ts < _CACHE_TTL_SEC:
|
||||
return result
|
||||
_cache.pop(key, None)
|
||||
self._store.pop(key, None)
|
||||
return None
|
||||
|
||||
async def set(self, key: str, result: PredictResult, ttl: int) -> None:
|
||||
self._store[key] = (time.time(), result)
|
||||
if len(self._store) > _CACHE_MAX_SIZE:
|
||||
oldest_key = min(self._store, key=lambda k: self._store[k][0])
|
||||
self._store.pop(oldest_key, None)
|
||||
|
||||
def _set_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str, result: PredictResult) -> None:
|
||||
# P1-5: 无锁写入。同上,dict set 原子。
|
||||
|
||||
class _RedisCache(_CacheBackend):
|
||||
"""可选 Redis 后端: PREDICT_CACHE_URL 非空时启用。
|
||||
|
||||
失败降级:读失败返回 None(跳过缓存),写失败打 warning;不中断预测主流程。
|
||||
不强制依赖 redis 包——未安装时启动回退内存并 warning。
|
||||
"""
|
||||
|
||||
def __init__(self, url: str) -> None:
|
||||
self._url = url
|
||||
self._redis = None # type: ignore[var-annotated]
|
||||
self._memory_fallback = _MemoryCache()
|
||||
self._available: bool | None = None # None=未探测,True=可用,False=不可用
|
||||
|
||||
async def _ensure_conn(self) -> bool:
|
||||
"""懒初始化 Redis 连接;失败返回 False 并降级内存。"""
|
||||
if self._available is not None:
|
||||
return self._available
|
||||
try:
|
||||
from redis.asyncio import Redis
|
||||
|
||||
self._redis = Redis.from_url(self._url, decode_responses=True, socket_timeout=2.0)
|
||||
await self._redis.ping()
|
||||
self._available = True
|
||||
logger.info("predict cache: Redis 后端已连接 %s", self._url.replace(self._url.split("@")[-1] if "@" in self._url else self._url, "***") if "://" in self._url else "redis")
|
||||
except Exception as e:
|
||||
self._available = False
|
||||
logger.warning("predict cache: Redis 连接失败(%s),降级内存缓存", e)
|
||||
return self._available
|
||||
|
||||
def get(self, key: str) -> PredictResult | None:
|
||||
# Redis get 是 async 的,此处统一由调用方走 async 路径;
|
||||
# 同步 get 仅用于不可降级场景——Redis 模式下直接返回 None,
|
||||
# 实际读取通过 get_async 完成。
|
||||
return None
|
||||
|
||||
async def get_async(self, key: str) -> PredictResult | None:
|
||||
if not await self._ensure_conn():
|
||||
return self._memory_fallback.get(key)
|
||||
try:
|
||||
import pickle
|
||||
|
||||
raw = await self._redis.get(key) # type: ignore[union-attr]
|
||||
if raw is None:
|
||||
return None
|
||||
return pickle.loads(raw.encode("latin-1")) if isinstance(raw, str) else pickle.loads(raw)
|
||||
except Exception as e:
|
||||
logger.warning("predict cache: Redis GET 失败(%s),跳过缓存", e)
|
||||
return None
|
||||
|
||||
async def set(self, key: str, result: PredictResult, ttl: int) -> None:
|
||||
if not await self._ensure_conn():
|
||||
await self._memory_fallback.set(key, result, ttl)
|
||||
return
|
||||
try:
|
||||
import pickle
|
||||
|
||||
payload = pickle.dumps(result).decode("latin-1")
|
||||
await self._redis.set(key, payload, ex=ttl) # type: ignore[union-attr]
|
||||
except Exception as e:
|
||||
logger.warning("predict cache: Redis SET 失败(%s),降级内存写入", e)
|
||||
await self._memory_fallback.set(key, result, ttl)
|
||||
|
||||
|
||||
def _build_cache_backend() -> _CacheBackend:
|
||||
url = getattr(settings, "PREDICT_CACHE_URL", None)
|
||||
if url:
|
||||
return _RedisCache(url)
|
||||
return _MemoryCache()
|
||||
|
||||
|
||||
_cache_backend: _CacheBackend = _build_cache_backend()
|
||||
|
||||
|
||||
async def _get_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> PredictResult | None:
|
||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||
_cache[key] = (time.time(), result)
|
||||
# P3-1: 超过上限时淘汰最旧条目(按时间戳排序)
|
||||
if len(_cache) > _CACHE_MAX_SIZE:
|
||||
oldest_key = min(_cache, key=lambda k: _cache[k][0])
|
||||
_cache.pop(oldest_key, None)
|
||||
if isinstance(_cache_backend, _RedisCache):
|
||||
return await _cache_backend.get_async(key)
|
||||
return _cache_backend.get(key)
|
||||
|
||||
|
||||
async def _set_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str, result: PredictResult) -> None:
|
||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||
await _cache_backend.set(key, result, _CACHE_TTL_SEC)
|
||||
|
||||
|
||||
def clear_prompt_cache() -> None:
|
||||
@@ -67,9 +165,12 @@ def clear_prompt_cache() -> None:
|
||||
|
||||
lru_cache 的模板缓存是进程级的,改完 .md 需要重启进程才能生效;
|
||||
提供显式清理入口,避免"改了模板却看不到变化"的困惑(见审查报告 P2-5)。
|
||||
同时清空预测响应缓存(内存后端);Redis 后端因共享不清除。
|
||||
"""
|
||||
_load_prompt_template.cache_clear()
|
||||
logger.info("prompt 模板缓存已清空")
|
||||
if isinstance(_cache_backend, _MemoryCache):
|
||||
_cache_backend._store.clear()
|
||||
logger.info("prompt 模板缓存 + 预测响应缓存(内存)已清空")
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
@@ -235,7 +336,7 @@ async def _predict_single(
|
||||
|
||||
# 0. 查缓存(同 match+provider+model+version+模板hash 5 分钟内直接返)
|
||||
if use_cache:
|
||||
cached = _get_cached(match_id, settings.LLM_PROVIDER, provider.model, version, tpl_hash)
|
||||
cached = await _get_cached(match_id, settings.LLM_PROVIDER, provider.model, version, tpl_hash)
|
||||
if cached is not None:
|
||||
logger.debug("predict cache hit match=%s", match_id)
|
||||
return cached
|
||||
@@ -334,7 +435,7 @@ async def _predict_single(
|
||||
|
||||
# 5. 写入缓存(仅当允许缓存时)
|
||||
if use_cache:
|
||||
_set_cached(match_id, settings.LLM_PROVIDER, provider.model, version, tpl_hash, result)
|
||||
await _set_cached(match_id, settings.LLM_PROVIDER, provider.model, version, tpl_hash, result)
|
||||
logger.info(
|
||||
"预测完成 match=%s mode=%s status=%s pred=%s:%s (%s) latency=%sms",
|
||||
match_id, "single", "success",
|
||||
|
||||
Reference in New Issue
Block a user