Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/autorouter_baseline_cache.py: 40%

169 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1from __future__ import annotations 

2 

3import asyncio 

4import hashlib 

5import json 

6import time 

7from collections.abc import Callable, Mapping 

8from dataclasses import dataclass, replace 

9from datetime import datetime 

10from types import MappingProxyType 

11from typing import TYPE_CHECKING, Final 

12 

13import httpx 

14from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter 

15 

16from litellm._logging import verbose_proxy_logger 

17from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY 

18from litellm.integrations.custom_logger import CustomLogger 

19from litellm.litellm_core_utils.core_helpers import ( 

20 get_litellm_metadata_from_kwargs, # pyright: ignore[reportUnknownVariableType] # legacy metadata boundary validated below 

21) 

22from litellm.llms.anthropic.prompt_cache_prediction import ( 

23 CountedPromptCachePlan, 

24 NativePredictionTarget, 

25 TokenCounter, 

26 UnsupportedCachePlan, 

27 UnsupportedPredictionTarget, 

28 count_cache_plan, 

29 count_prompt_tokens, 

30 parse_cache_plan, 

31 resolve_baseline_prediction_target, 

32 supported_baseline_recipient, 

33 supported_prediction_headers, 

34) 

35from litellm.proxy.spend_tracking.baseline_accounting import BaselineObservation 

36from litellm.proxy.spend_tracking.savings import ( 

37 _effective_model_info, # pyright: ignore[reportPrivateUsage] # existing deployment-price owner 

38 _proxy_llm_router, # pyright: ignore[reportPrivateUsage] # existing optional proxy-router owner 

39) 

40from litellm.types.router import BaselineRouteStamp 

41from litellm.types.utils import CallTypes, ModelInfo, Usage 

42from litellm.utils import get_prompt_cache_min_tokens 

43 

44if TYPE_CHECKING: 44 ↛ 45line 44 didn't jump to line 45 because the condition on line 44 was never true

45 from litellm.litellm_core_utils.litellm_logging import Logging 

46 from litellm.proxy.utils import PrismaClient 

47 from litellm.router import Router 

48 

49_METADATA: Final = TypeAdapter(Mapping[str, object]) 

50_PRICES: Final[TypeAdapter[ModelInfo | None]] = TypeAdapter(ModelInfo | None) 

51_JSON_BODY: Final = TypeAdapter(dict[str, JsonValue]) 

52_COUNT_TIMEOUT: Final = 3.0 

53_MAX_COUNTS: Final = 4096 

54 

55 

56class CapturedBaselineObservation(BaseModel): 

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

58 

59 scope: str 

60 api_key: str 

61 session_id: str 

62 router_name: str 

63 baseline_model: str 

64 model: str 

65 prices: ModelInfo | None 

66 observation: BaselineObservation 

67 

68 

69@dataclass(frozen=True, slots=True) 

70class BaselineCacheContext: 

71 collector: AutoRouterBaselineCache 

72 capture: CapturedBaselineObservation 

73 target: NativePredictionTarget | UnsupportedPredictionTarget 

74 baseline_deployment_id: str 

75 invalidated: str | None = None 

76 

77 

78class _Metadata(BaseModel): 

79 model_config = ConfigDict(strict=True, arbitrary_types_allowed=True) 

80 route: BaselineRouteStamp = Field(alias="_autorouter_baseline_route") 

81 user_api_key_hash: str = Field(min_length=1) 

82 session_id: str | None = None 

83 

84 

85class _WireEvent(BaseModel): 

86 model_config = ConfigDict(strict=True, arbitrary_types_allowed=True) 

87 httpx_response: httpx.Response 

88 api_call_start_time: datetime 

89 completion_start_time: datetime 

90 custom_llm_provider: str 

91 stream: bool = False 

92 prompt_cache_response_complete: bool = False 

93 

94 

95class _ResponseUsage(BaseModel): 

96 model_config = ConfigDict(strict=True, from_attributes=True) 

97 usage: Usage | None = None 

98 

99 

100def _digest(value: object) -> str: 

101 return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest() 

102 

103 

104class AutoRouterBaselineCache(CustomLogger): 

105 def __init__( 

106 self, 

107 prisma_client: PrismaClient | None, 

108 router: Callable[[], Router | None] = _proxy_llm_router, 

109 token_counter: TokenCounter | None = None, 

110 clock: Callable[[], float] = time.time, 

111 ) -> None: 

112 super().__init__() # pyright: ignore[reportUnknownMemberType] # legacy callback constructor 

113 self.router: Final = router 

114 self.token_counter: Final = token_counter 

115 self.clock: Final = clock 

116 self.count_slots: Final = asyncio.Semaphore(8) 

117 self.counts: Mapping[str, tuple[int, float]] = MappingProxyType({}) 

