Spaces:
Paused
Paused
| 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() | |
| 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)) | |
| 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())), | |
| ) | |
| 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 | |
| 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 | |
| def backend_name(self) -> str: | |
| if self.redis_url: | |
| return "tiered" | |
| return "memory" | |
| 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: ... | |
| 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, | |
| } | |
| 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 | |
| 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 | |