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

1import time 

2from collections.abc import Mapping 

3from types import MappingProxyType 

4from typing import Annotated, Final 

5 

6from fastapi import APIRouter, Depends, HTTPException, Request 

7from pydantic import BaseModel, JsonValue, TypeAdapter 

8 

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 

46 

47router: Final = APIRouter() 

48_REQUEST_DATA: Final = TypeAdapter(Mapping[str, object]) 

49 

50 

51class _CallerSettings(BaseModel): 

52 config: Mapping[str, object] | None = None 

53 

54 

55def has_request_transforms() -> bool: 

56 from litellm.proxy.hooks import PROXY_HOOKS 

57 

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 ) 

66 

67 

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 ) 

75 

76 

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 

80 

81 

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) 

91 

92 return count 

93 

94 

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) 

104 

105 

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 ) 

173 

174 

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. 

186 

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 

197 

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 )