from __future__ import annotations import asyncio import base64 import hashlib import json import os import time from collections import OrderedDict from collections.abc import Mapping from dataclasses import asdict, dataclass from typing import Any, Protocol import zlib try: import redis.asyncio as redis_async except ImportError: # pragma: no cover - optional dependency redis_async = None def _env_bool(name: str, default: bool) -> bool: raw = os.getenv(name) if raw is None: return default return raw.strip().lower() not in {"0", "false", "no", "off"} def _env_int(name: str, default: int) -> int: raw = os.getenv(name) if raw is None: return default try: return int(raw) except ValueError: return default def _stable_json(data: Any) -> str: return json.dumps(data, ensure_ascii=False, sort_keys=True, separators=(",", ":")) def _sha256(text: str) -> str: return hashlib.sha256(text.encode("utf-8")).hexdigest() @dataclass(frozen=True) class CompletionArtifact: schema_version: int raw_text: str model_id: str usage_input_tokens: int usage_output_tokens: int stop_reason: str stored_at: float def to_json(self) -> str: return _stable_json(asdict(self)) @classmethod def from_json(cls, raw: str) -> "CompletionArtifact": data = json.loads(raw) return cls( schema_version=int(data.get("schema_version", 1)), raw_text=str(data.get("raw_text", "")), model_id=str(data.get("model_id", "")), usage_input_tokens=int(data.get("usage_input_tokens", 0)), usage_output_tokens=int(data.get("usage_output_tokens", 0)), stop_reason=str(data.get("stop_reason", "end_turn")), stored_at=float(data.get("stored_at", time.time())), ) @dataclass(frozen=True) class CacheConfig: enabled: bool ttl_secs: int tool_ttl_secs: int max_entry_bytes: int memory_max_items: int redis_url: str | None redis_prefix: str @classmethod def from_env(cls) -> "CacheConfig": configured_max = _env_int("RESPONSE_CACHE_MAX_ENTRY_BYTES", 32 * 1024 * 1024) return cls( enabled=_env_bool("RESPONSE_CACHE_ENABLED", True), ttl_secs=max(1, _env_int("RESPONSE_CACHE_TTL_SECS", 300)), tool_ttl_secs=max(1, _env_int("RESPONSE_CACHE_TOOL_TTL_SECS", 120)), max_entry_bytes=0 if configured_max <= 0 else max(1024, configured_max), memory_max_items=max(1, _env_int("RESPONSE_CACHE_MEMORY_MAX_ITEMS", 256)), redis_url=os.getenv("RESPONSE_CACHE_REDIS_URL") or None, redis_prefix=os.getenv("RESPONSE_CACHE_REDIS_PREFIX", "p5js2api:resp:v1"), ) def ttl_for(self, has_tools: bool) -> int: return self.tool_ttl_secs if has_tools else self.ttl_secs @property def backend_name(self) -> str: if self.redis_url: return "tiered" return "memory" @dataclass(frozen=True) class CacheLookupResult: artifact: CompletionArtifact | None source: str | None ttl_secs: int | None = None class CacheBackend(Protocol): async def get(self, key: str) -> CompletionArtifact | None: ... async def set(self, key: str, artifact: CompletionArtifact, ttl_secs: int) -> None: ... async def delete(self, key: str) -> None: ... @dataclass class _MemoryEntry: artifact: CompletionArtifact expires_at: float size_bytes: int class InMemoryCacheBackend: def __init__(self, max_items: int): self._max_items = max_items self._entries: OrderedDict[str, _MemoryEntry] = OrderedDict() self._lock = asyncio.Lock() async def get(self, key: str) -> tuple[CompletionArtifact | None, int | None]: async with self._lock: self._prune_expired_locked() entry = self._entries.get(key) if not entry: return None, None if entry.expires_at <= time.time(): self._entries.pop(key, None) return None, None self._entries.move_to_end(key) ttl_secs = max(1, int(entry.expires_at - time.time())) return entry.artifact, ttl_secs async def set(self, key: str, artifact: CompletionArtifact, ttl_secs: int) -> None: raw = artifact.to_json() entry = _MemoryEntry( artifact=artifact, expires_at=time.time() + ttl_secs, size_bytes=len(raw.encode("utf-8")), ) async with self._lock: self._prune_expired_locked() self._entries[key] = entry self._entries.move_to_end(key) while len(self._entries) > self._max_items: self._entries.popitem(last=False) async def delete(self, key: str) -> None: async with self._lock: self._entries.pop(key, None) def _prune_expired_locked(self) -> None: now = time.time() expired = [key for key, entry in self._entries.items() if entry.expires_at <= now] for key in expired: self._entries.pop(key, None) class RedisCacheBackend: def __init__(self, redis_url: str): if redis_async is None: raise RuntimeError("redis package is not installed") self._client = redis_async.from_url(redis_url) async def get(self, key: str) -> tuple[CompletionArtifact | None, int | None]: pipeline = self._client.pipeline() pipeline.get(key) pipeline.ttl(key) raw, ttl = await pipeline.execute() if raw is None: return None, None if isinstance(raw, bytes): raw = raw.decode("utf-8") payload = json.loads(raw) encoding = payload.get("encoding", "plain") data = payload.get("data", "") if encoding == "zlib+base64": decoded = base64.b64decode(data.encode("ascii")) data = zlib.decompress(decoded).decode("utf-8") return CompletionArtifact.from_json(data), max(1, int(ttl)) if isinstance(ttl, int) and ttl > 0 else None async def set(self, key: str, artifact: CompletionArtifact, ttl_secs: int) -> None: raw_json = artifact.to_json().encode("utf-8") compressed = zlib.compress(raw_json, level=6) payload = { "encoding": "zlib+base64", "data": base64.b64encode(compressed).decode("ascii"), } await self._client.set(key, _stable_json(payload), ex=ttl_secs) async def delete(self, key: str) -> None: await self._client.delete(key) class TieredCacheBackend: def __init__(self, memory: InMemoryCacheBackend, redis_backend: RedisCacheBackend | None): self._memory = memory self._redis = redis_backend async def get(self, key: str) -> CacheLookupResult: artifact, ttl_secs = await self._memory.get(key) if artifact is not None: return CacheLookupResult(artifact=artifact, source="memory", ttl_secs=ttl_secs) if self._redis is None: return CacheLookupResult(artifact=None, source=None) artifact, ttl_secs = await self._redis.get(key) if artifact is None: return CacheLookupResult(artifact=None, source=None) await self._memory.set(key, artifact, ttl_secs=ttl_secs or 60) return CacheLookupResult(artifact=artifact, source="redis", ttl_secs=ttl_secs) async def set(self, key: str, artifact: CompletionArtifact, ttl_secs: int) -> None: await self._memory.set(key, artifact, ttl_secs) if self._redis is not None: await self._redis.set(key, artifact, ttl_secs) async def delete(self, key: str) -> None: await self._memory.delete(key) if self._redis is not None: await self._redis.delete(key) class InFlightRegistry: def __init__(self): self._lock = asyncio.Lock() self._futures: dict[str, asyncio.Future[CompletionArtifact]] = {} async def start(self, key: str) -> tuple[bool, asyncio.Future[CompletionArtifact]]: async with self._lock: current = self._futures.get(key) if current is not None: return False, current future: asyncio.Future[CompletionArtifact] = asyncio.get_running_loop().create_future() self._futures[key] = future return True, future async def resolve(self, key: str, artifact: CompletionArtifact) -> None: async with self._lock: future = self._futures.pop(key, None) if future is not None and not future.done(): future.set_result(artifact) async def reject(self, key: str, exc: BaseException) -> None: async with self._lock: future = self._futures.pop(key, None) if future is not None and not future.done(): future.set_exception(exc) future.add_done_callback(lambda f: f.exception()) class ResponseCacheService: def __init__(self, config: CacheConfig): self.config = config self.inflight = InFlightRegistry() self._memory = InMemoryCacheBackend(config.memory_max_items) redis_backend: RedisCacheBackend | None = None self._redis_init_error: str | None = None if config.redis_url and redis_async is not None: try: redis_backend = RedisCacheBackend(config.redis_url) except Exception as exc: redis_backend = None self._redis_init_error = str(exc) elif config.redis_url and redis_async is None: self._redis_init_error = "redis package unavailable" self._backend = TieredCacheBackend(self._memory, redis_backend) self._redis_backend = redis_backend self._stats_lock = asyncio.Lock() self._stats = { "hits_memory": 0, "hits_redis": 0, "misses": 0, "stores": 0, "bypasses": 0, "store_errors": 0, "oversize_skips": 0, "inflight_waits": 0, } async def get(self, key: str) -> CacheLookupResult: if not self.config.enabled: await self._record("bypasses") return CacheLookupResult(artifact=None, source=None) lookup = await self._backend.get(key) if lookup.artifact is None: await self._record("misses") elif lookup.source == "memory": await self._record("hits_memory") elif lookup.source == "redis": await self._record("hits_redis") return lookup async def set(self, key: str, artifact: CompletionArtifact, ttl_secs: int) -> bool: if not self.config.enabled: return False raw_bytes = len(artifact.to_json().encode("utf-8")) if self.config.max_entry_bytes > 0 and raw_bytes > self.config.max_entry_bytes: await self._record("oversize_skips") return False try: await self._backend.set(key, artifact, ttl_secs) except Exception: await self._record("store_errors") return False await self._record("stores") return True async def delete(self, key: str) -> None: await self._backend.delete(key) async def record_bypass(self) -> None: await self._record("bypasses") async def record_inflight_wait(self) -> None: await self._record("inflight_waits") def build_key( self, *, protocol_family: str, resolved_model: str, upstream_messages: list[dict], auth_scope: str, has_tools: bool, ) -> str: payload = { "schema_version": 1, "protocol_family": protocol_family, "resolved_model": resolved_model, "upstream_messages": upstream_messages, "auth_scope": auth_scope, "has_tools": has_tools, } return f"{self.config.redis_prefix}:{_sha256(_stable_json(payload))}" def auth_scope_from_headers(self, headers: Mapping[str, str]) -> str: token = self._header_value(headers, "x-api-key") or self._header_value(headers, "authorization") if not token: return "anon" token = token.removeprefix("Bearer ").strip() if not token: return "anon" return f"auth-{_sha256(token)[:16]}" def should_bypass(self, headers: Mapping[str, str]) -> bool: explicit = self._header_value(headers, "x-proxy-cache") if explicit and explicit.lower() in {"bypass", "off", "false", "no-cache"}: return True cache_control = self._header_value(headers, "cache-control") if cache_control and any(flag in cache_control.lower() for flag in ("no-cache", "no-store")): return True return False def describe(self) -> dict[str, Any]: return { "enabled": self.config.enabled, "backend": self.backend_name, "redis_configured": bool(self.config.redis_url), "redis_available": self._redis_backend is not None, "redis_error": self._redis_init_error, "ttl_secs": self.config.ttl_secs, "tool_ttl_secs": self.config.tool_ttl_secs, "max_entry_bytes": self.config.max_entry_bytes, "memory_max_items": self.config.memory_max_items, **self._stats, } @property def backend_name(self) -> str: if self._redis_backend is not None: return "tiered" return "memory" async def _record(self, key: str) -> None: async with self._stats_lock: self._stats[key] = self._stats.get(key, 0) + 1 @staticmethod def _header_value(headers: Mapping[str, str], name: str) -> str | None: if name in headers: return headers[name] lower_name = name.lower() for key, value in headers.items(): if key.lower() == lower_name: return value return None _CACHE_SERVICE: ResponseCacheService | None = None def get_cache_service() -> ResponseCacheService: global _CACHE_SERVICE if _CACHE_SERVICE is None: _CACHE_SERVICE = ResponseCacheService(CacheConfig.from_env()) return _CACHE_SERVICE