File size: 14,364 Bytes
1689f23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
671eec5
1689f23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
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