cachellm-proxy 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (43) hide show
  1. cachellm/__init__.py +13 -0
  2. cachellm/__main__.py +6 -0
  3. cachellm/api/__init__.py +5 -0
  4. cachellm/api/app.py +143 -0
  5. cachellm/api/auth.py +34 -0
  6. cachellm/api/deps.py +106 -0
  7. cachellm/api/routes_admin.py +265 -0
  8. cachellm/api/routes_chat.py +385 -0
  9. cachellm/api/sse.py +98 -0
  10. cachellm/cache/__init__.py +3 -0
  11. cachellm/cache/analytics.py +150 -0
  12. cachellm/cache/coalesce.py +63 -0
  13. cachellm/cache/entry.py +92 -0
  14. cachellm/cache/exact_store.py +33 -0
  15. cachellm/cache/keys.py +124 -0
  16. cachellm/cache/policy.py +134 -0
  17. cachellm/cache/redis_client.py +22 -0
  18. cachellm/cache/service.py +332 -0
  19. cachellm/cache/vector_store.py +217 -0
  20. cachellm/cli.py +122 -0
  21. cachellm/embeddings/__init__.py +19 -0
  22. cachellm/embeddings/base.py +38 -0
  23. cachellm/embeddings/fastembed_backend.py +75 -0
  24. cachellm/embeddings/hash_backend.py +42 -0
  25. cachellm/errors.py +72 -0
  26. cachellm/logging_setup.py +56 -0
  27. cachellm/models.py +181 -0
  28. cachellm/observability/__init__.py +6 -0
  29. cachellm/observability/metrics.py +147 -0
  30. cachellm/observability/tracing.py +107 -0
  31. cachellm/pricing.py +108 -0
  32. cachellm/providers/__init__.py +7 -0
  33. cachellm/providers/base.py +84 -0
  34. cachellm/providers/bedrock.py +238 -0
  35. cachellm/providers/fake.py +56 -0
  36. cachellm/providers/openai_compat.py +131 -0
  37. cachellm/providers/registry.py +96 -0
  38. cachellm/py.typed +0 -0
  39. cachellm/settings.py +230 -0
  40. cachellm_proxy-0.1.0.dist-info/METADATA +550 -0
  41. cachellm_proxy-0.1.0.dist-info/RECORD +43 -0
  42. cachellm_proxy-0.1.0.dist-info/WHEEL +4 -0
  43. cachellm_proxy-0.1.0.dist-info/entry_points.txt +3 -0