118 

119 async def async_pre_call_deployment_hook(self, kwargs: Mapping[str, object], call_type: CallTypes | None) -> None: 

120 from litellm.litellm_core_utils.litellm_logging import Logging 

121 

122 logging_obj: Final = kwargs.get("litellm_logging_obj") 

123 if not isinstance(logging_obj, Logging) or call_type != CallTypes.anthropic_messages: 123 ↛ 125line 123 didn't jump to line 125 because the condition on line 123 was always true

124 return 

125 try: 

126 metadata: Final = _METADATA.validate_python( 

127 get_litellm_metadata_from_kwargs( 

128 {"litellm_params": kwargs} # mutable-ok: legacy metadata owner requires a dictionary 

129 ) 

130 ) 

131 if metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY): 

132 return 

133 if logging_obj.baseline_cache_context is not None: 

134 await invalidate_baseline_cache(logging_obj, "retried_request") 

135 return 

136 request: Final = _Metadata.model_validate(metadata) 

137 session: Final = kwargs.get("litellm_session_id") or request.session_id or logging_obj.litellm_session_id 

138 if not isinstance(session, str) or not session or len(session) > 256: 

139 return 

140 router: Final = self.router() 

141 deployment: Final = router.get_deployment(request.route.baseline_deployment_id) if router else None 

142 if deployment is None: 

143 return 

144 target: Final = resolve_baseline_prediction_target(deployment.litellm_params) 

145 prices: Final = _PRICES.validate_python( 

146 _effective_model_info(router, request.route.baseline_deployment_id, request.route.baseline_model) 

147 ) 

148 scope: Final = "autorouter-baseline:v3:" + _digest( 

149 ( 

150 request.user_api_key_hash, 

151 session, 

152 request.route.router_name, 

153 request.route.baseline_deployment_id, 

154 deployment.litellm_params.model_dump(mode="json"), 

155 prices, 

156 ) 

157 ) 

158 started: Final = logging_obj.start_time.timestamp() 

159 capture: Final = CapturedBaselineObservation( 

160 scope=scope, 

161 api_key=request.user_api_key_hash, 

162 session_id=session, 

163 router_name=request.route.router_name, 

164 baseline_model=request.route.baseline_model, 

165 model=target.model if isinstance(target, NativePredictionTarget) else request.route.baseline_model, 

166 prices=prices, 

167 observation=BaselineObservation( 

168 request_id=logging_obj.litellm_call_id, 

169 started_at=started, 

170 available_at=started, 

171 outcome="uncertain", 

172 baseline_equivalent=False, 

173 reason="incomplete_response", 

174 ), 

175 ) 

176 logging_obj.baseline_cache_context = BaselineCacheContext( 

177 self, capture, target, request.route.baseline_deployment_id 

178 ) 

179 except Exception: # noqa: BLE001 # optional observation cannot fail inference 

180 verbose_proxy_logger.warning("Auto-router baseline observation could not be initialized") 

181 

182 async def _count(self, target: NativePredictionTarget, body: Mapping[str, JsonValue]) -> int | None: 

183 key: Final = _digest((target.model, target.api_key, target.api_base, _JSON_BODY.validate_python(body))) 

184 now: Final = self.clock() 

185 cached: Final = self.counts.get(key) 

186 if cached is not None and cached[1] > now: 

187 return cached[0] 

188 async with self.count_slots: 

189 tokens: Final = ( 

190 await self.token_counter(target.model, target.api_key, body) 

191 if self.token_counter is not None 

192 else await count_prompt_tokens(target.model, target.api_key, body, api_base=target.api_base) 

193 ) 

194 if tokens is None or tokens < 0: 

195 return None 

196 retained: Final = tuple((k, v) for k, v in self.counts.items() if v[1] > now and k != key)[-(_MAX_COUNTS - 1) :] 

197 self.counts = MappingProxyType(dict((*retained, (key, (tokens, now + 3600))))) 

198 return tokens 

199 

200 async def plan( 

201 self, target: NativePredictionTarget, wire: httpx.Request, body: Mapping[str, JsonValue], usage: Usage | None 

202 ) -> tuple[CountedPromptCachePlan | None, str | None]: 

203 if not supported_prediction_headers(wire.headers): 

204 return None, "unsupported_request_headers" 

205 plan: Final = parse_cache_plan(body) 

206 if isinstance(plan, UnsupportedCachePlan): 

207 return None, plan.reason 

208 details: Final = usage.prompt_tokens_details if usage is not None else None 

209 if ( 

210 not plan.breakpoints 

211 and details is not None 

212 and ((details.cached_tokens or 0) + (details.cache_creation_tokens or 0)) 

213 ): 

214 return None, "implicit_cache_without_breakpoints" 

215 

216 async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: 

