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.
- cachellm/__init__.py +13 -0
- cachellm/__main__.py +6 -0
- cachellm/api/__init__.py +5 -0
- cachellm/api/app.py +143 -0
- cachellm/api/auth.py +34 -0
- cachellm/api/deps.py +106 -0
- cachellm/api/routes_admin.py +265 -0
- cachellm/api/routes_chat.py +385 -0
- cachellm/api/sse.py +98 -0
- cachellm/cache/__init__.py +3 -0
- cachellm/cache/analytics.py +150 -0
- cachellm/cache/coalesce.py +63 -0
- cachellm/cache/entry.py +92 -0
- cachellm/cache/exact_store.py +33 -0
- cachellm/cache/keys.py +124 -0
- cachellm/cache/policy.py +134 -0
- cachellm/cache/redis_client.py +22 -0
- cachellm/cache/service.py +332 -0
- cachellm/cache/vector_store.py +217 -0
- cachellm/cli.py +122 -0
- cachellm/embeddings/__init__.py +19 -0
- cachellm/embeddings/base.py +38 -0
- cachellm/embeddings/fastembed_backend.py +75 -0
- cachellm/embeddings/hash_backend.py +42 -0
- cachellm/errors.py +72 -0
- cachellm/logging_setup.py +56 -0
- cachellm/models.py +181 -0
- cachellm/observability/__init__.py +6 -0
- cachellm/observability/metrics.py +147 -0
- cachellm/observability/tracing.py +107 -0
- cachellm/pricing.py +108 -0
- cachellm/providers/__init__.py +7 -0
- cachellm/providers/base.py +84 -0
- cachellm/providers/bedrock.py +238 -0
- cachellm/providers/fake.py +56 -0
- cachellm/providers/openai_compat.py +131 -0
- cachellm/providers/registry.py +96 -0
- cachellm/py.typed +0 -0
- cachellm/settings.py +230 -0
- cachellm_proxy-0.1.0.dist-info/METADATA +550 -0
- cachellm_proxy-0.1.0.dist-info/RECORD +43 -0
- cachellm_proxy-0.1.0.dist-info/WHEEL +4 -0
- 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
|
+
)
|