tokenbiryani 0.2.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 (51) hide show
  1. tokenbiryani/__init__.py +3 -0
  2. tokenbiryani/api/__init__.py +0 -0
  3. tokenbiryani/api/app.py +583 -0
  4. tokenbiryani/api/asgi.py +32 -0
  5. tokenbiryani/cli.py +1045 -0
  6. tokenbiryani/config.py +532 -0
  7. tokenbiryani/core/__init__.py +0 -0
  8. tokenbiryani/core/account.py +258 -0
  9. tokenbiryani/core/batch.py +135 -0
  10. tokenbiryani/core/breaker.py +53 -0
  11. tokenbiryani/core/cacheadvice.py +239 -0
  12. tokenbiryani/core/diagnostics.py +131 -0
  13. tokenbiryani/core/estimator.py +180 -0
  14. tokenbiryani/core/gateway.py +2395 -0
  15. tokenbiryani/core/handoff.py +87 -0
  16. tokenbiryani/core/keys.py +199 -0
  17. tokenbiryani/core/limits.py +440 -0
  18. tokenbiryani/core/oauth.py +222 -0
  19. tokenbiryani/core/pacing.py +320 -0
  20. tokenbiryani/core/queue.py +132 -0
  21. tokenbiryani/core/router.py +323 -0
  22. tokenbiryani/core/secrets.py +114 -0
  23. tokenbiryani/core/session.py +117 -0
  24. tokenbiryani/dashboard/__init__.py +56 -0
  25. tokenbiryani/dashboard/console.css +610 -0
  26. tokenbiryani/dashboard/console.html +3250 -0
  27. tokenbiryani/observability/__init__.py +0 -0
  28. tokenbiryani/observability/events.py +171 -0
  29. tokenbiryani/observability/usage.py +226 -0
  30. tokenbiryani/prices.yaml +77 -0
  31. tokenbiryani/providers/__init__.py +0 -0
  32. tokenbiryani/providers/anthropic_api.py +118 -0
  33. tokenbiryani/providers/base.py +173 -0
  34. tokenbiryani/providers/bedrock.py +182 -0
  35. tokenbiryani/providers/oauth.py +165 -0
  36. tokenbiryani/providers/oauth_credentials.py +293 -0
  37. tokenbiryani/providers/translate.py +35 -0
  38. tokenbiryani/providers/vertex.py +144 -0
  39. tokenbiryani/proxy/__init__.py +0 -0
  40. tokenbiryani/proxy/errors.py +169 -0
  41. tokenbiryani/proxy/sse.py +98 -0
  42. tokenbiryani/store/__init__.py +0 -0
  43. tokenbiryani/store/base.py +150 -0
  44. tokenbiryani/store/memory.py +149 -0
  45. tokenbiryani/store/redis_store.py +222 -0
  46. tokenbiryani/store/sqlite.py +336 -0
  47. tokenbiryani-0.2.0.dist-info/METADATA +697 -0
  48. tokenbiryani-0.2.0.dist-info/RECORD +51 -0
  49. tokenbiryani-0.2.0.dist-info/WHEEL +4 -0
  50. tokenbiryani-0.2.0.dist-info/entry_points.txt +2 -0
  51. tokenbiryani-0.2.0.dist-info/licenses/LICENSE +202 -0
