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

1"""Pure, chronological cache accounting for the recorded baseline comparison. 

2 

3Observation collection, pricing and durable publication belong to their existing 

4owners. Replaying these values in event order is independent of callback order. 

5""" 

6 

7from __future__ import annotations 

8 

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 

15 

16from pydantic import BaseModel, ConfigDict, Field 

17 

18from litellm.llms.anthropic.prompt_cache_prediction import CountedBreakpoint, CountedPromptCachePlan 

19from litellm.types.utils import CacheCreationTokenDetails, PromptTokensDetailsWrapper, Usage 

20 

21MAX_CACHE_TTL: Final = 3600 

22MAX_CACHE_ENTRIES: Final = 1024 

23 

24 

25class BaselineObservation(BaseModel): 

26 model_config = ConfigDict(extra="forbid", frozen=True, strict=True) 

27 

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 

38 

39 

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 

46 

47 

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 

57 

58 

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 

67 

68 

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 ) 

94 

95 

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 ) 

111 

112 

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 ) 

123 

124 

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 ) 

132 

133 

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 ) 

140 

141 

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 ) 

170 

171 

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 ) 

227 

228 

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 ) 

281 

282 

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 

285 

286 

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 ) 

303 

304 

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. 

310 

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