"""CacheMW: 响应缓存中间件;key 公式在此(算法一份),后端是笨 KV。 ARCH §7.5: key = sha256(canonical_json({model, messages_digest, namespace, salt}));多模态 part 先摘要再 hash(防 Video-Tree 整段 base64 进 hash 的 开销);namespace 必填防跨项目/租户毒化;只缓存阶梯通过的成功响应 (StructuredMW 在内层,能返回即已通过)。读写失败静默降级(铁律)。 """ from __future__ import annotations import asyncio import dataclasses import hashlib import json import uuid from typing import TYPE_CHECKING, Any from loguru import logger from polygateway.types import ChatRequest, LLMResponse if TYPE_CHECKING: from polygateway.ports import CacheBackend, CallNext, StructuredOutputStrategy _KEY_PREFIX = "pgw:cache:" _RESPONSE_FIELDS = {f.name for f in dataclasses.fields(LLMResponse)} def digest_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: """多模态 content part 先各自 sha256 摘要再参与序列化;文本原文参与。 与遥测落库共用同一函数(ARCH §7.8),保证缓存 key 与遥测口径一致。 """ digested = [] for msg in messages: content = msg.get("content") if isinstance(content, list): parts = [_digest_part(part) for part in content] digested.append({**msg, "content": parts}) else: digested.append(msg) return digested def _digest_part(part: Any) -> Any: if isinstance(part, dict) and part.get("type") == "image_url": url = str(part.get("image_url", {}).get("url", "")) return {"type": "image_url", "sha256": hashlib.sha256(url.encode()).hexdigest()} return part def build_cache_key( model_fingerprint: str, messages: list[dict[str, Any]], namespace: str, salt: str | None ) -> str: """缓存 key 公式;salt 仅非 None 时参与(VT 旧键语义: 不传 salt 键形不变)。""" key_obj: dict[str, Any] = { "model": model_fingerprint, "messages": digest_messages(messages), "namespace": namespace, } if salt is not None: key_obj["salt"] = salt payload = json.dumps(key_obj, sort_keys=True, ensure_ascii=False) return _KEY_PREFIX + hashlib.sha256(payload.encode("utf-8")).hexdigest() class CacheMW: """洋葱第二层(遥测内、结构化外)。 model_fingerprint 由装配层从源列表计算(多源 scope = 排序去重的 model 名合集);源集合变化 → key 变化 → 一次性冷启动,换取跨源集合零毒化。 """ def __init__( self, *, backend: CacheBackend, model_fingerprint: str, default_namespace: str, ttl_s: int, strategy: StructuredOutputStrategy | None = None, ) -> None: if not default_namespace.strip(): raise ValueError("缓存 namespace 不能为空(防跨项目/租户毒化)") if ttl_s <= 0: raise ValueError("缓存 TTL 必须 > 0(禁止永不过期)") self._backend = backend self._fingerprint = model_fingerprint self._namespace = default_namespace self._ttl_s = ttl_s self._strategy = strategy async def __call__(self, request: ChatRequest, call_next: CallNext) -> LLMResponse: namespace = request.cache_namespace or self._namespace key = build_cache_key(self._fingerprint, request.messages, namespace, request.cache_salt) cached = await self._safe_get(key) if cached is not None: hit = self._rehydrate(cached, request) if hit is not None: return hit response = await call_next(request) await self._safe_set(key, self._serialize(response)) return response # —— 命中路径 —— def _rehydrate(self, raw: str, request: ChatRequest) -> LLMResponse | None: """反序列化 + 按本次调用的 structured 档零网络重建;失败按未命中。""" try: data = json.loads(raw) fields = {k: v for k, v in data.items() if k in _RESPONSE_FIELDS} structured_data = self._rebuild_structured(fields.get("content", ""), request) fields.update( cache_hit=True, latency_ms=0, ttft_ms=None, max_inter_token_ms=None, call_id=str(uuid.uuid4()), structured_data=structured_data, ) return LLMResponse(**fields) except Exception as exc: logger.warning("缓存命中重建失败,按未命中回源: {}", exc) return None def _rebuild_structured(self, content: str, request: ChatRequest) -> Any | None: """对缓存 content 重跑阶梯②③(schema 变更后旧缓存自动重校验,设计 §2.1)。""" if request.structured is None: return None if self._strategy is None: raise ValueError("structured 调用需要装配 strategy 才能命中重建") parsed = self._strategy.parse(content) if request.structured == "json": return parsed return request.structured.model_validate(parsed) # —— 写路径与降级 —— def _serialize(self, response: LLMResponse) -> str: data = dataclasses.asdict(response) data.pop("structured_data", None) # pydantic 实例不可 JSON 往返(设计 §2.1) return json.dumps(data, ensure_ascii=False) async def _safe_get(self, key: str) -> str | None: try: return await self._backend.get(key) except asyncio.CancelledError: raise except Exception as exc: logger.warning("缓存读取失败,降级为未命中: {}", exc) return None async def _safe_set(self, key: str, value: str) -> None: try: await self._backend.set(key, value, self._ttl_s) except asyncio.CancelledError: raise except Exception as exc: logger.warning("缓存写入失败,跳过缓存: {}", exc)