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,92 @@
1
+ """The record stored in Redis for one cached answer."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import time
7
+ from dataclasses import asdict, dataclass, field
8
+ from typing import Any
9
+
10
+ import numpy as np
11
+
12
+
13
+ @dataclass
14
+ class CacheEntry:
15
+ entry_id: str
16
+ namespace: str
17
+ category: str
18
+ model: str
19
+ provider: str
20
+ prompt: str
21
+ response_text: str
22
+ prompt_tokens: int = 0
23
+ completion_tokens: int = 0
24
+ finish_reason: str = "stop"
25
+ created_at: float = field(default_factory=time.time)
26
+ ttl: int = 86_400
27
+ hits: int = 0
28
+ exact_hash: str = ""
29
+ embedding: np.ndarray | None = None
30
+ tool_calls: list[dict[str, Any]] | None = None
31
+
32
+ # ------------------------------------------------------------- serialisation
33
+ def to_redis_mapping(self) -> dict[str, str | bytes | int | float]:
34
+ """Flatten to a Redis hash. The vector goes in as raw float32 bytes."""
35
+ mapping: dict[str, str | bytes | int | float] = {
36
+ "entry_id": self.entry_id,
37
+ "namespace": self.namespace,
38
+ "category": self.category,
39
+ "model": self.model,
40
+ "provider": self.provider,
41
+ "prompt": self.prompt,
42
+ "response_text": self.response_text,
43
+ "prompt_tokens": int(self.prompt_tokens),
44
+ "completion_tokens": int(self.completion_tokens),
45
+ "finish_reason": self.finish_reason,
46
+ "created_at": float(self.created_at),
47
+ "ttl": int(self.ttl),
48
+ "hits": int(self.hits),
49
+ "exact_hash": self.exact_hash,
50
+ "tool_calls": json.dumps(self.tool_calls) if self.tool_calls else "",
51
+ }
52
+ if self.embedding is not None:
53
+ mapping["embedding"] = np.asarray(self.embedding, dtype=np.float32).tobytes()
54
+ return mapping
55
+
56
+ @staticmethod
57
+ def _s(raw: dict[Any, Any], key: str, default: str = "") -> str:
58
+ value = raw.get(key, raw.get(key.encode(), default))
59
+ if isinstance(value, bytes):
60
+ return value.decode("utf-8", "replace")
61
+ return str(value) if value is not None else default
62
+
63
+ @classmethod
64
+ def from_redis_mapping(cls, raw: dict[Any, Any]) -> CacheEntry:
65
+ get = cls._s
66
+ tool_calls_raw = get(raw, "tool_calls")
67
+ return cls(
68
+ entry_id=get(raw, "entry_id"),
69
+ namespace=get(raw, "namespace"),
70
+ category=get(raw, "category", "default"),
71
+ model=get(raw, "model"),
72
+ provider=get(raw, "provider"),
73
+ prompt=get(raw, "prompt"),
74
+ response_text=get(raw, "response_text"),
75
+ prompt_tokens=int(float(get(raw, "prompt_tokens", "0") or 0)),
76
+ completion_tokens=int(float(get(raw, "completion_tokens", "0") or 0)),
77
+ finish_reason=get(raw, "finish_reason", "stop"),
78
+ created_at=float(get(raw, "created_at", "0") or 0),
79
+ ttl=int(float(get(raw, "ttl", "0") or 0)),
80
+ hits=int(float(get(raw, "hits", "0") or 0)),
81
+ exact_hash=get(raw, "exact_hash"),
82
+ tool_calls=json.loads(tool_calls_raw) if tool_calls_raw else None,
83
+ )
84
+
85
+ def age_seconds(self) -> float:
86
+ return max(0.0, time.time() - self.created_at)
87
+
88
+ def summary(self) -> dict[str, Any]:
89
+ data = asdict(self)
90
+ data.pop("embedding", None)
91
+ data["age_seconds"] = round(self.age_seconds(), 1)
92
+ return data
@@ -0,0 +1,33 @@
1
+ """Level 1: exact-match lookup.
2
+
3
+ A literal repeat should never pay for an embedding. This tier is a plain
4
+ ``GET`` that resolves a normalised-prompt hash to an entry id, so it answers in
5
+ about a millisecond and takes the pressure off the vector index.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import redis.asyncio as aioredis
11
+
12
+ from cachellm.settings import Settings
13
+
14
+
15
+ class ExactStore:
16
+ def __init__(self, redis: aioredis.Redis, settings: Settings) -> None:
17
+ self._redis = redis
18
+ self._prefix = settings.exact_prefix
19
+
20
+ def key(self, namespace: str, exact_hash: str) -> str:
21
+ return f"{self._prefix}{namespace}:{exact_hash}"
22
+
23
+ async def get(self, namespace: str, exact_hash: str) -> str | None:
24
+ raw = await self._redis.get(self.key(namespace, exact_hash))
25
+ if raw is None:
26
+ return None
27
+ return raw.decode("utf-8") if isinstance(raw, bytes) else str(raw)
28
+
29
+ async def put(self, namespace: str, exact_hash: str, entry_id: str, ttl: int) -> None:
30
+ await self._redis.set(self.key(namespace, exact_hash), entry_id, ex=ttl)
31
+
32
+ async def delete(self, namespace: str, exact_hash: str) -> None:
33
+ await self._redis.delete(self.key(namespace, exact_hash))
cachellm/cache/keys.py ADDED
@@ -0,0 +1,124 @@
1
+ """Cache key derivation.
2
+
3
+ Two ideas do the heavy lifting here.
4
+
5
+ **Namespace.** Everything that changes what a *correct* answer looks like goes
6
+ into a namespace hash: provider, model, system prompt, temperature bucket,
7
+ max_tokens, response format. Two identical user questions asked under different
8
+ system prompts must never share an answer, and a model upgrade must not serve
9
+ stale text from the old model. Making that a tag on the index means
10
+ invalidation is one query, not a scan.
11
+
12
+ **Two levels.** The exact hash catches literal repeats in about a millisecond
13
+ without touching the embedding model. Only genuine near misses pay for a vector
14
+ search. In replayed traffic a large share of repeats are literal, so this tier
15
+ does real work.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import hashlib
21
+ import json
22
+ import re
23
+ import unicodedata
24
+
25
+ from cachellm.models import ChatCompletionRequest
26
+
27
+ _WS = re.compile(r"\s+")
28
+ _TRAILING_PUNCT = re.compile(r"[\s\.\?!,;:]+$")
29
+
30
+ # Words that carry no meaning for retrieval. Stripping them raises hit rate by
31
+ # collapsing polite phrasing onto the same point in vector space. Off by default
32
+ # because it also erases some genuine distinctions.
33
+ FILLER_WORDS = frozenset(
34
+ {
35
+ "please",
36
+ "could",
37
+ "would",
38
+ "can",
39
+ "you",
40
+ "kindly",
41
+ "hey",
42
+ "hi",
43
+ "hello",
44
+ "thanks",
45
+ "thank",
46
+ "sorry",
47
+ "just",
48
+ "actually",
49
+ "basically",
50
+ "really",
51
+ "simply",
52
+ "quick",
53
+ "quickly",
54
+ "tell",
55
+ "me",
56
+ "explain",
57
+ "about",
58
+ "i",
59
+ "want",
60
+ "to",
61
+ "know",
62
+ "wondering",
63
+ "help",
64
+ }
65
+ )
66
+
67
+
68
+ def normalise_text(text: str) -> str:
69
+ """Unicode-normalise, lowercase, collapse whitespace, drop trailing punctuation."""
70
+ out = unicodedata.normalize("NFKC", text or "").strip().lower()
71
+ out = _WS.sub(" ", out)
72
+ return _TRAILING_PUNCT.sub("", out)
73
+
74
+
75
+ _WORD = re.compile(r"[\w']+", re.UNICODE)
76
+
77
+
78
+ def strip_filler(text: str) -> str:
79
+ """Drop filler words, punctuation included, keeping the meaningful tokens."""
80
+ tokens = [t for t in _WORD.findall(normalise_text(text)) if t not in FILLER_WORDS]
81
+ return " ".join(tokens) if tokens else normalise_text(text)
82
+
83
+
84
+ def embedding_text(raw: str, *, strip_fillers: bool) -> str:
85
+ return strip_filler(raw) if strip_fillers else normalise_text(raw)
86
+
87
+
88
+ def temperature_bucket(temperature: float) -> str:
89
+ """Bucket to one decimal place.
90
+
91
+ 0.0 and 0.05 produce effectively the same answer distribution and should
92
+ share a cache entry; 0.2 and 0.9 should not.
93
+ """
94
+ return f"{round(float(temperature), 1):.1f}"
95
+
96
+
97
+ def namespace_for(request: ChatCompletionRequest, provider: str) -> str:
98
+ """Stable 16-hex-char id for "requests that may share answers"."""
99
+ response_format = request.response_format or {}
100
+ payload = {
101
+ "provider": provider,
102
+ "model": request.model.strip(),
103
+ "system": normalise_text(request.system_text()),
104
+ "temperature": temperature_bucket(request.effective_temperature),
105
+ "top_p": request.top_p,
106
+ "max_tokens": request.effective_max_tokens,
107
+ "response_format": response_format.get("type"),
108
+ "stop": request.stop,
109
+ }
110
+ blob = json.dumps(payload, sort_keys=True, separators=(",", ":"), default=str)
111
+ return hashlib.sha256(blob.encode("utf-8")).hexdigest()[:16]
112
+
113
+
114
+ def exact_hash(text: str) -> str:
115
+ return hashlib.sha256(normalise_text(text).encode("utf-8")).hexdigest()[:32]
116
+
117
+
118
+ def entry_id(namespace: str, exact: str) -> str:
119
+ return f"{namespace}-{exact[:20]}"
120
+
121
+
122
+ def prompt_fingerprint(text: str) -> str:
123
+ """Short, non-reversible id for a prompt, safe to put in logs and metrics."""
124
+ return hashlib.blake2b(normalise_text(text).encode("utf-8"), digest_size=6).hexdigest()
@@ -0,0 +1,134 @@
1
+ """Decides whether a request may be cached, in which category, for how long.
2
+
3
+ This module is the honesty layer. A semantic cache that caches everything will
4
+ eventually serve yesterday's stock price or one user's order status to another
5
+ user. Each rule below exists because of a specific way that goes wrong.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import re
11
+ from dataclasses import dataclass
12
+
13
+ from cachellm.models import ChatCompletionRequest
14
+ from cachellm.settings import Settings
15
+
16
+ # Anything that pins an answer to *now*. These still get cached, but with a
17
+ # short TTL, because "what is the weather in Jaipur" repeats a lot within an hour.
18
+ VOLATILE_PATTERNS = re.compile(
19
+ r"\b(today|tonight|right now|currently|current|latest|breaking|this (week|month|year|morning)|"
20
+ r"yesterday|tomorrow|as of|so far|up to date|live|news|stock price|share price|weather|"
21
+ r"forecast|who won|score|trending|recent|nowadays|these days)\b",
22
+ re.IGNORECASE,
23
+ )
24
+
25
+ # Open-ended generation: two different answers are both correct, so a cache hit
26
+ # that returns the same poem twice is a product bug, not a saving.
27
+ CREATIVE_PATTERNS = re.compile(
28
+ r"\b(write|compose|draft|generate|create|invent|imagine|brainstorm|come up with)\b.{0,40}"
29
+ r"\b(poem|story|song|lyrics|joke|essay|article|blog|caption|tagline|slogan|name|names|ideas?|"
30
+ r"script|email|letter|post)\b",
31
+ re.IGNORECASE,
32
+ )
33
+
34
+ # Constrained answer space, so near-duplicates are safe at a lower threshold.
35
+ CLASSIFICATION_PATTERNS = re.compile(
36
+ r"\b(classify|categorise|categorize|label|sentiment|is this|does this|yes or no|true or false|"
37
+ r"which category|tag this|detect|extract|rate this|score this)\b",
38
+ re.IGNORECASE,
39
+ )
40
+
41
+ # Personal data. If any of this is present the answer is about one person and
42
+ # must never be served to another, so we do not store it at all.
43
+ PII_PATTERNS: tuple[tuple[str, re.Pattern[str]], ...] = (
44
+ ("email", re.compile(r"\b[\w.+-]+@[\w-]+\.[\w.]{2,}\b")),
45
+ ("long_digits", re.compile(r"\b\d{9,}\b")), # phone, Aadhaar, account, order id
46
+ ("card", re.compile(r"\b(?:\d[ -]*?){13,16}\b")),
47
+ ("api_key", re.compile(r"\b(sk|pk|ghp|gho|AKIA|ASIA)[-_A-Za-z0-9]{12,}\b")),
48
+ (
49
+ "possessive",
50
+ re.compile(
51
+ r"\bmy (order|account|booking|policy|invoice|ticket|card|salary|"
52
+ r"password|address|phone|otp|pan|aadhaar|ssn)\b",
53
+ re.IGNORECASE,
54
+ ),
55
+ ),
56
+ )
57
+
58
+
59
+ @dataclass(frozen=True)
60
+ class PolicyDecision:
61
+ cacheable: bool
62
+ category: str
63
+ threshold: float
64
+ ttl: int
65
+ reason: str = ""
66
+
67
+ @property
68
+ def bypass_reason(self) -> str:
69
+ return "" if self.cacheable else self.reason
70
+
71
+
72
+ def _classify(text: str, request: ChatCompletionRequest) -> str:
73
+ if VOLATILE_PATTERNS.search(text):
74
+ return "volatile"
75
+ if CREATIVE_PATTERNS.search(text):
76
+ return "creative"
77
+ max_tokens = request.effective_max_tokens
78
+ if CLASSIFICATION_PATTERNS.search(text) or (max_tokens is not None and max_tokens <= 32):
79
+ return "classification"
80
+ if request.user_turn_count() > 1:
81
+ return "conversational"
82
+ return "factual"
83
+
84
+
85
+ def detect_pii(text: str) -> str | None:
86
+ for label, pattern in PII_PATTERNS:
87
+ if pattern.search(text):
88
+ return label
89
+ return None
90
+
91
+
92
+ def decide(
93
+ request: ChatCompletionRequest,
94
+ settings: Settings,
95
+ *,
96
+ cache_control: str = "",
97
+ ) -> PolicyDecision:
98
+ """Return the caching decision for one request."""
99
+ text = request.last_user_text()
100
+ category = _classify(text, request)
101
+ threshold = settings.threshold_for(category)
102
+ ttl = settings.ttl_for(category)
103
+
104
+ def no(reason: str) -> PolicyDecision:
105
+ return PolicyDecision(False, category, threshold, ttl, reason)
106
+
107
+ control = (cache_control or "").lower()
108
+ if "no-store" in control or "no-cache" in control:
109
+ return no("client_requested_bypass")
110
+ if not settings.enabled:
111
+ return no("cache_disabled")
112
+ if not text:
113
+ return no("no_user_message")
114
+ if len(text) > settings.max_prompt_chars:
115
+ return no("prompt_too_long")
116
+ if request.effective_temperature > settings.max_cacheable_temperature:
117
+ return no("temperature_too_high")
118
+ if (request.n or 1) > 1:
119
+ return no("multiple_completions_requested")
120
+ if request.tools and not settings.cache_tool_calls:
121
+ return no("tool_calls")
122
+ if any(m.has_non_text_parts() for m in request.messages):
123
+ return no("non_text_content")
124
+ fmt = (request.response_format or {}).get("type")
125
+ if fmt in ("json_object", "json_schema") and not settings.cache_json_mode:
126
+ return no("json_mode")
127
+ if request.user_turn_count() > 1 and not settings.cache_multi_turn:
128
+ return no("multi_turn_conversation")
129
+ if settings.pii_guard:
130
+ found = detect_pii(text)
131
+ if found:
132
+ return no(f"pii:{found}")
133
+
134
+ return PolicyDecision(True, category, threshold, ttl, "")
@@ -0,0 +1,22 @@
1
+ """Redis connection factory.
2
+
3
+ ``decode_responses`` is deliberately False. Embeddings are stored as raw
4
+ float32 bytes and a decoding client would corrupt them on the way back out;
5
+ text fields are decoded explicitly where they are read.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import redis.asyncio as aioredis
11
+
12
+ from cachellm.settings import Settings
13
+
14
+
15
+ def build_redis(settings: Settings) -> aioredis.Redis:
16
+ return aioredis.from_url(
17
+ settings.redis_url,
18
+ decode_responses=False,
19
+ health_check_interval=30,
20
+ socket_connect_timeout=5,
21
+ socket_keepalive=True,
22
+ )