p5jsai-api / response_cache.py
li2895's picture
Fix upstream delta text parsing
671eec5
Raw
History Blame Contribute Delete
14.4 kB
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