Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/prompt_cache_prediction.py: 32%
109 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
1import time
2from collections.abc import Mapping
3from types import MappingProxyType
4from typing import Annotated, Final
6from fastapi import APIRouter, Depends, HTTPException, Request
7from pydantic import BaseModel, JsonValue, TypeAdapter
9import litellm
10from litellm._internal_context import current_billing_time, pinned_billing_time
11from litellm.caching.caching import DualCache
12from litellm.integrations.custom_logger import CustomLogger
13from litellm.llms.anthropic.prompt_cache_prediction import (
14 PromptPrefix,
15 TokenCounter,
16 UnsupportedPredictionTarget,
17 cache_scope,
18 count_prompt_tokens,
19 parse_prompt,
20 resolve_prediction_target,
21 supported_prediction_headers,
22)
23from litellm.proxy._types import UserAPIKeyAuth
24from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
25from litellm.proxy.auth.auth_utils import get_cache_prediction_deployments
26from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
27from litellm.proxy.common_utils.http_parsing_utils import (
28 _read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary
29)
30from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens
31from litellm.proxy.hooks.parallel_request_limiter_v3 import (
32 _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner
33)
34from litellm.proxy.hooks.prompt_cache_prediction import lookup
35from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
36from litellm.types.management_endpoints.prompt_cache_prediction import (
37 CacheCostScenario,
38 CacheEvidence,
39 CachePredictionArm,
40 CachePredictionRequest,
41 CachePredictionResponse,
42 CacheTokenBuckets,
43)
44from litellm.types.router import Deployment
45from litellm.utils import get_prompt_cache_min_tokens
47router: Final = APIRouter()
48_REQUEST_DATA: Final = TypeAdapter(Mapping[str, object])
51class _CallerSettings(BaseModel):
52 config: Mapping[str, object] | None = None
55def has_request_transforms() -> bool:
56 from litellm.proxy.hooks import PROXY_HOOKS
58 builtins: Final = frozenset(PROXY_HOOKS.values())
59 hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook")
60 callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger)
61 return any(
62 type(callback) not in builtins
63 and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks)
64 for callback in callbacks
65 )
68def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets:
69 return CacheTokenBuckets(
70 uncached_input_tokens=suffix_tokens,
71 cache_read_input_tokens=read_tokens,
72 cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0,
73 cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0,
74 )
77def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None:
78 cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens)
79 return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None
82def _capacity_counter(
83 limiter: _PROXY_MaxParallelRequestsHandler_v3,
84 caller: UserAPIKeyAuth,
85 model_name: str,
86 request_data: Mapping[str, object],
87) -> TokenCounter:
88 async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None:
89 async with limiter.request_capacity(caller, model_name, request_data=request_data):
90 return await count_prompt_tokens(model, api_key, body)
92 return count
95def _capacity_request_data(
96 http_request: Request, caller: UserAPIKeyAuth, request_data: Mapping[str, object]
97) -> Mapping[str, object]:
98 # The parsed-body cache retains only original top-level keys. Replay the
99 # shared idempotent tag merges on limiter-only data when auth added metadata.
100 data: Final = dict(request_data) # mutable-ok: the existing tag merge owners accept a dictionary out-param
101 LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(http_request, data, caller) # pyright: ignore[reportUnknownMemberType] # legacy tag owner takes the validated capacity dictionary
102 LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(data, caller) # pyright: ignore[reportUnknownMemberType] # legacy tag owner merges trusted key tags into capacity metadata
103 return MappingProxyType(data)
106async def predict_arm(
107 deployment: Deployment,
108 body: Mapping[str, JsonValue],
109 prefix: PromptPrefix,
110 caller_key_hash: str,
111 cache: DualCache,
112 token_counter: TokenCounter,
113) -> CachePredictionArm:
114 deployment_id: Final = deployment.model_info.id or ""
115 params: Final = deployment.litellm_params
116 unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model)
117 if deployment.model_info.blocked:
118 return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"}))
119 target: Final = resolve_prediction_target(params)
120 if isinstance(target, UnsupportedPredictionTarget):
121 return unknown.model_copy(update=MappingProxyType({"reason": target.reason}))
122 model: Final = target.model
123 api_key: Final = target.api_key
124 total_count: Final = await token_counter(model, api_key, body)
125 prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body)
126 if total_count is None or prefix_count is None or total_count < prefix_count:
127 return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"}))
128 scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model)
129 observation: Final = await lookup(cache, scope, prefix)
130 exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint
131 cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count
132 if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable):
133 return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"}))
134 suffix: Final = total_count - cacheable
135 evidence: Final = (
136 CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at)
137 if observation is not None
138 else None
139 )
140 if cacheable < get_prompt_cache_min_tokens(params.model):
141 disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count))
142 if disabled is None:
143 return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"}))
144 return CachePredictionArm(
145 deployment_id=deployment_id,
146 model=model,
147 cache_state="disabled",
148 reason="below_cache_minimum",
149 estimate=disabled,
150 cold=disabled,
151 warm=disabled,
152 token_count_source="anthropic_count_tokens",
153 )
154 fresh: Final = observation is not None and observation.expires_at > time.time()
155 read: Final = observation.cached_tokens if fresh and observation is not None else 0
156 with pinned_billing_time(current_billing_time()):
157 cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds))
158 warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds))
159 estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds))
160 if cold is None or warm is None or estimate is None:
161 return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"}))
162 return CachePredictionArm(
163 deployment_id=deployment_id,
164 model=model,
165 cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown",
166 reason=None if fresh else "observation_expired" if observation else "no_compatible_observation",
167 estimate=estimate,
168 cold=cold,
169 warm=warm,
170 evidence=evidence,
171 token_count_source="anthropic_count_tokens",
172 )
175@router.post(
176 "/cost/predict-cache",
177 tags=["Cost Tracking"], # mutable-ok: FastAPI requires a list for OpenAPI tags
178 response_model=CachePredictionResponse,
179)
180async def predict_cache_cost(
181 request: CachePredictionRequest,
182 http_request: Request,
183 user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
184) -> CachePredictionResponse:
185 """Compare the next native Anthropic request on two configured deployment IDs.
187 Estimates use provider token counting and recent successful cache telemetry for this key.
188 Unknown cache state uses the cold scenario when prices/counts are available. Cache observations
189 do not guarantee retention. v0 supports one message-content breakpoint, text and client tools;
190 system/tool-only breakpoints, thinking, images, nondefault Anthropic versions, beta headers and
191 request transforms are unknown.
192 Each provider count consumes one RPM unit and holds concurrency capacity; a comparison uses
193 up to four counts. The legacy rate limiter returns unknown without contacting the provider.
194 This endpoint does not generate tokens, prewarm caches, choose a model or alter routing.
195 """
196 from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
198 if llm_router is None: 198 ↛ 199line 198 didn't jump to line 199 because the condition on line 198 was never true
199 raise HTTPException(status_code=503, detail="Model router is unavailable")
200 deployments: Final = get_cache_prediction_deployments(
201 current_deployment_id=request.current_deployment_id,
202 candidate_deployment_id=request.candidate_deployment_id,
203 llm_router=llm_router,
204 team_id=user_api_key_dict.team_id,
205 )
206 if deployments is None: 206 ↛ 208line 206 didn't jump to line 208 because the condition on line 206 was always true
207 raise HTTPException(status_code=404, detail="Deployment not found")
208 current, candidate = deployments
209 for deployment in (current, candidate):
210 await can_key_call_resolved_model(
211 model=deployment.model_name,
212 llm_model_list=llm_router.get_model_list(),
213 valid_token=user_api_key_dict,
214 llm_router=llm_router,
215 )
216 prefix: Final = parse_prompt(request.request)
217 caller: Final = user_api_key_dict.api_key
218 caller_settings: Final = _CallerSettings.model_validate(user_api_key_dict, from_attributes=True)
219 unsupported_transform: Final = bool(caller_settings.config) or has_request_transforms()
220 unsupported_headers: Final = not supported_prediction_headers(http_request.headers)
221 limiter: Final = proxy_logging_obj.get_proxy_hook("parallel_request_limiter")
222 if (
223 prefix is None
224 or not caller
225 or unsupported_transform
226 or unsupported_headers
227 or not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3)
228 ):
229 reason: Final = (
230 "unsupported_provider_headers"
231 if unsupported_headers
232 else "unsupported_request_transform"
233 if unsupported_transform
234 else "unsupported_prompt_shape"
235 if prefix is None
236 else "caller_identity_unavailable"
237 if not caller
238 else "limiter_unavailable"
239 )
240 return CachePredictionResponse(
241 stay=CachePredictionArm(deployment_id=request.current_deployment_id, reason=reason),
242 switch=CachePredictionArm(deployment_id=request.candidate_deployment_id, reason=reason),
243 switch_delta=None,
244 cache_rebuild_penalty=None,
245 )
246 request_data: Final = _capacity_request_data(
247 http_request, user_api_key_dict, _REQUEST_DATA.validate_python(await _read_request_body(http_request))
248 )
249 stay: Final = await predict_arm(
250 current,
251 request.request,
252 prefix,
253 caller,
254 proxy_logging_obj.internal_usage_cache.dual_cache,
255 _capacity_counter(limiter, user_api_key_dict, current.model_name, request_data),
256 )
257 switch: Final = (
258 stay
259 if current.model_info.id == candidate.model_info.id
260 else await predict_arm(
261 candidate,
262 request.request,
263 prefix,
264 caller,
265 proxy_logging_obj.internal_usage_cache.dual_cache,
266 _capacity_counter(limiter, user_api_key_dict, candidate.model_name, request_data),
267 )
268 )
269 return CachePredictionResponse(
270 stay=stay,
271 switch=switch,
272 switch_delta=(switch.estimate.input_cost - stay.estimate.input_cost)
273 if switch.estimate is not None and stay.estimate is not None
274 else None,
275 cache_rebuild_penalty=(switch.estimate.input_cost - switch.warm.input_cost)
276 if switch.estimate is not None and switch.warm is not None
277 else None,
278 )