217 return await self._count(target, body) 

218 

219 try: 

220 counted: Final = await asyncio.wait_for( 

221 count_cache_plan(target.model, target.api_key, plan, token_counter=count), timeout=_COUNT_TIMEOUT 

222 ) 

223 return (None, counted.reason) if isinstance(counted, UnsupportedCachePlan) else (counted, None) 

224 except TimeoutError: 

225 return None, "token_count_timeout" 

226 except Exception: # noqa: BLE001 # token counting cannot fail a completed request 

227 return None, "token_count_unavailable" 

228 

229 

230async def invalidate_baseline_cache(logging_obj: Logging, reason: str, *, completed: bool = False) -> None: 

231 context: Final = logging_obj.baseline_cache_context 

232 if context is not None: 

233 logging_obj.baseline_cache_context = replace(context, invalidated=reason) 

234 logging_obj.baseline_observation = context.capture.model_copy( 

235 update=MappingProxyType( 

236 { 

237 "observation": context.capture.observation.model_copy( 

238 update=MappingProxyType( 

239 { 

240 "available_at": max(context.capture.observation.started_at, context.collector.clock()), 

241 "reason": reason, 

242 } 

243 ) 

244 ), 

245 } 

246 ) 

247 ) 

248 

249 

250async def finalize_baseline_cache(logging_obj: Logging, response_obj: object) -> None: 

251 context: Final = logging_obj.baseline_cache_context 

252 if context is None: 

253 return 

254 try: 

255 capture: Final = await _capture(context, logging_obj, response_obj) 

256 if logging_obj.baseline_cache_context is context: 

257 logging_obj.baseline_observation = capture # rebind-ok: attach only to the captured request owner 

258 except Exception: # noqa: BLE001 # observation failures must preserve inference and billing 

259 await invalidate_baseline_cache(logging_obj, "observation_unavailable") 

260 

261 

262async def _capture( 

263 context: BaselineCacheContext, logging_obj: Logging, response_obj: object 

264) -> CapturedBaselineObservation: 

265 original: Final = context.capture.observation 

266 details: Final = _METADATA.validate_python(logging_obj.model_call_details) 

267 if details.get("cache_hit") is True: 

268 return context.capture.model_copy( 

269 update=MappingProxyType( 

270 { 

271 "observation": original.model_copy( 

272 update=MappingProxyType({"outcome": "response_cache", "reason": "response_cache_hit"}) 

273 ) 

274 } 

275 ) 

276 ) 

277 event: Final = _WireEvent.model_validate(details) 

278 wire: Final = event.httpx_response.request 

279 usage: Final = _ResponseUsage.model_validate(response_obj).usage 

280 complete: Final = ( 

281 event.custom_llm_provider == "anthropic" 

282 and event.httpx_response.status_code == 200 

283 and (not event.stream or event.prompt_cache_response_complete) 

284 ) 

285 started: Final = original.started_at 

286 available: Final = event.completion_start_time.timestamp() 

287 if context.invalidated or not complete or not started <= available <= context.collector.clock(): 

288 return context.capture.model_copy( 

289 update=MappingProxyType( 

290 { 

291 "observation": original.model_copy( 

292 update=MappingProxyType( 

293 { 

294 "available_at": max(started, context.collector.clock()), 

295 "reason": context.invalidated or "incomplete_response", 

296 } 

297 ) 

298 ) 

299 } 

300 ) 

301 ) 

302 target: Final = context.target 

303 if isinstance(target, UnsupportedPredictionTarget) or not supported_baseline_recipient(target, wire): 

304 return context.capture.model_copy( 

305 update=MappingProxyType( 

306 { 

307 "observation": original.model_copy( 

308 update=MappingProxyType( 

309 { 

310 "available_at": available, 

311 "reason": target.reason 

312 if isinstance(target, UnsupportedPredictionTarget) 

313 else "unsupported_baseline_recipient", 

314 } 

315 ) 

316 ) 

317 } 

318 ) 

319 ) 

320 body: Final = _JSON_BODY.validate_json(wire.content) 

321 same: Final = ( 

322 logging_obj.get_router_model_id() == context.baseline_deployment_id and body.get("model") == target.model 

323 ) 

324 plan, reason = await context.collector.plan(target, wire, body, usage) 

325 minimum: Final = get_prompt_cache_min_tokens(target.model) 

326 return context.capture.model_copy( 

327 update=MappingProxyType( 

328 { 

329 "observation": BaselineObservation( 

330 request_id=original.request_id, 

331 started_at=started, 

332 available_at=available, 

333 outcome="complete", 

334 baseline_equivalent=same, 

335 usage=usage, 

336 plan=plan, 

337 minimum_cache_tokens=minimum, 

338 reason=reason, 

339 ) 

340 } 

341 ) 

342 )