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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from __future__ import annotations
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
13import httpx
14from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter
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
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
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
56class CapturedBaselineObservation(BaseModel):
57 model_config = ConfigDict(extra="forbid", frozen=True, strict=True)
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
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
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
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
95class _ResponseUsage(BaseModel):
96 model_config = ConfigDict(strict=True, from_attributes=True)
97 usage: Usage | None = None
100def _digest(value: object) -> str:
101 return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
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({})
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
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")
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
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"
216 async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None:
217 return await self._count(target, body)
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"
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 )
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")
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 )