Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/spend_tracking/baseline_accounting.py: 37%
129 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Pure, chronological cache accounting for the recorded baseline comparison.
3Observation collection, pricing and durable publication belong to their existing
4owners. Replaying these values in event order is independent of callback order.
5"""
7from __future__ import annotations
9from collections.abc import Sequence
10from dataclasses import dataclass
11from itertools import groupby
12from math import isfinite
13from types import MappingProxyType
14from typing import Final, Literal
16from pydantic import BaseModel, ConfigDict, Field
18from litellm.llms.anthropic.prompt_cache_prediction import CountedBreakpoint, CountedPromptCachePlan
19from litellm.types.utils import CacheCreationTokenDetails, PromptTokensDetailsWrapper, Usage
21MAX_CACHE_TTL: Final = 3600
22MAX_CACHE_ENTRIES: Final = 1024
25class BaselineObservation(BaseModel):
26 model_config = ConfigDict(extra="forbid", frozen=True, strict=True)
28 version: Literal[3] = 3
29 request_id: str = Field(min_length=1)
30 started_at: float = Field(allow_inf_nan=False, ge=0)
31 available_at: float = Field(allow_inf_nan=False, ge=0)
32 outcome: Literal["complete", "uncertain", "response_cache"]
33 baseline_equivalent: bool
34 usage: Usage | None = None
35 plan: CountedPromptCachePlan | None = None
36 minimum_cache_tokens: int = Field(default=0, ge=0)
37 reason: str | None = None
40@dataclass(frozen=True, slots=True)
41class BaselineEstimate:
42 request_id: str
43 reason: str
44 provenance: Literal["observed_identical", "modeled"] | None = None
45 usage: Usage | None = None
48@dataclass(frozen=True, slots=True)
49class CacheEntry:
50 fingerprint: str
51 content_fingerprint: str
52 tokens: int
53 ttl_seconds: int
54 available_at: float
55 expires_at: float
56 uncertain: bool = False
59@dataclass(frozen=True, slots=True)
60class BaselineHistory:
61 first_at: float | None = None
62 last_at: float | None = None
63 equivalent: bool = True
64 uncertain_before: float = 0.0
65 entries: tuple[CacheEntry, ...] = ()
66 blocked_until: float = 0.0
69def _complete_usage(usage: Usage | None) -> bool:
70 if usage is None or usage.prompt_tokens < 0 or usage.completion_tokens < 0:
71 return False
72 details: Final = usage.prompt_tokens_details
73 if details is None:
74 return False
75 values: Final = (details.text_tokens, details.cached_tokens, details.cache_creation_tokens)
76 if any(value is None or value < 0 for value in values):
77 return False
78 split: Final = details.cache_creation_token_details
79 writes: Final = details.cache_creation_tokens or 0
80 return (
81 usage.total_tokens == usage.prompt_tokens + usage.completion_tokens
82 and sum(value or 0 for value in values) == usage.prompt_tokens
83 and (
84 writes == 0
85 or (
86 split is not None
87 and split.ephemeral_5m_input_tokens is not None
88 and split.ephemeral_1h_input_tokens is not None
89 and min(split.ephemeral_5m_input_tokens, split.ephemeral_1h_input_tokens) >= 0
90 and split.ephemeral_5m_input_tokens + split.ephemeral_1h_input_tokens == writes
91 )
92 )
93 )
96def _valid_plan(plan: CountedPromptCachePlan | None) -> bool:
97 if plan is None or plan.total_tokens < 0 or len(plan.breakpoints) > 4:
98 return False
99 return all(
100 marker.fingerprint
101 and marker.content_fingerprint
102 and marker.fingerprint in marker.lookback_fingerprints
103 and marker.content_fingerprint in marker.lookback_content_fingerprints
104 and marker.ttl_seconds in (300, 3600)
105 and 0 <= marker.prefix_tokens <= plan.total_tokens
106 for marker in plan.breakpoints
107 ) and all(
108 left.prefix_tokens <= right.prefix_tokens and left.ttl_seconds >= right.ttl_seconds
109 for left, right in zip(plan.breakpoints, plan.breakpoints[1:])
110 )
113def _markers(observation: BaselineObservation) -> tuple[CountedBreakpoint, ...]:
114 return (
115 tuple(
116 marker
117 for marker in observation.plan.breakpoints
118 if marker.prefix_tokens >= observation.minimum_cache_tokens
119 )
120 if observation.plan is not None
121 else ()
122 )
125def _matches(entry: CacheEntry, markers: tuple[CountedBreakpoint, ...], started: float) -> bool:
126 return entry.available_at <= started < entry.expires_at and any(
127 entry.fingerprint in marker.lookback_fingerprints
128 and entry.tokens <= marker.prefix_tokens
129 and entry.ttl_seconds == marker.ttl_seconds
130 for marker in markers
131 )
134def _ambiguous(entry: CacheEntry, markers: tuple[CountedBreakpoint, ...], started: float) -> bool:
135 return entry.available_at <= started < entry.expires_at and any(
136 entry.content_fingerprint in marker.lookback_content_fingerprints
137 and (entry.uncertain or entry.ttl_seconds != marker.ttl_seconds)
138 for marker in markers
139 )
142def _usage_with_cache(usage: Usage, total: int, read: int, write_5m: int, write_1h: int) -> Usage:
143 writes: Final = write_5m + write_1h
144 original_details: Final = usage.prompt_tokens_details or PromptTokensDetailsWrapper()
145 details: Final = original_details.model_copy(
146 deep=True,
147 update=MappingProxyType(
148 {
149 "text_tokens": total - read - writes,
150 "cached_tokens": read,
151 "cache_creation_tokens": writes,
152 "cache_write_tokens": writes,
153 "cache_creation_token_details": CacheCreationTokenDetails(
154 ephemeral_5m_input_tokens=write_5m,
155 ephemeral_1h_input_tokens=write_1h,
156 ),
157 }
158 ),
159 )
160 return Usage.model_validate(
161 { # mutable-ok: Usage only runs its normalizing constructor for a plain dictionary
162 **usage.model_dump(),
163 "prompt_tokens": total,
164 "total_tokens": total + usage.completion_tokens,
165 "prompt_tokens_details": details,
166 "cache_read_input_tokens": read,
167 "cache_creation_input_tokens": writes,
168 },
169 )
172def _estimate(history: BaselineHistory, observation: BaselineObservation, equivalent: bool) -> BaselineEstimate:
173 if observation.outcome != "complete" or not _complete_usage(observation.usage):
174 return BaselineEstimate(observation.request_id, observation.reason or observation.outcome)
175 usage: Final = observation.usage
176 if usage is None:
177 return BaselineEstimate(observation.request_id, "missing_usage")
178 if equivalent and observation.baseline_equivalent:
179 return BaselineEstimate(
180 observation.request_id, "identical_baseline_path", "observed_identical", usage.model_copy(deep=True)
181 )
182 if observation.started_at < history.blocked_until:
183 return BaselineEstimate(observation.request_id, "concurrent_uncertainty")
184 plan: Final = observation.plan
185 if not _valid_plan(plan) or plan is None:
186 return BaselineEstimate(observation.request_id, observation.reason or "unsupported_cache_plan")
187 markers: Final = _markers(observation)
188 if any(_ambiguous(entry, markers, observation.started_at) for entry in history.entries):
189 return BaselineEstimate(observation.request_id, "cache_ttl_changed")
190 read: Final = max(
191 (
192 entry.tokens
193 for entry in history.entries
194 if not entry.uncertain and _matches(entry, markers, observation.started_at)
195 ),
196 default=0,
197 )
198 end: Final = markers[-1].prefix_tokens if markers else 0
199 if read < end and observation.started_at < history.uncertain_before + max(marker.ttl_seconds for marker in markers):
200 return BaselineEstimate(observation.request_id, "history_unavailable")
201 one_hour: Final = max(
202 (marker.prefix_tokens for marker in markers if marker.ttl_seconds == 3600 and marker.prefix_tokens > read),
203 default=read,
204 )
205 expired: Final = any(
206 entry.expires_at <= observation.started_at
207 and any(entry.fingerprint in marker.lookback_fingerprints for marker in markers)
208 for entry in history.entries
209 )
210 reason: Final = (
211 "cache_prefix_available"
212 if read
213 else "cache_prefix_expired"
214 if expired
215 else "cache_prefix_cold"
216 if markers
217 else "below_cache_minimum"
218 if plan.breakpoints
219 else "no_cache_breakpoints"
220 )
221 return BaselineEstimate(
222 observation.request_id,
223 reason,
224 "modeled",
225 _usage_with_cache(usage, plan.total_tokens, read, end - one_hour, one_hour - read),
226 )
229def _writes(history: BaselineHistory, observation: BaselineObservation) -> tuple[CacheEntry, ...]:
230 if (
231 observation.outcome != "complete"
232 or observation.started_at < history.blocked_until
233 or not _complete_usage(observation.usage)
234 or not _valid_plan(observation.plan)
235 ):
236 return ()
237 markers: Final = _markers(observation)
238 ambiguous: Final = tuple(entry for entry in history.entries if _ambiguous(entry, markers, observation.started_at))
239 hit: Final = (
240 max(
241 (
242 entry
243 for entry in history.entries
244 if not entry.uncertain and _matches(entry, markers, observation.started_at)
245 ),
246 key=lambda entry: entry.tokens,
247 default=None,
248 )
249 if not ambiguous
250 else None
251 )
252 refresh: Final = (
253 (
254 CacheEntry(
255 hit.fingerprint,
256 hit.content_fingerprint,
257 hit.tokens,
258 hit.ttl_seconds,
259 observation.available_at,
260 observation.started_at + hit.ttl_seconds,
261 ),
262 )
263 if hit is not None and all(marker.fingerprint != hit.fingerprint for marker in markers)
264 else ()
265 )
266 return (
267 *refresh,
268 *(
269 CacheEntry(
270 marker.fingerprint,
271 marker.content_fingerprint,
272 marker.prefix_tokens,
273 marker.ttl_seconds,
274 observation.available_at,
275 observation.started_at + max((marker.ttl_seconds, *(entry.ttl_seconds for entry in ambiguous))),
276 uncertain=bool(ambiguous),
277 )
278 for marker in markers
279 ),
280 )
283def _entry_key(entry: CacheEntry) -> tuple[str, str, int, int, bool]:
284 return entry.fingerprint, entry.content_fingerprint, entry.tokens, entry.ttl_seconds, entry.uncertain
287def _compact_entries(entries: tuple[CacheEntry, ...], started: float) -> tuple[CacheEntry, ...]:
288 ordered: Final = sorted((entry for entry in entries if entry.expires_at >= started - MAX_CACHE_TTL), key=_entry_key)
289 return tuple(
290 retained
291 for _, values in groupby(ordered, key=_entry_key)
292 for group in (tuple(values),)
293 for retained in (
294 max(
295 (entry for entry in group if entry.available_at <= started),
296 key=lambda entry: entry.expires_at,
297 default=None,
298 ),
299 *(entry for entry in group if entry.available_at > started),
300 )
301 if retained is not None
302 )
305def advance_baseline_history(
306 history: BaselineHistory,
307 simultaneous: Sequence[BaselineObservation],
308) -> tuple[BaselineHistory, tuple[BaselineEstimate, ...]]:
309 """Apply one request-start timestamp; ties cannot manufacture initial equality.
311 The storage owner groups and orders observations before calling this function.
312 Equal timestamps are evaluated against the same preceding cache snapshot.
313 """
314 if not simultaneous:
315 return history, ()
316 started: Final = simultaneous[0].started_at
317 valid_order: Final = (
318 isfinite(started)
319 and all(item.started_at == started and item.available_at >= started for item in simultaneous)
320 and (history.last_at is None or started > history.last_at)
321 )
322 if not valid_order:
323 return history, tuple(BaselineEstimate(item.request_id, "invalid_observation_order") for item in simultaneous)
324 first: Final = started if history.first_at is None else history.first_at
325 uncertain: Final = max(history.uncertain_before, first)
326 relevant: Final = tuple(item for item in simultaneous if item.outcome != "response_cache")
327 equivalent: Final = history.equivalent and all(item.baseline_equivalent for item in relevant)
328 before: Final = BaselineHistory(
329 first, history.last_at, equivalent, uncertain, history.entries, history.blocked_until
330 )
331 estimates: Final = tuple(_estimate(before, item, equivalent) for item in simultaneous)
332 invalidated: Final = any(
333 item.outcome != "complete" or not _complete_usage(item.usage) or not _valid_plan(item.plan) for item in relevant
334 )
335 blocked: Final = max((history.blocked_until, *(item.available_at for item in relevant if invalidated)))
336 entries: Final = _compact_entries(
337 () if invalidated else (*history.entries, *(entry for item in relevant for entry in _writes(before, item))),
338 started,
339 )
340 overflow: Final = len(entries) > MAX_CACHE_ENTRIES
341 return BaselineHistory(
342 first_at=first,
343 last_at=started,
344 equivalent=equivalent,
345 uncertain_before=max(started, blocked) if invalidated or overflow else uncertain,
346 entries=() if overflow else entries,
347 blocked_until=blocked,
348 ), estimates