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.
- tokenbiryani/__init__.py +3 -0
- tokenbiryani/api/__init__.py +0 -0
- tokenbiryani/api/app.py +583 -0
- tokenbiryani/api/asgi.py +32 -0
- tokenbiryani/cli.py +1045 -0
- tokenbiryani/config.py +532 -0
- tokenbiryani/core/__init__.py +0 -0
- tokenbiryani/core/account.py +258 -0
- tokenbiryani/core/batch.py +135 -0
- tokenbiryani/core/breaker.py +53 -0
- tokenbiryani/core/cacheadvice.py +239 -0
- tokenbiryani/core/diagnostics.py +131 -0
- tokenbiryani/core/estimator.py +180 -0
- tokenbiryani/core/gateway.py +2395 -0
- tokenbiryani/core/handoff.py +87 -0
- tokenbiryani/core/keys.py +199 -0
- tokenbiryani/core/limits.py +440 -0
- tokenbiryani/core/oauth.py +222 -0
- tokenbiryani/core/pacing.py +320 -0
- tokenbiryani/core/queue.py +132 -0
- tokenbiryani/core/router.py +323 -0
- tokenbiryani/core/secrets.py +114 -0
- tokenbiryani/core/session.py +117 -0
- tokenbiryani/dashboard/__init__.py +56 -0
- tokenbiryani/dashboard/console.css +610 -0
- tokenbiryani/dashboard/console.html +3250 -0
- tokenbiryani/observability/__init__.py +0 -0
- tokenbiryani/observability/events.py +171 -0
- tokenbiryani/observability/usage.py +226 -0
- tokenbiryani/prices.yaml +77 -0
- tokenbiryani/providers/__init__.py +0 -0
- tokenbiryani/providers/anthropic_api.py +118 -0
- tokenbiryani/providers/base.py +173 -0
- tokenbiryani/providers/bedrock.py +182 -0
- tokenbiryani/providers/oauth.py +165 -0
- tokenbiryani/providers/oauth_credentials.py +293 -0
- tokenbiryani/providers/translate.py +35 -0
- tokenbiryani/providers/vertex.py +144 -0
- tokenbiryani/proxy/__init__.py +0 -0
- tokenbiryani/proxy/errors.py +169 -0
- tokenbiryani/proxy/sse.py +98 -0
- tokenbiryani/store/__init__.py +0 -0
- tokenbiryani/store/base.py +150 -0
- tokenbiryani/store/memory.py +149 -0
- tokenbiryani/store/redis_store.py +222 -0
- tokenbiryani/store/sqlite.py +336 -0
- tokenbiryani-0.2.0.dist-info/METADATA +697 -0
- tokenbiryani-0.2.0.dist-info/RECORD +51 -0
- tokenbiryani-0.2.0.dist-info/WHEEL +4 -0
- tokenbiryani-0.2.0.dist-info/entry_points.txt +2 -0
- 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"
|