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,38 @@
1
+ from __future__ import annotations
2
+
3
+ import abc
4
+
5
+ import numpy as np
6
+
7
+
8
+ class Embedder(abc.ABC):
9
+ """Turns text into a unit-length float32 vector.
10
+
11
+ Vectors are L2-normalised on the way out, which means the cosine distance
12
+ Redis reports can be turned into a similarity with ``1 - distance`` and no
13
+ extra maths at query time.
14
+ """
15
+
16
+ name: str
17
+ dim: int
18
+
19
+ @abc.abstractmethod
20
+ async def embed(self, text: str) -> np.ndarray: ...
21
+
22
+ @abc.abstractmethod
23
+ async def embed_batch(self, texts: list[str]) -> list[np.ndarray]: ...
24
+
25
+ async def warmup(self) -> None:
26
+ await self.embed("warmup")
27
+
28
+ @staticmethod
29
+ def normalise(vector: np.ndarray) -> np.ndarray:
30
+ vec = np.asarray(vector, dtype=np.float32)
31
+ norm = float(np.linalg.norm(vec))
32
+ if norm == 0.0:
33
+ return vec
34
+ return (vec / norm).astype(np.float32)
35
+
36
+ @staticmethod
37
+ def to_bytes(vector: np.ndarray) -> bytes:
38
+ return np.asarray(vector, dtype=np.float32).tobytes()
@@ -0,0 +1,75 @@
1
+ """Local ONNX embeddings via fastembed.
2
+
3
+ Running the embedding model *inside* the proxy is the single most important
4
+ latency decision in the project. A hosted embedding API costs 80-200 ms from
5
+ India, which would become the floor for every cache hit and destroy the whole
6
+ point. bge-small on CPU is single-digit milliseconds and costs nothing.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import asyncio
12
+ from collections import OrderedDict
13
+ from typing import Any
14
+
15
+ import numpy as np
16
+
17
+ from cachellm.embeddings.base import Embedder
18
+
19
+
20
+ class FastEmbedEmbedder(Embedder):
21
+ def __init__(self, model_name: str, dim: int, cache_size: int = 2048) -> None:
22
+ self.name = model_name
23
+ self.dim = dim
24
+ self._cache_size = cache_size
25
+ self._cache: OrderedDict[str, np.ndarray] = OrderedDict()
26
+ self._model: Any = None
27
+ self._lock = asyncio.Lock()
28
+
29
+ async def _ensure_model(self) -> Any:
30
+ """Load the ONNX model once, off the event loop.
31
+
32
+ Loading takes a second or two and is CPU-bound, so it goes to a worker
33
+ thread; the lock stops a burst of concurrent first requests from each
34
+ loading their own copy.
35
+ """
36
+ if self._model is None:
37
+ async with self._lock:
38
+ if self._model is None:
39
+ from fastembed import TextEmbedding
40
+
41
+ def load() -> Any:
42
+ return TextEmbedding(model_name=self.name)
43
+
44
+ self._model = await asyncio.to_thread(load)
45
+ return self._model
46
+
47
+ def _cache_get(self, text: str) -> np.ndarray | None:
48
+ vec = self._cache.get(text)
49
+ if vec is not None:
50
+ self._cache.move_to_end(text)
51
+ return vec
52
+
53
+ def _cache_put(self, text: str, vector: np.ndarray) -> None:
54
+ self._cache[text] = vector
55
+ self._cache.move_to_end(text)
56
+ while len(self._cache) > self._cache_size:
57
+ self._cache.popitem(last=False)
58
+
59
+ async def embed(self, text: str) -> np.ndarray:
60
+ cached = self._cache_get(text)
61
+ if cached is not None:
62
+ return cached
63
+ vectors = await self.embed_batch([text])
64
+ return vectors[0]
65
+
66
+ async def embed_batch(self, texts: list[str]) -> list[np.ndarray]:
67
+ if not texts:
68
+ return []
69
+ model = await self._ensure_model()
70
+ pending = [t for t in texts if self._cache_get(t) is None]
71
+ if pending:
72
+ raw = await asyncio.to_thread(lambda: list(model.embed(pending)))
73
+ for text, vector in zip(pending, raw, strict=True):
74
+ self._cache_put(text, self.normalise(np.asarray(vector, dtype=np.float32)))
75
+ return [self._cache[t] for t in texts]
@@ -0,0 +1,42 @@
1
+ """Deterministic hashing embedder.
2
+
3
+ No model download, no network, identical output on every machine, which makes
4
+ it the right backend for unit tests and CI. It uses the classic hashing trick
5
+ (signed token buckets), so it still produces meaningful lexical similarity:
6
+ shared words pull vectors together. It does NOT understand meaning, so it is
7
+ not a substitute for a real model in the evaluation numbers.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import hashlib
13
+ import re
14
+
15
+ import numpy as np
16
+
17
+ from cachellm.embeddings.base import Embedder
18
+
19
+ _TOKEN = re.compile(r"[a-z0-9']+")
20
+
21
+
22
+ class HashEmbedder(Embedder):
23
+ def __init__(self, dim: int = 384) -> None:
24
+ self.name = "hash-embedder"
25
+ self.dim = dim
26
+
27
+ def _vector(self, text: str) -> np.ndarray:
28
+ vec = np.zeros(self.dim, dtype=np.float32)
29
+ tokens = _TOKEN.findall(text.lower())
30
+ for token in tokens:
31
+ digest = hashlib.blake2b(token.encode("utf-8"), digest_size=8).digest()
32
+ value = int.from_bytes(digest, "big")
33
+ bucket = value % self.dim
34
+ sign = 1.0 if (value >> 63) & 1 else -1.0
35
+ vec[bucket] += sign
36
+ return self.normalise(vec)
37
+
38
+ async def embed(self, text: str) -> np.ndarray:
39
+ return self._vector(text)
40
+
41
+ async def embed_batch(self, texts: list[str]) -> list[np.ndarray]:
42
+ return [self._vector(t) for t in texts]
cachellm/errors.py ADDED
@@ -0,0 +1,72 @@
1
+ """OpenAI-shaped error envelope.
2
+
3
+ A drop-in proxy has to fail the way the thing it replaces fails, otherwise
4
+ client SDK error handling breaks. The official clients read
5
+ ``{"error": {"message", "type", "param", "code"}}``.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import Any
11
+
12
+ from fastapi import HTTPException
13
+
14
+
15
+ class CacheLLMError(HTTPException):
16
+ def __init__(
17
+ self,
18
+ status_code: int,
19
+ message: str,
20
+ err_type: str = "invalid_request_error",
21
+ code: str | None = None,
22
+ param: str | None = None,
23
+ ) -> None:
24
+ super().__init__(status_code=status_code, detail=message)
25
+ self.message = message
26
+ self.err_type = err_type
27
+ self.code = code
28
+ self.param = param
29
+
30
+ def envelope(self) -> dict[str, Any]:
31
+ return {
32
+ "error": {
33
+ "message": self.message,
34
+ "type": self.err_type,
35
+ "param": self.param,
36
+ "code": self.code,
37
+ }
38
+ }
39
+
40
+
41
+ class AuthError(CacheLLMError):
42
+ def __init__(self, message: str = "Invalid API key provided.") -> None:
43
+ super().__init__(401, message, "invalid_request_error", code="invalid_api_key")
44
+
45
+
46
+ class UpstreamError(CacheLLMError):
47
+ def __init__(self, message: str, status_code: int = 502) -> None:
48
+ super().__init__(status_code, message, "api_error", code="upstream_error")
49
+
50
+
51
+ class ModelNotFoundError(CacheLLMError):
52
+ def __init__(self, model: str) -> None:
53
+ super().__init__(
54
+ 404,
55
+ f"The model `{model}` does not exist or you do not have access to it.",
56
+ "invalid_request_error",
57
+ code="model_not_found",
58
+ param="model",
59
+ )
60
+
61
+
62
+ class CacheMissError(CacheLLMError):
63
+ """Raised when a client sent ``X-Cache-Control: only-if-cached`` and missed."""
64
+
65
+ def __init__(self) -> None:
66
+ super().__init__(
67
+ 504,
68
+ "No cached response above the similarity threshold and the request "
69
+ "was sent with X-Cache-Control: only-if-cached.",
70
+ "api_error",
71
+ code="cache_miss",
72
+ )
@@ -0,0 +1,56 @@
1
+ """structlog configuration.
2
+
3
+ Prompt and completion text is treated as sensitive by default. A proxy sees
4
+ every question every user asks, so logging it wholesale is a privacy incident
5
+ waiting to happen. Set ``CACHELLM_LOG_PROMPTS=true`` only in local development.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import logging
11
+ import sys
12
+ from typing import Any
13
+
14
+ import structlog
15
+
16
+ from cachellm.settings import Settings
17
+
18
+ SENSITIVE_KEYS = {"prompt", "matched_prompt", "response_text", "messages", "content", "answer"}
19
+
20
+
21
+ def _redact_prompts(_logger: Any, _name: str, event_dict: dict[str, Any]) -> dict[str, Any]:
22
+ for key in list(event_dict):
23
+ if key in SENSITIVE_KEYS:
24
+ value = event_dict[key]
25
+ if isinstance(value, str):
26
+ event_dict[key] = f"<redacted {len(value)} chars>"
27
+ else:
28
+ event_dict[key] = "<redacted>"
29
+ return event_dict
30
+
31
+
32
+ def configure_logging(settings: Settings) -> None:
33
+ level = getattr(logging, settings.log_level.upper(), logging.INFO)
34
+ logging.basicConfig(format="%(message)s", stream=sys.stdout, level=level)
35
+
36
+ processors: list[Any] = [
37
+ structlog.contextvars.merge_contextvars,
38
+ structlog.processors.add_log_level,
39
+ structlog.processors.TimeStamper(fmt="iso", utc=True),
40
+ ]
41
+ if not settings.log_prompts:
42
+ processors.append(_redact_prompts)
43
+ processors.append(structlog.processors.StackInfoRenderer())
44
+ processors.append(structlog.processors.format_exc_info)
45
+ processors.append(
46
+ structlog.processors.JSONRenderer()
47
+ if settings.log_json
48
+ else structlog.dev.ConsoleRenderer(colors=sys.stdout.isatty())
49
+ )
50
+
51
+ structlog.configure(
52
+ processors=processors,
53
+ wrapper_class=structlog.make_filtering_bound_logger(level),
54
+ logger_factory=structlog.PrintLoggerFactory(),
55
+ cache_logger_on_first_use=True,
56
+ )
cachellm/models.py ADDED
@@ -0,0 +1,181 @@
1
+ """Request and response schemas mirroring the OpenAI chat completions contract.
2
+
3
+ Unknown fields are preserved rather than rejected: OpenAI adds parameters
4
+ regularly and a proxy that 422s on an unrecognised key is not a drop-in.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import time
10
+ import uuid
11
+ from typing import Any, Literal
12
+
13
+ from pydantic import BaseModel, ConfigDict, Field
14
+
15
+ Role = Literal["system", "user", "assistant", "tool", "developer", "function"]
16
+
17
+
18
+ class ChatMessage(BaseModel):
19
+ model_config = ConfigDict(extra="allow")
20
+
21
+ role: Role
22
+ content: str | list[dict[str, Any]] | None = None
23
+ name: str | None = None
24
+ tool_calls: list[dict[str, Any]] | None = None
25
+ tool_call_id: str | None = None
26
+
27
+ def as_text(self) -> str:
28
+ """Flatten multi-part content down to plain text for hashing/embedding."""
29
+ if self.content is None:
30
+ return ""
31
+ if isinstance(self.content, str):
32
+ return self.content
33
+ parts: list[str] = []
34
+ for part in self.content:
35
+ if not isinstance(part, dict):
36
+ continue
37
+ if part.get("type") in (None, "text") and isinstance(part.get("text"), str):
38
+ parts.append(part["text"])
39
+ return "\n".join(parts)
40
+
41
+ def has_non_text_parts(self) -> bool:
42
+ if isinstance(self.content, list):
43
+ return any(
44
+ isinstance(p, dict) and p.get("type") not in (None, "text") for p in self.content
45
+ )
46
+ return False
47
+
48
+
49
+ class ChatCompletionRequest(BaseModel):
50
+ model_config = ConfigDict(extra="allow", protected_namespaces=())
51
+
52
+ model: str
53
+ messages: list[ChatMessage]
54
+ temperature: float | None = None
55
+ top_p: float | None = None
56
+ n: int | None = None
57
+ stream: bool | None = None
58
+ stream_options: dict[str, Any] | None = None
59
+ stop: str | list[str] | None = None
60
+ max_tokens: int | None = None
61
+ max_completion_tokens: int | None = None
62
+ presence_penalty: float | None = None
63
+ frequency_penalty: float | None = None
64
+ seed: int | None = None
65
+ user: str | None = None
66
+ response_format: dict[str, Any] | None = None
67
+ tools: list[dict[str, Any]] | None = None
68
+ tool_choice: str | dict[str, Any] | None = None
69
+
70
+ # ------------------------------------------------------------------ helpers
71
+ @property
72
+ def effective_temperature(self) -> float:
73
+ return 1.0 if self.temperature is None else float(self.temperature)
74
+
75
+ @property
76
+ def effective_max_tokens(self) -> int | None:
77
+ return self.max_completion_tokens or self.max_tokens
78
+
79
+ def system_text(self) -> str:
80
+ return "\n".join(
81
+ m.as_text() for m in self.messages if m.role in ("system", "developer")
82
+ ).strip()
83
+
84
+ def conversation_messages(self) -> list[ChatMessage]:
85
+ return [m for m in self.messages if m.role not in ("system", "developer")]
86
+
87
+ def last_user_text(self) -> str:
88
+ for message in reversed(self.messages):
89
+ if message.role == "user":
90
+ return message.as_text().strip()
91
+ return ""
92
+
93
+ def user_turn_count(self) -> int:
94
+ return sum(1 for m in self.messages if m.role == "user")
95
+
96
+ def cache_text(self, include_history: bool) -> str:
97
+ """The text that gets embedded and hashed."""
98
+ if not include_history:
99
+ return self.last_user_text()
100
+ turns = [f"{m.role}: {m.as_text().strip()}" for m in self.conversation_messages()]
101
+ return "\n".join(t for t in turns if t.strip())
102
+
103
+
104
+ class Usage(BaseModel):
105
+ model_config = ConfigDict(extra="allow")
106
+
107
+ prompt_tokens: int = 0
108
+ completion_tokens: int = 0
109
+ total_tokens: int = 0
110
+
111
+
112
+ class ChoiceMessage(BaseModel):
113
+ model_config = ConfigDict(extra="allow")
114
+
115
+ role: Literal["assistant"] = "assistant"
116
+ content: str | None = None
117
+ tool_calls: list[dict[str, Any]] | None = None
118
+
119
+
120
+ class Choice(BaseModel):
121
+ model_config = ConfigDict(extra="allow")
122
+
123
+ index: int = 0
124
+ message: ChoiceMessage
125
+ finish_reason: str | None = "stop"
126
+ logprobs: None = None
127
+
128
+
129
+ class ChatCompletionResponse(BaseModel):
130
+ model_config = ConfigDict(extra="allow", protected_namespaces=())
131
+
132
+ id: str = Field(default_factory=lambda: f"chatcmpl-{uuid.uuid4().hex[:24]}")
133
+ object: Literal["chat.completion"] = "chat.completion"
134
+ created: int = Field(default_factory=lambda: int(time.time()))
135
+ model: str
136
+ choices: list[Choice]
137
+ usage: Usage = Field(default_factory=Usage)
138
+ system_fingerprint: str | None = None
139
+
140
+ @classmethod
141
+ def from_text(
142
+ cls,
143
+ *,
144
+ model: str,
145
+ text: str,
146
+ prompt_tokens: int = 0,
147
+ completion_tokens: int = 0,
148
+ finish_reason: str = "stop",
149
+ ) -> ChatCompletionResponse:
150
+ return cls(
151
+ model=model,
152
+ choices=[
153
+ Choice(
154
+ index=0,
155
+ message=ChoiceMessage(role="assistant", content=text),
156
+ finish_reason=finish_reason,
157
+ )
158
+ ],
159
+ usage=Usage(
160
+ prompt_tokens=prompt_tokens,
161
+ completion_tokens=completion_tokens,
162
+ total_tokens=prompt_tokens + completion_tokens,
163
+ ),
164
+ )
165
+
166
+ def text(self) -> str:
167
+ if not self.choices:
168
+ return ""
169
+ return self.choices[0].message.content or ""
170
+
171
+
172
+ class ModelCard(BaseModel):
173
+ id: str
174
+ object: Literal["model"] = "model"
175
+ created: int = Field(default_factory=lambda: int(time.time()))
176
+ owned_by: str = "cachellm"
177
+
178
+
179
+ class ModelList(BaseModel):
180
+ object: Literal["list"] = "list"
181
+ data: list[ModelCard]
@@ -0,0 +1,6 @@
1
+ from __future__ import annotations
2
+
3
+ from cachellm.observability.metrics import Metrics, get_metrics, reset_metrics
4
+ from cachellm.observability.tracing import setup_tracing, span
5
+
6
+ __all__ = ["Metrics", "get_metrics", "reset_metrics", "setup_tracing", "span"]
@@ -0,0 +1,147 @@
1
+ """Prometheus metrics.
2
+
3
+ The metric that matters most is ``cachellm_cost_usd_total{kind="saved"}``. It is
4
+ computed once, at serve time, from the token counts stored on the cache entry
5
+ multiplied by that model's price. Everything downstream (the Grafana panel, the
6
+ admin endpoint, the README headline) reads the same number, so they can never
7
+ disagree with each other.
8
+
9
+ Label cardinality is kept deliberately low: model and provider are bounded
10
+ sets, category has six values, result has five. Prompt text never becomes a
11
+ label.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ from prometheus_client import (
17
+ CONTENT_TYPE_LATEST,
18
+ CollectorRegistry,
19
+ Counter,
20
+ Gauge,
21
+ Histogram,
22
+ generate_latest,
23
+ )
24
+
25
+ # Buckets tuned for this workload: cache hits land in single-digit milliseconds,
26
+ # provider calls in hundreds of milliseconds to seconds.
27
+ LATENCY_BUCKETS = (0.001, 0.002, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0)
28
+ SIMILARITY_BUCKETS = (0.5, 0.6, 0.7, 0.75, 0.8, 0.85, 0.88, 0.9, 0.92, 0.94, 0.96, 0.98, 0.99, 1.0)
29
+
30
+
31
+ class Metrics:
32
+ def __init__(self, registry: CollectorRegistry | None = None) -> None:
33
+ self.registry = registry or CollectorRegistry()
34
+
35
+ self.requests = Counter(
36
+ "cachellm_requests_total",
37
+ "Chat completion requests handled, by cache outcome.",
38
+ ["result", "category", "provider", "model"],
39
+ registry=self.registry,
40
+ )
41
+ self.request_duration = Histogram(
42
+ "cachellm_request_duration_seconds",
43
+ "End-to-end request duration.",
44
+ ["result"],
45
+ buckets=LATENCY_BUCKETS,
46
+ registry=self.registry,
47
+ )
48
+ self.lookup_duration = Histogram(
49
+ "cachellm_lookup_duration_seconds",
50
+ "Time spent deciding hit or miss (policy, embedding, vector search).",
51
+ ["tier"],
52
+ buckets=LATENCY_BUCKETS,
53
+ registry=self.registry,
54
+ )
55
+ self.embed_duration = Histogram(
56
+ "cachellm_embed_duration_seconds",
57
+ "Time spent embedding the prompt.",
58
+ buckets=LATENCY_BUCKETS,
59
+ registry=self.registry,
60
+ )
61
+ self.similarity = Histogram(
62
+ "cachellm_similarity",
63
+ "Best similarity score seen at lookup time.",
64
+ ["result"],
65
+ buckets=SIMILARITY_BUCKETS,
66
+ registry=self.registry,
67
+ )
68
+ self.tokens = Counter(
69
+ "cachellm_tokens_total",
70
+ "Tokens, split by whether they were spent upstream or avoided.",
71
+ ["kind", "direction", "model"],
72
+ registry=self.registry,
73
+ )
74
+ self.cost = Counter(
75
+ "cachellm_cost_usd_total",
76
+ "Modelled USD, split into spent upstream and saved by the cache.",
77
+ ["kind", "model"],
78
+ registry=self.registry,
79
+ )
80
+ self.entries = Gauge(
81
+ "cachellm_cache_entries",
82
+ "Documents currently in the vector index.",
83
+ registry=self.registry,
84
+ )
85
+ self.provider_errors = Counter(
86
+ "cachellm_provider_errors_total",
87
+ "Upstream provider failures.",
88
+ ["provider"],
89
+ registry=self.registry,
90
+ )
91
+ self.coalesced = Counter(
92
+ "cachellm_coalesced_total",
93
+ "Requests that joined an in-flight identical call instead of duplicating it.",
94
+ registry=self.registry,
95
+ )
96
+ self.near_misses = Counter(
97
+ "cachellm_near_miss_total",
98
+ "Lookups that landed just below the similarity threshold.",
99
+ ["category"],
100
+ registry=self.registry,
101
+ )
102
+ self.shadow_hits = Counter(
103
+ "cachellm_shadow_hits_total",
104
+ "Requests the cache would have served while running in shadow mode.",
105
+ ["category"],
106
+ registry=self.registry,
107
+ )
108
+
109
+ # ------------------------------------------------------------------ helpers
110
+ def observe_tokens(
111
+ self, model: str, prompt_tokens: int, completion_tokens: int, kind: str
112
+ ) -> None:
113
+ if prompt_tokens:
114
+ self.tokens.labels(kind=kind, direction="input", model=model).inc(prompt_tokens)
115
+ if completion_tokens:
116
+ self.tokens.labels(kind=kind, direction="output", model=model).inc(completion_tokens)
117
+
118
+ def observe_cost(self, model: str, usd: float, kind: str) -> None:
119
+ if usd:
120
+ self.cost.labels(kind=kind, model=model).inc(usd)
121
+
122
+ def render(self) -> tuple[bytes, str]:
123
+ """Serialise for a Prometheus scrape.
124
+
125
+ The content type has to match the body. `generate_latest` emits the
126
+ Prometheus text format; advertising OpenMetrics alongside it makes
127
+ Prometheus reject the scrape with "data does not end with # EOF",
128
+ because OpenMetrics requires that terminator and the text format has
129
+ no such thing. Nothing but a real scrape catches this.
130
+ """
131
+ return generate_latest(self.registry), CONTENT_TYPE_LATEST
132
+
133
+
134
+ _metrics: Metrics | None = None
135
+
136
+
137
+ def get_metrics() -> Metrics:
138
+ global _metrics
139
+ if _metrics is None:
140
+ _metrics = Metrics()
141
+ return _metrics
142
+
143
+
144
+ def reset_metrics() -> None:
145
+ """Test helper: start from a clean registry."""
146
+ global _metrics
147
+ _metrics = None