diff --git a/src/core/config.py b/src/core/config.py index fa6ef49..00dc846 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -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" diff --git a/src/llm/predict.py b/src/llm/predict.py index c302e95..c582df4 100644 --- a/src/llm/predict.py +++ b/src/llm/predict.py @@ -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 穿插。 +# ── 缓存后端:内存(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 + 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) + + +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) - entry = _cache.get(key) - if entry is not None: - ts, result = entry - if time.time() - ts < _CACHE_TTL_SEC: - return result - _cache.pop(key, None) - return None + if isinstance(_cache_backend, _RedisCache): + return await _cache_backend.get_async(key) + return _cache_backend.get(key) -def _set_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str, result: PredictResult) -> None: - # P1-5: 无锁写入。同上,dict set 原子。 +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) - _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) + 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",