@@ -0,0 +1,258 @@
1
+ """Account runtime: health state machine, rolling stats, and cost accounting."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections import deque
6
+ from dataclasses import dataclass, field
7
+ from enum import Enum
8
+ from typing import Any, Deque, Dict, Optional, Set
9
+
10
+ from ..config import AccountConfig, ModelPrice
11
+ from ..proxy.errors import AccountAction, Classification
12
+ from .breaker import BreakerState, CircuitBreaker
13
+ from .limits import LimitMirror, TokenEstimate
14
+
15
+ WINDOW = 50
16
+
17
+
18
+ class AccountState(str, Enum):
19
+ #: serving traffic, headroom available
20
+ READY = "ready"
21
+ #: waiting on a reset; self-heals, nothing for an operator to do
22
+ COOLING = "cooling"
23
+ #: auth failure, spend cap, or a tripped breaker — needs a look
24
+ DISABLED = "disabled"
25
+
26
+
27
+ @dataclass
28
+ class Usage:
29
+ input_tokens: int = 0
30
+ output_tokens: int = 0
31
+ cache_read_tokens: int = 0
32
+ cache_creation_tokens: int = 0
33
+
34
+ @classmethod
35
+ def from_body(cls, body: Optional[Dict[str, Any]]) -> Usage:
36
+ raw = (body or {}).get("usage") or {}
37
+
38
+ def get(name: str) -> int:
39
+ try:
40
+ return int(raw.get(name) or 0)
41
+ except (TypeError, ValueError):
42
+ return 0
43
+
44
+ return cls(
45
+ input_tokens=get("input_tokens"),
46
+ output_tokens=get("output_tokens"),
47
+ cache_read_tokens=get("cache_read_input_tokens"),
48
+ cache_creation_tokens=get("cache_creation_input_tokens"),
49
+ )
50
+
51
+ @property
52
+ def billed_input(self) -> int:
53
+ return self.input_tokens + self.cache_read_tokens + self.cache_creation_tokens
54
+
55
+ def cost(self, price: Optional[ModelPrice]) -> Optional[float]:
56
+ if price is None:
57
+ return None
58
+ return (
59
+ self.input_tokens * price.input
60
+ + self.output_tokens * price.output
61
+ + self.cache_read_tokens * price.cache_read
62
+ + self.cache_creation_tokens * price.cache_write
63
+ ) / 1_000_000.0
64
+
65
+
66
+ @dataclass
67
+ class AccountRuntime:
68
+ config: AccountConfig
69
+ mirror: LimitMirror = field(default_factory=LimitMirror)
70
+ breaker: CircuitBreaker = field(default_factory=CircuitBreaker)
71
+
72
+ cooling_until: float = 0.0
73
+ disabled_reason: Optional[str] = None
74
+ unsupported_models: Set[str] = field(default_factory=set)
75
+
76
+ #: The `anthropic-ratelimit-*` headers seen on the last real response from this
77
+ #: account, kept verbatim so the header check has something true to inspect.
78
+ #:
79
+ #: It cannot use a `/v1/models` probe: that endpoint carries no limit headers, so
80
+ #: checking against it reports all nine missing for a perfectly healthy account.
81
+ #: Empty means no request has been through this account yet — which is a third
82
+ #: answer, distinct from "checked and fine" and "checked and broken".
83
+ last_limit_headers: Dict[str, str] = field(default_factory=dict)
84
+
85
+ inflight: int = 0
86
+ outcomes: Deque[bool] = field(default_factory=lambda: deque(maxlen=WINDOW))
87
+ latencies: Deque[float] = field(default_factory=lambda: deque(maxlen=WINDOW))
88
+
89
+ requests_total: int = 0
90
+ failures_total: int = 0
91
+ error_kinds: Dict[str, int] = field(default_factory=dict)
92
+ spend_usd: float = 0.0
93
+ input_tokens_total: int = 0
94
+ output_tokens_total: int = 0
95
+ cache_read_total: int = 0
96
+ cache_creation_total: int = 0
97
+
98
+ @property
99
+ def id(self) -> str:
100
+ return self.config.id
101
+
102
+ # ---- state ---------------------------------------------------------------
103
+
104
+ def observe_headers(self, headers: Dict[str, str], now: float) -> None:
105
+ """Update the mirror, and remember the limit headers it was given."""
106
+ self.mirror.update_from_headers(headers, now)
107
+ seen = {
108
+ str(k).lower(): str(v)
109
+ for k, v in headers.items()
110
+ if str(k).lower().startswith("anthropic-ratelimit-")
111
+ }
112
+ if seen:
113
+ self.last_limit_headers = seen
114
+
115
+ def state(self, now: float) -> AccountState:
116
+ if not self.config.enabled:
117
+ return AccountState.DISABLED
118
+ if self.disabled_reason:
119
+ return AccountState.DISABLED
120
+ if self.breaker.state(now) is BreakerState.OPEN:
121
+ return AccountState.DISABLED
122
+ if now < self.cooling_until:
123
+ return AccountState.COOLING
124
+ return AccountState.READY
125
+
126
+ def cooling_for(self, now: float) -> Optional[float]:
127
+ """Seconds until this account is usable again, from any known cause."""
128
+ candidates = []
129
+ if now < self.cooling_until:
130
+ candidates.append(self.cooling_until - now)
131
+ if self.breaker.state(now) is BreakerState.OPEN and self.breaker.opened_at is not None:
132
+ candidates.append(self.breaker.opened_at + self.breaker.cooldown_seconds - now)
133
+ reset = self.mirror.next_reset(now)
134
+ if reset is not None and not self.mirror.headroom(now):
135
+ candidates.append(reset)
136
+ return min(candidates) if candidates else None
137
+
138
+ def eligible(self, model: str, estimate: TokenEstimate, now: float) -> Optional[str]:
139
+ """Return None if this account may serve the request, else why not."""
140
+ if self.state(now) is AccountState.DISABLED:
141
+ return self.disabled_reason or "disabled"
142
+ if self.state(now) is AccountState.COOLING:
143
+ remaining = self.cooling_for(now) or 0.0
144
+ return f"cooling, {remaining:.0f}s remaining"
145
+ if not self.config.supports_model(model):
146
+ return "model not in account allowlist"
147
+ if model in self.unsupported_models:
148
+ return "model rejected by upstream"
149
+ if self.inflight >= self.config.max_concurrency:
150
+ return f"at max concurrency ({self.config.max_concurrency})"
151
+ if self.config.spend_cap_usd is not None and self.spend_usd >= self.config.spend_cap_usd:
152
+ return "spend cap reached"
153
+ if not self.breaker.allow(now):
154
+ return "circuit open"
155
+ if not self.mirror.can_serve(estimate, now):
156
+ return "insufficient headroom"
157
+ return None
158
+
159
+ # ---- outcomes ------------------------------------------------------------
160
+
161
+ def apply(self, classification: Classification, now: float, model: str = "") -> None:
162
+ action = classification.account_action
163
+ if action is AccountAction.NONE:
164
+ return
165
+ if action is AccountAction.ERROR_TICK:
166
+ self.breaker.record_failure(now)
167
+ elif action is AccountAction.COOLDOWN:
168
+ seconds = classification.cooldown_seconds
169
+ if seconds is None:
170
+ seconds = self.mirror.next_reset(now) or 60.0
171
+ self.cooling_until = max(self.cooling_until, now + float(seconds))
172
+ elif action is AccountAction.SHORT_COOLDOWN:
173
+ self.cooling_until = max(self.cooling_until, now + 5.0)
174
+ self.breaker.record_failure(now)
175
+ elif action is AccountAction.DISABLE:
176
+ self.disabled_reason = classification.kind
177
+ elif action is AccountAction.MARK_MODEL_UNSUPPORTED and model:
178
+ self.unsupported_models.add(model)
179
+ self.outcomes.append(False)
180
+ self.failures_total += 1
181
+ self.error_kinds[classification.kind] = self.error_kinds.get(classification.kind, 0) + 1
182
+
183
+ def record_success(
184
+ self,
185
+ latency: float,
186
+ usage: Usage,
187
+ price: Optional[ModelPrice],
188
+ now: float,
189
+ cost_multiplier: float = 1.0,
190
+ ) -> Optional[float]:
191
+ self.breaker.record_success()
192
+ self.outcomes.append(True)
193
+ self.latencies.append(latency)
194
+ self.requests_total += 1
195
+ self.input_tokens_total += usage.input_tokens
196
+ self.output_tokens_total += usage.output_tokens
197
+ self.cache_read_total += usage.cache_read_tokens
198
+ self.cache_creation_total += usage.cache_creation_tokens
199
+ cost = usage.cost(price)
200
+ if cost is not None:
201
+ cost *= cost_multiplier
202
+ if cost is not None:
203
+ self.spend_usd += cost
204
+ return cost
205
+
206
+ # ---- derived stats -------------------------------------------------------
207
+
208
+ def error_rate(self) -> float:
209
+ if not self.outcomes:
210
+ return 0.0
211
+ return sum(1 for ok in self.outcomes if not ok) / float(len(self.outcomes))
212
+
213
+ def p95_latency(self) -> Optional[float]:
214
+ if not self.latencies:
215
+ return None
216
+ ordered = sorted(self.latencies)
217
+ index = min(len(ordered) - 1, int(round(0.95 * (len(ordered) - 1))))
218
+ return ordered[index]
219
+
220
+ def cache_hit_rate(self) -> Optional[float]:
221
+ billed = self.input_tokens_total + self.cache_read_total + self.cache_creation_total
222
+ if billed <= 0:
223
+ return None
224
+ return self.cache_read_total / float(billed)
225
+
226
+ def load(self) -> float:
227
+ if self.config.max_concurrency <= 0:
228
+ return 1.0
229
+ return min(1.0, self.inflight / float(self.config.max_concurrency))
230
+
231
+ def snapshot(self, now: float) -> Dict[str, Any]:
232
+ cooling = self.cooling_for(now)
233
+ p95 = self.p95_latency()
234
+ cache_rate = self.cache_hit_rate()
235
+ return {
236
+ "id": self.id,
237
+ "type": self.config.type,
238
+ "state": self.state(now).value,
239
+ "disabled_reason": self.disabled_reason,
240
+ "cooling_for": None if cooling is None else round(cooling, 1),
241
+ "inflight": self.inflight,
242
+ "priority": self.config.priority,
243
+ "cost_tier": self.config.cost_tier,
244
+ "limits": self.mirror.snapshot(now),
245
+ "error_rate": round(self.error_rate(), 4),
246
+ "p95_latency": None if p95 is None else round(p95, 3),
247
+ "cache_hit_rate": None if cache_rate is None else round(cache_rate, 4),
248
+ "requests_total": self.requests_total,
249
+ "failures_total": self.failures_total,
250
+ "error_kinds": dict(self.error_kinds),
251
+ "spend_usd": round(self.spend_usd, 4),
252
+ "tokens": {
253
+ "input": self.input_tokens_total,
254
+ "output": self.output_tokens_total,
255
+ "cache_read": self.cache_read_total,
256
+ "cache_write": self.cache_creation_total,
257
+ },
258
+ }
@@ -0,0 +1,135 @@
1
+ """The spill lane: one Messages request run through the Message Batches API.
2
+
3
+ Batches are cheaper but asynchronous. The gateway holds the client's connection
4
+ while it polls, bounded by that request's own wait budget — never longer. A batch
5
+ that outlives the budget is cancelled and the id is handed back, so nothing is
6
+ silently abandoned upstream.
7
+
8
+ Only `batch` priority spills, and never a streaming request: batches do not stream.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import asyncio
14
+ import json
15
+ import time
16
+ from dataclasses import dataclass
17
+ from typing import Any, Dict, Mapping, Optional
18
+
19
+ import httpx
20
+
21
+ from ..providers.base import Upstream
22
+
23
+ CUSTOM_ID = "tokenbiryani-spill"
24
+
25
+ #: Terminal processing_status values from the Batches API.
26
+ ENDED = "ended"
27
+ TERMINAL = {"ended", "canceled", "cancelled", "expired", "errored"}
28
+
29
+
30
+ class BatchUnavailable(Exception):
31
+ """The spill lane could not serve this request; fall back to queueing."""
32
+
33
+
34
+ class BatchTimeout(Exception):
35
+ """The batch outlived the request's wait budget. Carries the id for follow-up."""
36
+
37
+ def __init__(self, batch_id: Optional[str]) -> None:
38
+ super().__init__("batch did not finish within the request's wait budget")
39
+ self.batch_id = batch_id
40
+
41
+
42
+ @dataclass
43
+ class BatchOutcome:
44
+ message: Dict[str, Any]
45
+ batch_id: str
46
+ polls: int
47
+ waited: float
48
+
49
+
50
+ def _first_result(payload: bytes) -> Optional[Dict[str, Any]]:
51
+ """Batch results are JSONL. We submit one request, so we read one line."""
52
+ for line in payload.decode("utf-8", "replace").splitlines():
53
+ line = line.strip()
54
+ if not line:
55
+ continue
56
+ try:
57
+ parsed = json.loads(line)
58
+ except ValueError:
59
+ continue
60
+ if isinstance(parsed, dict):
61
+ return parsed
62
+ return None
63
+
64
+
65
+ async def run_single(
66
+ client: httpx.AsyncClient,
67
+ upstream: Upstream,
68
+ params: Mapping[str, Any],
69
+ headers: Mapping[str, str],
70
+ poll_interval: float,
71
+ deadline: float,
72
+ ) -> BatchOutcome:
73
+ """Submit one request as a batch and wait for it, within `deadline`."""
74
+ started = time.time()
75
+ payload = {k: v for k, v in params.items() if k != "stream"}
76
+
77
+ submitted = await upstream.submit_batch(
78
+ client, [{"custom_id": CUSTOM_ID, "params": payload}], headers
79
+ )
80
+ if submitted.status >= 300 or not submitted.body:
81
+ raise BatchUnavailable(
82
+ f"batch submission returned {submitted.status}"
83
+ )
84
+ batch_id = str(submitted.body.get("id") or "")
85
+ if not batch_id:
86
+ raise BatchUnavailable("batch submission returned no id")
87
+
88
+ polls = 0
89
+ status = str(submitted.body.get("processing_status") or "in_progress")
90
+ while status not in TERMINAL:
91
+ if time.time() >= deadline:
92
+ await _cancel(client, upstream, batch_id, headers)
93
+ raise BatchTimeout(batch_id)
94
+ await asyncio.sleep(min(poll_interval, max(0.0, deadline - time.time())))
95
+ polls += 1
96
+ polled = await upstream.poll_batch(client, batch_id, headers)
97
+ if polled.status >= 300 or not polled.body:
98
+ raise BatchUnavailable(f"batch poll returned {polled.status}")
99
+ status = str(polled.body.get("processing_status") or "in_progress")
100
+
101
+ if status != ENDED:
102
+ raise BatchUnavailable(f"batch finished as {status}")
103
+
104
+ fetched = await upstream.fetch_batch_results(client, batch_id, headers)
105
+ if fetched.status >= 300:
106
+ raise BatchUnavailable(f"batch results returned {fetched.status}")
107
+
108
+ entry = _first_result(fetched.raw)
109
+ if entry is None:
110
+ raise BatchUnavailable("batch results were empty")
111
+
112
+ result = entry.get("result") or {}
113
+ if result.get("type") != "succeeded":
114
+ raise BatchUnavailable(
115
+ "batch request {}".format(result.get("type") or "did not succeed")
116
+ )
117
+ message = result.get("message")
118
+ if not isinstance(message, dict):
119
+ raise BatchUnavailable("batch result carried no message")
120
+
121
+ return BatchOutcome(
122
+ message=message, batch_id=batch_id, polls=polls, waited=time.time() - started
123
+ )
124
+
125
+
126
+ async def _cancel(
127
+ client: httpx.AsyncClient,
128
+ upstream: Upstream,
129
+ batch_id: str,
130
+ headers: Mapping[str, str],
131
+ ) -> None:
132
+ try:
133
+ await upstream.cancel_batch(client, batch_id, headers)
134
+ except Exception: # noqa: BLE001 - a failed cancel must not mask the timeout
135
+ pass
@@ -0,0 +1,53 @@
1
+ """Per-account circuit breaker: closed -> open -> half-open -> closed."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+ from enum import Enum
7
+ from typing import Optional
8
+
9
+
10
+ class BreakerState(str, Enum):
11
+ CLOSED = "closed"
12
+ OPEN = "open"
13
+ HALF_OPEN = "half_open"
14
+
15
+
16
+ @dataclass
17
+ class CircuitBreaker:
18
+ failure_threshold: int = 5
19
+ cooldown_seconds: float = 30.0
20
+
21
+ consecutive_failures: int = 0
22
+ opened_at: Optional[float] = None
23
+ _probe_in_flight: bool = False
24
+
25
+ def state(self, now: float) -> BreakerState:
26
+ if self.opened_at is None:
27
+ return BreakerState.CLOSED
28
+ if now - self.opened_at >= self.cooldown_seconds:
29
+ return BreakerState.HALF_OPEN
30
+ return BreakerState.OPEN
31
+
32
+ def allow(self, now: float) -> bool:
33
+ """Half-open lets exactly one probe through at a time."""
34
+ state = self.state(now)
35
+ if state is BreakerState.CLOSED:
36
+ return True
37
+ if state is BreakerState.OPEN:
38
+ return False
39
+ if self._probe_in_flight:
40
+ return False
41
+ self._probe_in_flight = True
42
+ return True
43
+
44
+ def record_success(self) -> None:
45
+ self.consecutive_failures = 0
46
+ self.opened_at = None
47
+ self._probe_in_flight = False
48
+
49
+ def record_failure(self, now: float) -> None:
50
+ self.consecutive_failures += 1
51
+ self._probe_in_flight = False
52
+ if self.consecutive_failures >= self.failure_threshold:
53
+ self.opened_at = now
@@ -0,0 +1,239 @@
1
+ """Whether the caller ever asked for caching, and what it costs when they did not.
2
+
3
+ The gateway already measures the *outcome* — `cache_read_input_tokens` comes back on
4
+ every response and the console draws a hit rate from it. What it could not do is tell
5
+ the two reasons for a 0% hit rate apart:
6
+
7
+ the client never marked a breakpoint -> caching never engages at all, and no
8
+ routing change can fix it
9
+ the breakpoint is there, affinity broke -> a routing problem, already diagnosed
10
+
11
+ Both render as "cache hit 0%", and an operator seeing that reaches for the strategy
12
+ knob, which cannot help in the first case. Anthropic's cache only engages where the
13
+ request carries `cache_control`, so the presence of that marker is the fact that
14
+ separates them, and it is knowable before the request is even sent.
15
+
16
+ Nothing here keeps prompt content. It reads the body in memory, records booleans and
17
+ token counts, and drops it — the promise that bodies are never logged is not
18
+ weakened by knowing whether one had a breakpoint in it.
19
+ """
20
+
21
+ from __future__ import annotations
22
+
23
+ import copy
24
+ from collections import defaultdict
25
+ from typing import Any, Dict, Iterable, List, Mapping, Optional, Tuple
26
+
27
+ #: Anthropic will not cache a prefix shorter than this, so a breakpoint on a small
28
+ #: request buys nothing and advising one would be noise. The real minimum is
29
+ #: model-dependent (larger for Haiku); this is the smaller of them, because
30
+ #: over-reporting a missing breakpoint is worse than staying quiet.
31
+ MIN_CACHEABLE_TOKENS = 1024
32
+
33
+ #: Rough characters per token, matching `limits.CHARS_PER_TOKEN`.
34
+ CHARS_PER_TOKEN = 3.5
35
+
36
+ #: Below this many requests a group is not worth an opinion.
37
+ MIN_REQUESTS_FOR_ADVICE = 20
38
+
39
+ #: A group with fewer than this fraction of requests marked is treated as unmarked.
40
+ BREAKPOINT_PRESENT = 0.5
41
+
42
+ CACHE_CONTROL = "cache_control"
43
+ EPHEMERAL = {"type": "ephemeral"}
44
+
45
+
46
+ def _blocks(value: Any) -> Iterable[Mapping[str, Any]]:
47
+ """Yield the mapping blocks in a `system`, `tools` or message `content` field."""
48
+ if isinstance(value, Mapping):
49
+ yield value
50
+ elif isinstance(value, list):
51
+ for item in value:
52
+ if isinstance(item, Mapping):
53
+ yield item
54
+
55
+
56
+ def has_cache_breakpoint(body: Mapping[str, Any]) -> bool:
57
+ """True when anything in the request is marked for caching.
58
+
59
+ Checks the three places a breakpoint may legally sit: a system block, a tool
60
+ definition, and a content block inside a message.
61
+ """
62
+ for field in ("system", "tools"):
63
+ for block in _blocks(body.get(field)):
64
+ if block.get(CACHE_CONTROL):
65
+ return True
66
+ messages = body.get("messages")
67
+ if isinstance(messages, list):
68
+ for message in messages:
69
+ if not isinstance(message, Mapping):
70
+ continue
71
+ for block in _blocks(message.get("content")):
72
+ if block.get(CACHE_CONTROL):
73
+ return True
74
+ return False
75
+
76
+
77
+ def cacheable_prefix_tokens(body: Mapping[str, Any]) -> int:
78
+ """Roughly how much of this request is a stable, cacheable head.
79
+
80
+ The system prompt and the tool definitions: the part an agent resends unchanged
81
+ on every turn, and the part worth a breakpoint.
82
+ """
83
+ characters = 0
84
+ for field in ("system", "tools"):
85
+ value = body.get(field)
86
+ if isinstance(value, str):
87
+ characters += len(value)
88
+ continue
89
+ for block in _blocks(value):
90
+ for key in ("text", "description", "name"):
91
+ piece = block.get(key)
92
+ if isinstance(piece, str):
93
+ characters += len(piece)
94
+ schema = block.get("input_schema")
95
+ if isinstance(schema, Mapping):
96
+ characters += len(str(schema))
97
+ return int(characters / CHARS_PER_TOKEN)
98
+
99
+
100
+ def insert_cache_breakpoint(body: Mapping[str, Any]) -> Tuple[Dict[str, Any], bool]:
101
+ """Return a copy of the body with one breakpoint at the end of the stable head.
102
+
103
+ Off unless `cache.auto_breakpoint` is set, because this is the gateway editing a
104
+ caller's request — the thing it otherwise refuses to do. It is the third
105
+ documented exception, after the two fields Bedrock and Vertex need.
106
+
107
+ The breakpoint goes on the last tool if there are tools, otherwise on the last
108
+ system block; those are the end of the prefix an agent resends unchanged. A
109
+ string `system` is promoted to a one-block list, which is the same prompt in the
110
+ other legal spelling and the only way to attach the marker at all.
111
+ """
112
+ if has_cache_breakpoint(body):
113
+ return dict(body), False
114
+ if cacheable_prefix_tokens(body) < MIN_CACHEABLE_TOKENS:
115
+ return dict(body), False
116
+
117
+ updated = copy.deepcopy(dict(body))
118
+
119
+ tools = updated.get("tools")
120
+ if isinstance(tools, list) and tools and isinstance(tools[-1], dict):
121
+ tools[-1][CACHE_CONTROL] = dict(EPHEMERAL)
122
+ return updated, True
123
+
124
+ system = updated.get("system")
125
+ if isinstance(system, str) and system:
126
+ updated["system"] = [
127
+ {"type": "text", "text": system, CACHE_CONTROL: dict(EPHEMERAL)}
128
+ ]
129
+ return updated, True
130
+ if isinstance(system, list) and system and isinstance(system[-1], dict):
131
+ system[-1][CACHE_CONTROL] = dict(EPHEMERAL)
132
+ return updated, True
133
+
134
+ return updated, False
135
+
136
+
137
+ class _Group:
138
+ """Counters for one (virtual key, model) pair. No content, only arithmetic."""
139
+
140
+ __slots__ = (
141
+ "requests", "marked", "prefix_tokens", "input_tokens",
142
+ "cache_read_tokens", "cache_creation_tokens",
143
+ )
144
+
145
+ def __init__(self) -> None:
146
+ self.requests = 0
147
+ self.marked = 0
148
+ self.prefix_tokens = 0
149
+ self.input_tokens = 0
150
+ self.cache_read_tokens = 0
151
+ self.cache_creation_tokens = 0
152
+
153
+
154
+ class CacheAdvisor:
155
+ """Aggregates breakpoint presence against realised cache hits, per key and model."""
156
+
157
+ def __init__(
158
+ self,
159
+ min_requests: int = MIN_REQUESTS_FOR_ADVICE,
160
+ min_cacheable_tokens: int = MIN_CACHEABLE_TOKENS,
161
+ ) -> None:
162
+ self.min_requests = max(1, int(min_requests))
163
+ self.min_cacheable_tokens = max(1, int(min_cacheable_tokens))
164
+ self._groups: Dict[Tuple[str, str], _Group] = defaultdict(_Group)
165
+
166
+ def reconfigure(self, min_requests: int) -> None:
167
+ """Apply new settings without discarding the counters already gathered."""
168
+ self.min_requests = max(1, int(min_requests))
169
+
170
+ def observe(
171
+ self,
172
+ key_name: str,
173
+ model: str,
174
+ marked: Optional[bool],
175
+ prefix_tokens: int,
176
+ input_tokens: int,
177
+ cache_read_tokens: int,
178
+ cache_creation_tokens: int,
179
+ ) -> None:
180
+ if marked is None:
181
+ return
182
+ group = self._groups[(key_name or "", model or "")]
183
+ group.requests += 1
184
+ group.marked += 1 if marked else 0
185
+ group.prefix_tokens += max(0, int(prefix_tokens))
186
+ group.input_tokens += max(0, int(input_tokens))
187
+ group.cache_read_tokens += max(0, int(cache_read_tokens))
188
+ group.cache_creation_tokens += max(0, int(cache_creation_tokens))
189
+
190
+ def advice(self) -> Dict[str, Any]:
191
+ findings: List[Dict[str, Any]] = []
192
+ for (key_name, model), group in sorted(self._groups.items()):
193
+ if group.requests < self.min_requests:
194
+ continue
195
+ billed = (
196
+ group.input_tokens + group.cache_read_tokens + group.cache_creation_tokens
197
+ )
198
+ marked_rate = group.marked / float(group.requests)
199
+ mean_prefix = group.prefix_tokens // group.requests
200
+ hit_rate = (group.cache_read_tokens / billed) if billed else 0.0
201
+
202
+ findings.append({
203
+ "key": key_name,
204
+ "model": model,
205
+ "requests": group.requests,
206
+ "breakpoint_rate": round(marked_rate, 4),
207
+ "mean_prefix_tokens": mean_prefix,
208
+ "cache_hit_rate": round(hit_rate, 4),
209
+ "uncached_input_tokens": group.input_tokens,
210
+ "verdict": self._verdict(marked_rate, mean_prefix, hit_rate),
211
+ })
212
+ return {
213
+ "min_requests": self.min_requests,
214
+ "min_cacheable_tokens": self.min_cacheable_tokens,
215
+ "findings": findings,
216
+ }
217
+
218
+ def _verdict(self, marked_rate: float, mean_prefix: int, hit_rate: float) -> str:
219
+ """One sentence an operator can act on, or one saying there is nothing to do."""
220
+ if marked_rate < BREAKPOINT_PRESENT:
221
+ if mean_prefix < self.min_cacheable_tokens:
222
+ return (
223
+ "no cache_control breakpoint, and the stable prefix is under "
224
+ f"{self.min_cacheable_tokens} tokens — too small to cache, so "
225
+ "there is nothing to fix here"
226
+ )
227
+ return (
228
+ "no cache_control breakpoint on a stable prefix of about "
229
+ f"{mean_prefix} tokens: caching never engages for this traffic, and "
230
+ "no routing strategy can recover it. The client has to mark the "
231
+ "prefix, or set cache.auto_breakpoint"
232
+ )
233
+ if hit_rate < 0.2:
234
+ return (
235
+ "breakpoints are being sent but the hit rate is low — this is a "
236
+ "routing problem, not a client one. Check cache breaks and whether "
237
+ "the owner keeps going cooling"
238
+ )
239
+ return "caching is engaged and working"