@@ -0,0 +1,332 @@
1
+ """The cache orchestrator: one lookup path, one store path, one invalidate path.
2
+
3
+ Everything the proxy knows about caching lives behind this class, so the HTTP
4
+ layer stays thin and the same logic can be reused by an in-process client
5
+ wrapper without a server.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import time
11
+ from dataclasses import dataclass, field
12
+ from typing import Any, Literal
13
+
14
+ import numpy as np
15
+ import structlog
16
+
17
+ from cachellm.cache.analytics import Analytics, NearMiss
18
+ from cachellm.cache.entry import CacheEntry
19
+ from cachellm.cache.exact_store import ExactStore
20
+ from cachellm.cache.keys import (
21
+ embedding_text,
22
+ exact_hash,
23
+ namespace_for,
24
+ prompt_fingerprint,
25
+ )
26
+ from cachellm.cache.keys import (
27
+ entry_id as make_entry_id,
28
+ )
29
+ from cachellm.cache.policy import PolicyDecision, decide
30
+ from cachellm.cache.vector_store import VectorStore
31
+ from cachellm.embeddings.base import Embedder
32
+ from cachellm.models import ChatCompletionRequest
33
+ from cachellm.pricing import estimate_cost
34
+ from cachellm.settings import Settings
35
+
36
+ log = structlog.get_logger(__name__)
37
+
38
+ LookupStatus = Literal["hit", "miss", "bypass", "shadow_hit"]
39
+
40
+
41
+ @dataclass
42
+ class LookupResult:
43
+ status: LookupStatus
44
+ decision: PolicyDecision
45
+ namespace: str = ""
46
+ similarity: float = 0.0
47
+ tier: str = "" # "exact" | "semantic"
48
+ entry: CacheEntry | None = None
49
+ embedding: np.ndarray | None = None
50
+ exact: str = ""
51
+ cache_text: str = ""
52
+ lookup_ms: float = 0.0
53
+ neighbours: list[tuple[str, float]] = field(default_factory=list)
54
+
55
+ @property
56
+ def served_from_cache(self) -> bool:
57
+ return self.status == "hit" and self.entry is not None
58
+
59
+ @property
60
+ def category(self) -> str:
61
+ return self.decision.category
62
+
63
+ def saved_usd(self, model: str) -> float:
64
+ if self.entry is None:
65
+ return 0.0
66
+ return estimate_cost(model, self.entry.prompt_tokens, self.entry.completion_tokens)
67
+
68
+
69
+ class CacheService:
70
+ def __init__(
71
+ self,
72
+ *,
73
+ settings: Settings,
74
+ embedder: Embedder,
75
+ vectors: VectorStore,
76
+ exact: ExactStore,
77
+ analytics: Analytics,
78
+ ) -> None:
79
+ self.settings = settings
80
+ self.embedder = embedder
81
+ self.vectors = vectors
82
+ self.exact = exact
83
+ self.analytics = analytics
84
+
85
+ # ------------------------------------------------------------------ lookup
86
+ async def lookup(
87
+ self,
88
+ request: ChatCompletionRequest,
89
+ provider_name: str,
90
+ *,
91
+ cache_control: str = "",
92
+ ) -> LookupResult:
93
+ started = time.perf_counter()
94
+ decision = decide(request, self.settings, cache_control=cache_control)
95
+ if not decision.cacheable:
96
+ return LookupResult(
97
+ status="bypass", decision=decision, lookup_ms=(time.perf_counter() - started) * 1000
98
+ )
99
+
100
+ namespace = namespace_for(request, provider_name)
101
+ raw_text = request.cache_text(include_history=self.settings.cache_multi_turn)
102
+ exact_key = exact_hash(raw_text)
103
+
104
+ # --- level 1: exact match, no embedding needed -----------------------
105
+ entry_id = await self.exact.get(namespace, exact_key)
106
+ if entry_id:
107
+ entry = await self.vectors.get(entry_id)
108
+ if entry is not None:
109
+ status: LookupStatus = "shadow_hit" if self.settings.shadow_mode else "hit"
110
+ return LookupResult(
111
+ status=status,
112
+ decision=decision,
113
+ namespace=namespace,
114
+ similarity=1.0,
115
+ tier="exact",
116
+ entry=entry,
117
+ exact=exact_key,
118
+ cache_text=raw_text,
119
+ lookup_ms=(time.perf_counter() - started) * 1000,
120
+ )
121
+ # Pointer outlived its entry (entry TTL shorter, or manual delete).
122
+ await self.exact.delete(namespace, exact_key)
123
+
124
+ # --- level 2: semantic nearest neighbour -----------------------------
125
+ text_to_embed = embedding_text(raw_text, strip_fillers=self.settings.strip_filler_words)
126
+ vector = await self.embedder.embed(text_to_embed)
127
+ matches = await self.vectors.search(namespace, vector, self.settings.top_k)
128
+ neighbours = [(m.entry_id, round(sim, 4)) for m, sim in matches]
129
+
130
+ if matches:
131
+ best_entry, best_sim = matches[0]
132
+ if best_sim >= decision.threshold:
133
+ status = "shadow_hit" if self.settings.shadow_mode else "hit"
134
+ return LookupResult(
135
+ status=status,
136
+ decision=decision,
137
+ namespace=namespace,
138
+ similarity=best_sim,
139
+ tier="semantic",
140
+ entry=best_entry,
141
+ embedding=vector,
142
+ exact=exact_key,
143
+ cache_text=raw_text,
144
+ neighbours=neighbours,
145
+ lookup_ms=(time.perf_counter() - started) * 1000,
146
+ )
147
+ if best_sim >= decision.threshold - self.settings.near_miss_margin:
148
+ await self.analytics.record_near_miss(
149
+ NearMiss(
150
+ prompt=raw_text,
151
+ matched_prompt=best_entry.prompt,
152
+ similarity=round(best_sim, 4),
153
+ threshold=decision.threshold,
154
+ category=decision.category,
155
+ namespace=namespace,
156
+ model=request.model,
157
+ at=time.time(),
158
+ )
159
+ )
160
+ await self.analytics.incr("near_misses")
161
+
162
+ return LookupResult(
163
+ status="miss",
164
+ decision=decision,
165
+ namespace=namespace,
166
+ similarity=matches[0][1] if matches else 0.0,
167
+ embedding=vector,
168
+ exact=exact_key,
169
+ cache_text=raw_text,
170
+ neighbours=neighbours,
171
+ lookup_ms=(time.perf_counter() - started) * 1000,
172
+ )
173
+
174
+ # ------------------------------------------------------------------- store
175
+ async def store(
176
+ self,
177
+ *,
178
+ lookup: LookupResult,
179
+ request: ChatCompletionRequest,
180
+ provider_name: str,
181
+ response_text: str,
182
+ prompt_tokens: int,
183
+ completion_tokens: int,
184
+ finish_reason: str,
185
+ tool_calls: list[dict[str, Any]] | None = None,
186
+ ) -> str | None:
187
+ """Persist a completed answer. Only clean, finished responses are stored."""
188
+ if not lookup.decision.cacheable or lookup.status == "bypass":
189
+ return None
190
+
191
+ complete = finish_reason in ("stop", "end_turn", "eos", "stop_sequence", "complete")
192
+ if not complete and not self.settings.cache_truncated:
193
+ # The most common cause of a mysteriously low hit rate: max_tokens is
194
+ # set below what the model wants to say, every answer comes back
195
+ # truncated, and nothing is cacheable. Counted so /admin/stats can
196
+ # say so out loud instead of leaving it to be guessed at.
197
+ await self.analytics.incr("stores_skipped_truncated")
198
+ log.debug("skip_store_incomplete", finish_reason=finish_reason)
199
+ return None
200
+ if not response_text.strip():
201
+ await self.analytics.incr("stores_skipped_empty")
202
+ return None
203
+
204
+ vector = lookup.embedding
205
+ if vector is None:
206
+ text_to_embed = embedding_text(
207
+ lookup.cache_text
208
+ or request.cache_text(include_history=self.settings.cache_multi_turn),
209
+ strip_fillers=self.settings.strip_filler_words,
210
+ )
211
+ vector = await self.embedder.embed(text_to_embed)
212
+
213
+ namespace = lookup.namespace or namespace_for(request, provider_name)
214
+ exact_key = lookup.exact or exact_hash(lookup.cache_text)
215
+ eid = make_entry_id(namespace, exact_key)
216
+ entry = CacheEntry(
217
+ entry_id=eid,
218
+ namespace=namespace,
219
+ category=lookup.decision.category,
220
+ model=request.model,
221
+ provider=provider_name,
222
+ prompt=lookup.cache_text or request.last_user_text(),
223
+ response_text=response_text,
224
+ prompt_tokens=prompt_tokens,
225
+ completion_tokens=completion_tokens,
226
+ finish_reason=finish_reason,
227
+ ttl=lookup.decision.ttl,
228
+ exact_hash=exact_key,
229
+ embedding=vector,
230
+ tool_calls=tool_calls,
231
+ )
232
+ await self.vectors.put(entry)
233
+ await self.exact.put(namespace, exact_key, eid, lookup.decision.ttl)
234
+ await self.analytics.incr("entries_written")
235
+ log.debug(
236
+ "cache_store",
237
+ entry_id=eid,
238
+ category=entry.category,
239
+ fingerprint=prompt_fingerprint(entry.prompt),
240
+ ttl=entry.ttl,
241
+ )
242
+ return eid
243
+
244
+ # -------------------------------------------------------------- maintenance
245
+ async def register_hit(self, lookup: LookupResult, model: str) -> float:
246
+ """Bump hit bookkeeping and return the modelled dollars saved."""
247
+ if lookup.entry is None:
248
+ return 0.0
249
+ await self.vectors.touch_hit(lookup.entry.entry_id)
250
+ saved = lookup.saved_usd(model)
251
+ await self.analytics.bulk(
252
+ {
253
+ "hits": 1,
254
+ f"hits_{lookup.tier}": 1,
255
+ "tokens_saved_prompt": lookup.entry.prompt_tokens,
256
+ "tokens_saved_completion": lookup.entry.completion_tokens,
257
+ },
258
+ {"usd_saved": saved},
259
+ )
260
+ return saved
261
+
262
+ async def invalidate(
263
+ self, *, namespace: str | None = None, model: str | None = None, drop_all: bool = False
264
+ ) -> int:
265
+ removed = await self.vectors.invalidate(namespace=namespace, model=model, drop_all=drop_all)
266
+ await self.analytics.incr("invalidations")
267
+ log.info(
268
+ "cache_invalidate", namespace=namespace, model=model, drop_all=drop_all, removed=removed
269
+ )
270
+ return removed
271
+
272
+ async def stats(self) -> dict[str, Any]:
273
+ counters = await self.analytics.counters()
274
+ requests = counters.get("requests", 0.0)
275
+ hits = counters.get("hits", 0.0)
276
+ cacheable = counters.get("cacheable_requests", 0.0)
277
+ misses = counters.get("misses", 0.0)
278
+ skipped_truncated = counters.get("stores_skipped_truncated", 0.0)
279
+ diagnostics: list[str] = []
280
+ if misses and skipped_truncated / misses > 0.2:
281
+ diagnostics.append(
282
+ f"{skipped_truncated / misses:.0%} of misses were not cached because the "
283
+ f"response hit max_tokens. Raise max_tokens so answers finish, or set "
284
+ f"CACHELLM_CACHE_TRUNCATED=true to cache truncated answers anyway."
285
+ )
286
+ if requests and counters.get("bypass", 0.0) / requests > 0.5:
287
+ diagnostics.append(
288
+ "Over half of requests bypassed the cache. Check X-Cache-Bypass-Reason: "
289
+ "temperature, tool calls, JSON mode and multi-turn are skipped by default."
290
+ )
291
+ if self.settings.embedding_backend == "hash":
292
+ diagnostics.append(
293
+ "Running the hash embedding backend. It is a deterministic test double "
294
+ "with no understanding of meaning, so the semantic tier will rarely fire. "
295
+ "Set CACHELLM_EMBEDDING_BACKEND=fastembed for real matching."
296
+ )
297
+ elif not self.settings.is_calibrated:
298
+ diagnostics.append(
299
+ f"No measured threshold for embedding model "
300
+ f"{self.settings.embedding_model!r}; using {self.settings.calibrated_threshold}. "
301
+ f"Run `python -m bench.compare_models` on your own pairs."
302
+ )
303
+
304
+ return {
305
+ "entries": await self.vectors.count(),
306
+ "requests": int(requests),
307
+ "cacheable_requests": int(cacheable),
308
+ "hits": int(hits),
309
+ "hits_exact": int(counters.get("hits_exact", 0.0)),
310
+ "hits_semantic": int(counters.get("hits_semantic", 0.0)),
311
+ "misses": int(counters.get("misses", 0.0)),
312
+ "bypassed": int(counters.get("bypass", 0.0)),
313
+ "shadow_hits": int(counters.get("shadow_hits", 0.0)),
314
+ "near_misses": int(counters.get("near_misses", 0.0)),
315
+ "coalesced": int(counters.get("coalesced", 0.0)),
316
+ "errors": int(counters.get("provider_errors", 0.0)),
317
+ "hit_rate": round(hits / requests, 4) if requests else 0.0,
318
+ "hit_rate_of_cacheable": round(hits / cacheable, 4) if cacheable else 0.0,
319
+ "stores_skipped_truncated": int(skipped_truncated),
320
+ "stores_skipped_empty": int(counters.get("stores_skipped_empty", 0.0)),
321
+ "diagnostics": diagnostics,
322
+ "usd_saved": round(counters.get("usd_saved", 0.0), 6),
323
+ "usd_spent": round(counters.get("usd_spent", 0.0), 6),
324
+ "tokens_saved": int(
325
+ counters.get("tokens_saved_prompt", 0.0)
326
+ + counters.get("tokens_saved_completion", 0.0)
327
+ ),
328
+ "latency_ms": {
329
+ "hit": await self.analytics.latency_percentiles("hit"),
330
+ "miss": await self.analytics.latency_percentiles("miss"),
331
+ },
332
+ }
@@ -0,0 +1,217 @@
1
+ """Level 2: semantic nearest-neighbour lookup backed by Redis + RedisVL.
2
+
3
+ The index is HNSW over cosine distance with a ``namespace`` tag filter, so a
4
+ query only ever considers entries that are allowed to answer it. Redis returns
5
+ cosine *distance*; because every vector is unit length, similarity is simply
6
+ ``1 - distance``.
7
+
8
+ Entries are written with the plain Redis client rather than through RedisVL so
9
+ that the hash write and its TTL land in one pipeline, and so an expiring key
10
+ takes itself out of the index with no sweeper process.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ from collections.abc import Mapping
16
+ from typing import Any, cast
17
+
18
+ import numpy as np
19
+ import redis.asyncio as aioredis
20
+ import structlog
21
+ from redis.exceptions import RedisError
22
+ from redisvl.index import AsyncSearchIndex
23
+ from redisvl.query import FilterQuery, VectorQuery
24
+ from redisvl.query.filter import Tag
25
+
26
+ from cachellm.cache.entry import CacheEntry
27
+ from cachellm.settings import Settings
28
+
29
+ log = structlog.get_logger(__name__)
30
+
31
+ RETURN_FIELDS = [
32
+ "entry_id",
33
+ "namespace",
34
+ "category",
35
+ "model",
36
+ "provider",
37
+ "prompt",
38
+ "response_text",
39
+ "prompt_tokens",
40
+ "completion_tokens",
41
+ "finish_reason",
42
+ "created_at",
43
+ "ttl",
44
+ "hits",
45
+ "exact_hash",
46
+ "tool_calls",
47
+ ]
48
+
49
+
50
+ class VectorStore:
51
+ def __init__(self, redis: aioredis.Redis, settings: Settings) -> None:
52
+ self._redis = redis
53
+ self._settings = settings
54
+ self._prefix = settings.entry_prefix
55
+ self._index: AsyncSearchIndex | None = None
56
+
57
+ # ------------------------------------------------------------------ schema
58
+ def schema(self) -> dict[str, Any]:
59
+ return {
60
+ "index": {
61
+ "name": self._settings.index_name,
62
+ "prefix": self._prefix,
63
+ "storage_type": "hash",
64
+ },
65
+ "fields": [
66
+ {"name": "namespace", "type": "tag"},
67
+ {"name": "category", "type": "tag"},
68
+ {"name": "model", "type": "tag"},
69
+ {"name": "provider", "type": "tag"},
70
+ {"name": "created_at", "type": "numeric", "attrs": {"sortable": True}},
71
+ {
72
+ "name": "embedding",
73
+ "type": "vector",
74
+ "attrs": {
75
+ "dims": self._settings.expected_dim(),
76
+ "distance_metric": "cosine",
77
+ "algorithm": "hnsw",
78
+ "datatype": "float32",
79
+ },
80
+ },
81
+ ],
82
+ }
83
+
84
+ async def connect(self) -> None:
85
+ index = AsyncSearchIndex.from_dict(self.schema(), redis_client=self._redis)
86
+ await index.create(overwrite=False)
87
+ self._index = index
88
+ log.info(
89
+ "vector_index_ready",
90
+ index=self._settings.index_name,
91
+ dims=self._settings.expected_dim(),
92
+ )
93
+
94
+ @property
95
+ def index(self) -> AsyncSearchIndex:
96
+ if self._index is None:
97
+ raise RuntimeError("VectorStore.connect() has not been awaited")
98
+ return self._index
99
+
100
+ def key(self, entry_id: str) -> str:
101
+ return f"{self._prefix}{entry_id}"
102
+
103
+ # ------------------------------------------------------------------- writes
104
+ async def put(self, entry: CacheEntry) -> None:
105
+ key = self.key(entry.entry_id)
106
+ # redis-py types the mapping key as a wide union, and Mapping is
107
+ # invariant in its key type, so a plain dict[str, ...] will not match.
108
+ # The values really are Redis-compatible; this is a variance nit.
109
+ mapping = cast("Mapping[Any, Any]", entry.to_redis_mapping())
110
+ async with self._redis.pipeline(transaction=True) as pipe:
111
+ pipe.hset(key, mapping=mapping)
112
+ if entry.ttl > 0:
113
+ pipe.expire(key, entry.ttl)
114
+ await pipe.execute()
115
+
116
+ async def touch_hit(self, entry_id: str) -> None:
117
+ """Record that an entry served a request. Never resurrects an expired key."""
118
+ await self._redis.hincrby(self.key(entry_id), "hits", 1)
119
+
120
+ async def get(self, entry_id: str) -> CacheEntry | None:
121
+ raw = await self._redis.hgetall(self.key(entry_id))
122
+ if not raw:
123
+ return None
124
+ return CacheEntry.from_redis_mapping(raw)
125
+
126
+ async def delete(self, entry_id: str) -> int:
127
+ return int(await self._redis.delete(self.key(entry_id)))
128
+
129
+ # ------------------------------------------------------------------ queries
130
+ async def search(
131
+ self, namespace: str, vector: np.ndarray, k: int
132
+ ) -> list[tuple[CacheEntry, float]]:
133
+ query = VectorQuery(
134
+ vector=np.asarray(vector, dtype=np.float32).tolist(),
135
+ vector_field_name="embedding",
136
+ num_results=k,
137
+ filter_expression=Tag("namespace") == namespace,
138
+ return_fields=RETURN_FIELDS,
139
+ return_score=True,
140
+ )
141
+ rows = await self.index.query(query)
142
+ out: list[tuple[CacheEntry, float]] = []
143
+ for row in rows:
144
+ distance = float(row.get("vector_distance", 1.0))
145
+ similarity = 1.0 - distance
146
+ out.append((CacheEntry.from_redis_mapping(row), similarity))
147
+ out.sort(key=lambda pair: pair[1], reverse=True)
148
+ return out
149
+
150
+ async def count(self) -> int:
151
+ try:
152
+ info = await self.index.info()
153
+ except RedisError: # pragma: no cover - index may not exist yet
154
+ return 0
155
+ return int(info.get("num_docs", 0) or 0)
156
+
157
+ # ------------------------------------------------------------- invalidation
158
+ async def scan_keys(self, match: str | None = None) -> list[str]:
159
+ pattern = match or f"{self._prefix}*"
160
+ keys: list[str] = []
161
+ async for key in self._redis.scan_iter(match=pattern, count=500):
162
+ keys.append(key.decode() if isinstance(key, bytes) else str(key))
163
+ return keys
164
+
165
+ async def invalidate(
166
+ self, *, namespace: str | None = None, model: str | None = None, drop_all: bool = False
167
+ ) -> int:
168
+ """Delete entries by namespace, by model, or everything.
169
+
170
+ Namespace and model are indexed tags, so this is a query, not a scan of
171
+ the whole keyspace. That is what makes "the system prompt changed, drop
172
+ its cache" a one-line operation.
173
+ """
174
+ if drop_all:
175
+ keys = await self.scan_keys()
176
+ keys += await self._scan_prefix(self._settings.exact_prefix)
177
+ if not keys:
178
+ return 0
179
+ return int(await self._redis.delete(*keys))
180
+
181
+ if namespace:
182
+ expr = Tag("namespace") == namespace
183
+ elif model:
184
+ expr = Tag("model") == model
185
+ else:
186
+ return 0
187
+
188
+ # A filter-only query, not a nearest-neighbour search. Namespace and
189
+ # model are indexed tags, so "drop everything for this system prompt"
190
+ # or "this model was upgraded" is an index lookup, never a keyspace
191
+ # scan. RedisVL escapes tag values, which matters because model ids are
192
+ # full of dots, slashes and colons.
193
+ to_delete: list[str] = []
194
+ query = FilterQuery(
195
+ filter_expression=expr,
196
+ return_fields=["entry_id", "namespace", "exact_hash"],
197
+ num_results=500,
198
+ )
199
+ async for page in self.index.paginate(query, page_size=500):
200
+ for row in page:
201
+ entry_id = CacheEntry._s(row, "entry_id")
202
+ ns = CacheEntry._s(row, "namespace")
203
+ exact = CacheEntry._s(row, "exact_hash")
204
+ if entry_id:
205
+ to_delete.append(self.key(entry_id))
206
+ if ns and exact:
207
+ to_delete.append(f"{self._settings.exact_prefix}{ns}:{exact}")
208
+
209
+ if not to_delete:
210
+ return 0
211
+ return int(await self._redis.delete(*to_delete))
212
+
213
+ async def _scan_prefix(self, prefix: str) -> list[str]:
214
+ keys: list[str] = []
215
+ async for key in self._redis.scan_iter(match=f"{prefix}*", count=500):
216
+ keys.append(key.decode() if isinstance(key, bytes) else str(key))
217
+ return keys
cachellm/cli.py ADDED
@@ -0,0 +1,122 @@
1
+ """Command line interface."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import json
7
+ from pathlib import Path
8
+ from typing import Annotated
9
+
10
+ import typer
11
+
12
+ from cachellm import __version__
13
+ from cachellm.settings import get_settings
14
+
15
+ app = typer.Typer(
16
+ name="cachellm",
17
+ help="A drop-in semantic cache for OpenAI-compatible LLM APIs.",
18
+ no_args_is_help=True,
19
+ add_completion=False,
20
+ )
21
+
22
+
23
+ @app.command()
24
+ def version() -> None:
25
+ """Print the installed version."""
26
+ typer.echo(__version__)
27
+
28
+
29
+ @app.command()
30
+ def serve(
31
+ host: Annotated[str, typer.Option(help="Bind address.")] = "",
32
+ port: Annotated[int, typer.Option(help="Bind port.")] = 0,
33
+ reload: Annotated[bool, typer.Option(help="Auto-reload on code changes.")] = False,
34
+ ) -> None:
35
+ """Run the proxy."""
36
+ import uvicorn
37
+
38
+ settings = get_settings()
39
+ uvicorn.run(
40
+ "cachellm.api.app:create_app",
41
+ factory=True,
42
+ host=host or settings.host,
43
+ port=port or settings.port,
44
+ reload=reload,
45
+ log_config=None,
46
+ )
47
+
48
+
49
+ @app.command()
50
+ def config() -> None:
51
+ """Show the effective configuration."""
52
+ settings = get_settings()
53
+ data = settings.model_dump()
54
+ data["api_keys"] = f"<{len(settings.client_keys)} key(s) configured>"
55
+ data["openai_api_key"] = "<set>" if settings.openai_api_key else "<unset>"
56
+ typer.echo(json.dumps(data, indent=2, default=str))
57
+
58
+
59
+ @app.command()
60
+ def stats() -> None:
61
+ """Print cache statistics straight from Redis."""
62
+
63
+ async def run() -> None:
64
+ from cachellm.api.deps import build_state, shutdown_state
65
+
66
+ settings = get_settings()
67
+ state = await build_state(settings)
68
+ try:
69
+ if state.cache is None:
70
+ typer.secho(f"cache unavailable: {state.degraded_reason}", fg="red")
71
+ raise typer.Exit(1)
72
+ typer.echo(json.dumps(await state.cache.stats(), indent=2))
73
+ finally:
74
+ await shutdown_state(state)
75
+
76
+ asyncio.run(run())
77
+
78
+
79
+ @app.command()
80
+ def invalidate(
81
+ namespace: Annotated[str, typer.Option(help="Namespace hash to drop.")] = "",
82
+ model: Annotated[str, typer.Option(help="Drop every entry for this model.")] = "",
83
+ all_entries: Annotated[bool, typer.Option("--all", help="Drop the whole cache.")] = False,
84
+ ) -> None:
85
+ """Remove cache entries by namespace, model, or all of them."""
86
+ if not (namespace or model or all_entries):
87
+ typer.secho("pass --namespace, --model or --all", fg="red")
88
+ raise typer.Exit(2)
89
+
90
+ async def run() -> None:
91
+ from cachellm.api.deps import build_state, shutdown_state
92
+
93
+ state = await build_state(get_settings())
94
+ try:
95
+ if state.cache is None:
96
+ typer.secho(f"cache unavailable: {state.degraded_reason}", fg="red")
97
+ raise typer.Exit(1)
98
+ removed = await state.cache.invalidate(
99
+ namespace=namespace or None, model=model or None, drop_all=all_entries
100
+ )
101
+ typer.echo(f"removed {removed} keys")
102
+ finally:
103
+ await shutdown_state(state)
104
+
105
+ asyncio.run(run())
106
+
107
+
108
+ @app.command()
109
+ def tune(
110
+ pairs_file: Annotated[Path, typer.Argument(help="JSONL of {a, b, duplicate} objects.")],
111
+ output: Annotated[Path, typer.Option(help="Write the sweep to this JSON file.")] = Path(
112
+ "threshold_sweep.json"
113
+ ),
114
+ ) -> None:
115
+ """Sweep similarity thresholds against labelled prompt pairs."""
116
+ from bench.tune_threshold import run_sweep
117
+
118
+ asyncio.run(run_sweep(pairs_file, output))
119
+
120
+
121
+ if __name__ == "__main__": # pragma: no cover
122
+ app()
@@ -0,0 +1,19 @@
1
+ from __future__ import annotations
2
+
3
+ from cachellm.embeddings.base import Embedder
4
+ from cachellm.embeddings.hash_backend import HashEmbedder
5
+ from cachellm.settings import Settings
6
+
7
+ __all__ = ["Embedder", "HashEmbedder", "build_embedder"]
8
+
9
+
10
+ def build_embedder(settings: Settings) -> Embedder:
11
+ if settings.embedding_backend == "hash":
12
+ return HashEmbedder(dim=settings.embedding_dim)
13
+ from cachellm.embeddings.fastembed_backend import FastEmbedEmbedder
14
+
15
+ return FastEmbedEmbedder(
16
+ model_name=settings.embedding_model,
17
+ dim=settings.expected_dim(),
18
+ cache_size=settings.embedding_cache_size,
19
+ )