Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/spend_tracking/budget_reservation.py: 26%
644 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 json
5import math
6import time
7from collections.abc import Mapping, Sequence
8from dataclasses import dataclass
9from datetime import datetime, timedelta, timezone
10from types import MappingProxyType
11from typing import Final, NoReturn, SupportsFloat, SupportsIndex, SupportsInt, cast
13from fastapi import HTTPException, status
15import litellm
16from litellm._logging import verbose_proxy_logger
17from litellm.litellm_core_utils.duration_parser import duration_in_seconds
18from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import select_tier_for_input, tier_rate
19from litellm.proxy._types import (
20 Litellm_EntityType,
21 LiteLLM_TeamMembership,
22 LiteLLM_TeamTable,
23 LiteLLM_UserTable,
24 UserAPIKeyAuth,
25)
26from litellm.proxy.auth.auth_utils import get_model_from_request
27from litellm.proxy.auth.budget_throttle import should_throttle_budget_exceeded
28from litellm.proxy.auth.route_checks import RouteChecks
29from litellm.proxy.common_utils.user_api_key_cache import (
30 UserApiKeyCache,
31 end_user_cache_key,
32 model_access_group_cache_key,
33 model_access_group_spend_counter_key,
34 project_cache_key,
35 project_spend_counter_key,
36 tag_cache_key,
37 team_membership_reservation_cache_key,
38)
39from litellm.proxy.spend_tracking.input_tokens import count_input_tokens, count_input_tokens_for_model
40from litellm.proxy.spend_tracking.spend_counter_batch import PendingSpendIncrement, spend_counter_batch_scope
41from litellm.proxy.utils import PrismaClient, ProxyLogging
42from litellm.router import Router
43from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
44from litellm.types.router import DeploymentTypedDict
47@dataclass
48class _BudgetCounter:
49 counter_key: str
50 max_budget: float
51 fallback_spend: float
52 entity_type: str
53 entity_id: str
54 source_cache_key: str | None = None
55 spend_log_entity_id: str | None = None
56 window_duration: str | None = None
57 window_start: datetime | None = None
60_COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = {
61 "Key": Litellm_EntityType.KEY.value,
62 "Team": Litellm_EntityType.TEAM.value,
63 "TeamMember": Litellm_EntityType.TEAM_MEMBER.value,
64 "User": Litellm_EntityType.USER.value,
65 "EndUser": Litellm_EntityType.END_USER.value,
66 "Tag": Litellm_EntityType.TAG.value,
67 "Model access group": Litellm_EntityType.MODEL_ACCESS_GROUP.value,
68 "Organization": Litellm_EntityType.ORGANIZATION.value,
69 "Project": Litellm_EntityType.PROJECT.value,
70}
73class _CounterReservationUnavailable(Exception):
74 def __init__(
75 self,
76 touched_counter: bool = False,
77 counter_invalidated: bool = False,
78 ) -> None:
79 self.touched_counter = touched_counter
80 self.counter_invalidated = counter_invalidated
81 super().__init__("Counter reservation unavailable")
84def _raise_reservation_unavailable(counter_key: str) -> NoReturn:
85 verbose_proxy_logger.warning(
86 "fail_closed_budget_enforcement: rejecting request — budget reservation for %s could not be written",
87 counter_key,
88 )
89 raise HTTPException(
90 status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
91 detail=(
92 "Budget enforcement unavailable: the budget reservation could not "
93 "be written to the spend counter backend, and "
94 "fail_closed_budget_enforcement is enabled, so the request was "
95 "rejected to avoid exceeding the configured budget. Retry shortly."
96 ),
97 )
100def get_reserved_counter_keys(budget_reservation: dict | None) -> set:
101 if not budget_reservation:
102 return set()
103 entries: Final = budget_reservation.get("entries") or []
104 return {
105 entry["counter_key"] for entry in entries if isinstance(entry, dict) and entry.get("counter_key") is not None
106 }
109_lease_renewals: Final[set[asyncio.Task[None]]] = set() # mutable-ok: asyncio only weak-refs pending tasks
112def _start_reservation_lease_renewal(budget_reservation: Mapping[str, object], counter_keys: frozenset[str]) -> None:
113 """A reservation lives inside spend counter keys that expire on their Redis TTL. Renew the TTL
114 while the request is in flight so a request longer than the TTL does not drop its
115 reservation and admit concurrent requests against the DB floor on any worker."""
116 from litellm.proxy.proxy_server import spend_counter_cache
118 if spend_counter_cache.redis_cache is None or not counter_keys:
119 return
120 task: Final = asyncio.create_task(
121 _renew_reservation_lease(
122 budget_reservation=budget_reservation,
123 counter_keys=counter_keys,
124 interval=spend_counter_cache.redis_cache.default_ttl / 2,
125 request_task=asyncio.current_task(),
126 )
127 )
128 _lease_renewals.add(task)
129 task.add_done_callback(_lease_renewals.discard)
132async def _renew_reservation_lease(
133 budget_reservation: Mapping[str, object],
134 counter_keys: frozenset[str],
135 interval: float,
136 request_task: asyncio.Task[object] | None,
137) -> None:
138 """Stops on finalization or once the request task that took the reservation is gone, so a
139 disconnect path that skipped reconciliation falls back to the plain counter TTL."""
140 from litellm.proxy.proxy_server import refresh_spend_counter_ttl
142 deadline: Final = time.monotonic() + litellm.request_timeout
143 while time.monotonic() < deadline:
144 await asyncio.sleep(interval)
145 if budget_reservation.get("finalized") is True or (request_task is not None and request_task.done()):
146 return
147 for counter_key in counter_keys:
148 await refresh_spend_counter_ttl(counter_key=counter_key)
151def _key_reservation_should_release_for_throttle(counter_key: str, valid_token: UserAPIKeyAuth | None) -> bool:
152 """
153 Whether an over-budget key's own ``max_budget`` reservation should be
154 released rather than blocked, because the key opted into throttling: the
155 rate limiter slows it instead. Only the key's own ``max_budget`` counter is
156 exempt; team/user/window counters still enforce normally, and under-budget
157 requests never reach this branch so their concurrent-overspend protection is
158 untouched.
159 """
160 if valid_token is None:
161 return False
162 return counter_key == f"spend:key:{valid_token.token}" and should_throttle_budget_exceeded(valid_token)
165async def _apply_over_budget_reservation_policy(
166 counter: _BudgetCounter,
167 valid_token: UserAPIKeyAuth | None,
168 entry: dict[str, float | str],
169 applied_entries: list[dict[str, float | str]],
170 reservation_cost: float,
171 current_spend: float,
172 fail_closed_budget_enforcement: bool = False,
173) -> float:
174 """
175 Decide what to do when a counter is over budget, and return the reservation
176 cost to carry into the next counter. Three outcomes: an over-budget key that
177 opted into throttling releases its own reservation (the rate limiter slows
178 it) and keeps the cost; a partially-remaining budget resizes the reservation
179 down to what is left, unless strict enforcement is on, because the known
180 estimate already does not fit; anything else hard-blocks by raising.
181 """
182 if _key_reservation_should_release_for_throttle(counter.counter_key, valid_token):
183 await _release_applied_entries_best_effort(entries=[entry], default_reserved_cost=reservation_cost)
184 applied_entries.remove(entry)
185 return reservation_cost
187 remaining_before_reservation: Final = counter.max_budget - (current_spend - reservation_cost)
188 if remaining_before_reservation <= 1e-12:
189 _raise_counter_budget_exceeded(counter=counter, current_cost=current_spend)
190 if fail_closed_budget_enforcement and current_spend - counter.max_budget > 1e-12:
191 _raise_counter_budget_exceeded(
192 counter=counter,
193 current_cost=current_spend - reservation_cost,
194 estimated_cost=reservation_cost,
195 )
196 await _resize_applied_reservation(
197 entries=applied_entries,
198 current_reserved_cost=reservation_cost,
199 new_reserved_cost=remaining_before_reservation,
200 )
201 return remaining_before_reservation
204def _raise_counter_budget_exceeded(
205 counter: _BudgetCounter,
206 current_cost: float,
207 estimated_cost: float | None = None,
208) -> NoReturn:
209 estimate_detail: Final = "" if estimated_cost is None else f"Estimated request cost: {estimated_cost}, "
210 raise litellm.BudgetExceededError(
211 current_cost=current_cost,
212 max_budget=counter.max_budget,
213 message=(
214 "Budget has been exceeded! "
215 f"{counter.entity_type}={counter.entity_id} "
216 f"Current cost: {current_cost}, "
217 f"{estimate_detail}"
218 f"Max budget: {counter.max_budget}"
219 ),
220 entity_type=_COUNTER_ENTITY_TYPES.get(counter.entity_type),
221 entity_id=counter.spend_log_entity_id or counter.entity_id,
222 )
225_UNBILLED_ROUTES: Final[frozenset[str]] = frozenset(
226 {
227 "/models",
228 "/v1/models",
229 "/utils/token_counter",
230 "/responses/input_tokens",
231 "/v1/responses/input_tokens",
232 "/openai/v1/responses/input_tokens",
233 }
234)
235_TOKEN_COUNTING_SEGMENTS: Final[frozenset[str]] = frozenset({"count_tokens", "count-tokens"})
236_TOKEN_COUNTING_ACTION: Final = "countTokens"
239def _is_token_counting_route(route: str) -> bool:
240 resource, _, action = route.rsplit("/", 1)[-1].partition(":")
241 return resource in _TOKEN_COUNTING_SEGMENTS or action == _TOKEN_COUNTING_ACTION
244def _is_unbilled_route(route: str) -> bool:
245 return route in _UNBILLED_ROUTES or _is_token_counting_route(route)
248async def reserve_budget_for_request(
249 request_body: dict,
250 route: str,
251 llm_router: Router | None,
252 valid_token: UserAPIKeyAuth | None,
253 team_object: LiteLLM_TeamTable | None,
254 user_object: LiteLLM_UserTable | None,
255 prisma_client: PrismaClient | None,
256 user_api_key_cache: UserApiKeyCache,
257 proxy_logging_obj: ProxyLogging,
258 end_user_id: str | None = None,
259 end_user_object: object = None,
260 apply_user_budget_to_team_keys: bool = False,
261 fail_closed_budget_enforcement: bool = False,
262 raw_body: bytes | None = None,
263) -> dict | None:
264 if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
265 return None
266 if _is_unbilled_route(route):
267 return None
268 if get_model_from_request(request_body, route, llm_router=llm_router) is None:
269 return None
271 counters: Final = await _get_budget_counters(
272 request_body=request_body,
273 valid_token=valid_token,
274 team_object=team_object,
275 user_object=user_object,
276 prisma_client=prisma_client,
277 user_api_key_cache=user_api_key_cache,
278 proxy_logging_obj=proxy_logging_obj,
279 end_user_id=end_user_id,
280 end_user_object=end_user_object,
281 apply_user_budget_to_team_keys=apply_user_budget_to_team_keys,
282 )
283 if not counters: 283 ↛ 286line 283 didn't jump to line 286 because the condition on line 283 was always true
284 return None
286 input_token_counts: Final = await count_request_input_tokens(
287 request_body=request_body,
288 route=route,
289 llm_router=llm_router,
290 raw_body=raw_body,
291 )
293 current_spend_by_counter_key: Final[dict[str, float]] = {}
294 reservation_cost = estimate_request_max_cost(
295 request_body=request_body,
296 route=route,
297 llm_router=llm_router,
298 input_token_counts=input_token_counts,
299 )
300 # estimate_request_max_cost still returns None when the model is unknown
301 # to the cost map (no token-priced cost fields, e.g. image/audio routes).
302 # In that case we fall back to read-time enforcement only.
303 if reservation_cost is None or reservation_cost <= 0:
304 return None
306 applied_entries: Final[list[dict[str, float | str]]] = []
307 try:
308 with _counters_batch_scope(frozenset(counter.counter_key for counter in counters)):
309 for counter in counters:
310 entry = _counter_to_reservation_entry(
311 counter=counter,
312 reserved_cost=reservation_cost,
313 )
314 applied_entries.append(entry)
315 try:
316 reserved_value = await _reserve_counter(
317 counter=counter,
318 reservation_cost=reservation_cost,
319 )
320 except _CounterReservationUnavailable as exc:
321 if exc.touched_counter and not exc.counter_invalidated:
322 await _release_applied_entries_best_effort(
323 entries=[entry],
324 default_reserved_cost=reservation_cost,
325 )
326 applied_entries.remove(entry)
327 if fail_closed_budget_enforcement:
328 _raise_reservation_unavailable(counter_key=counter.counter_key)
329 continue
331 if reserved_value is not None:
332 current_spend = reserved_value
333 else:
334 cached_spend = current_spend_by_counter_key.get(counter.counter_key)
335 if cached_spend is None:
336 cached_spend = await _get_current_counter_value(counter=counter)
337 current_spend = cached_spend + reservation_cost
338 if current_spend > counter.max_budget:
339 reservation_cost = await _apply_over_budget_reservation_policy(
340 counter=counter,
341 valid_token=valid_token,
342 entry=entry,
343 applied_entries=applied_entries,
344 reservation_cost=reservation_cost,
345 current_spend=current_spend,
346 fail_closed_budget_enforcement=fail_closed_budget_enforcement,
347 )
348 continue
349 except Exception:
350 await _release_applied_entries_best_effort(
351 entries=applied_entries,
352 default_reserved_cost=reservation_cost,
353 )
354 raise
356 if not applied_entries:
357 return None
359 input_cost: Final = estimate_request_input_cost(
360 request_body=request_body,
361 route=route,
362 llm_router=llm_router,
363 input_token_counts=input_token_counts,
364 )
365 budget_reservation: Final = {
366 "reserved_cost": reservation_cost,
367 "entries": applied_entries,
368 "finalized": False,
369 "callback_bound": False,
370 "input_cost": min(float(input_cost or 0.0), reservation_cost),
371 "input_tokens": max(input_token_counts.values(), default=None),
372 }
373 _start_reservation_lease_renewal(
374 budget_reservation=budget_reservation,
375 counter_keys=frozenset(get_reserved_counter_keys(budget_reservation=budget_reservation)),
376 )
377 return budget_reservation
380async def reconcile_budget_reservation(
381 budget_reservation: dict | None,
382 actual_cost: float | None,
383 finalize: bool = True,
384) -> None:
385 if not budget_reservation or budget_reservation.get("finalized") is True:
386 return
388 reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0)
389 actual: Final = float(actual_cost or 0.0)
390 await _set_reserved_entries_actual_cost(
391 entries=budget_reservation.get("entries") or [],
392 actual_cost=actual,
393 default_reserved_cost=reserved_cost,
394 )
395 if finalize:
396 budget_reservation["finalized"] = True
399async def release_budget_reservation(budget_reservation: dict | None) -> None:
400 await reconcile_budget_reservation(
401 budget_reservation=budget_reservation,
402 actual_cost=0.0,
403 )
406async def release_budget_reservation_on_cancel(
407 budget_reservation: dict | None,
408) -> None:
409 """Reconcile a still-open reservation when the request is cancelled mid-flight.
411 A client disconnect or timeout cancels the request task, which surfaces as
412 CancelledError / GeneratorExit rather than a normal exception, so neither the
413 success cost callback nor the failure hook runs and the pre-call reservation
414 is never reconciled. Left alone it pins the spend counter above real spend
415 and 429s subsequent requests until the counter's TTL expires.
417 Reconcile to the request's input-token cost rather than refunding to zero:
418 by the time a request is cancelled in-flight the provider call was already
419 dispatched, so the input tokens were billed even if no chunk reached the
420 client. Refunding to zero would let a caller abort pre-token to dodge that
421 charge; the worst-case output portion of the reservation is still released.
423 asyncio.shield keeps the reconcile running to completion even though the
424 surrounding task is being cancelled. The `finalized` guard makes this a no-op
425 when success/failure handling already reconciled, so calling it on every
426 cancellation path is safe.
427 """
428 if not budget_reservation or budget_reservation.get("finalized") is True:
429 return
430 incurred_cost: Final = float(budget_reservation.get("input_cost") or 0.0)
431 try:
432 await asyncio.shield(
433 reconcile_budget_reservation(budget_reservation=budget_reservation, actual_cost=incurred_cost)
434 )
435 except (asyncio.CancelledError, Exception):
436 pass
439async def invalidate_budget_reservation_counters(
440 budget_reservation: dict | None,
441) -> None:
442 if budget_reservation is None:
443 return
445 from litellm.proxy.proxy_server import _invalidate_spend_counter
447 for counter_key in get_reserved_counter_keys(budget_reservation=budget_reservation):
448 await _invalidate_spend_counter(counter_key=counter_key)
451async def release_or_invalidate_budget_reservation(
452 budget_reservation: dict | None, # mutable-ok: stamps finalized on the caller's shared reservation dict
453) -> None:
454 """Reconcile a still-open reservation on a terminal path that settles no cost.
456 A failed or upstream-refused request never runs the success cost callback, so
457 its pre-call reservation stays open and keeps the spend counter pinned above
458 real spend until the counter's TTL expires, 429ing later requests on the same
459 key. Release it to zero; if the release itself fails (e.g. the counter store is
460 unreachable) drop the reserved counters directly and mark the reservation
461 finalized so nothing reprocesses it. Idempotent: the finalized guard makes a
462 second call a no-op once success or failure handling already reconciled.
463 """
464 if budget_reservation is None or budget_reservation.get("finalized") is True:
465 return
466 try:
467 await asyncio.shield(release_budget_reservation(budget_reservation=budget_reservation))
468 except Exception: # noqa: BLE001 # a cleanup failure must not pin the counter; drop it directly instead
469 verbose_proxy_logger.exception("Failed to release budget reservation; invalidating counters")
470 try:
471 await invalidate_budget_reservation_counters(budget_reservation=budget_reservation)
472 except Exception: # noqa: BLE001 # nothing left to try; the finalized stamp below keeps it from being reprocessed
473 verbose_proxy_logger.exception("Failed to invalidate budget reservation counters after release failed")
474 finally:
475 budget_reservation["finalized"] = True
478async def release_unbound_budget_reservation(budget_reservation: Mapping[str, object]) -> None:
479 """Release a reservation no logging callback took ownership of, once the request ended.
481 A handler whose litellm call never builds a logging object (batch cancel, file
482 content, anything without the client decorator) runs no cost callback, so nothing
483 else would ever reconcile its reservation. A bound reservation is left alone: its
484 success or failure handler settles it, possibly after the response has been sent.
485 """
486 if not isinstance(budget_reservation, dict) or budget_reservation.get("callback_bound") is True:
487 return
488 await release_or_invalidate_budget_reservation(budget_reservation=budget_reservation)
491async def _get_budget_counters(
492 request_body: dict,
493 valid_token: UserAPIKeyAuth,
494 team_object: LiteLLM_TeamTable | None,
495 user_object: LiteLLM_UserTable | None,
496 prisma_client: PrismaClient | None,
497 user_api_key_cache: UserApiKeyCache,
498 proxy_logging_obj: ProxyLogging,
499 end_user_id: str | None = None,
500 end_user_object: object = None,
501 apply_user_budget_to_team_keys: bool = False,
502) -> list[_BudgetCounter]:
503 counters: Final[list[_BudgetCounter]] = []
505 if valid_token.token is not None: 505 ↛ 527line 505 didn't jump to line 527 because the condition on line 505 was always true
506 if valid_token.max_budget is not None and valid_token.max_budget > 0: 506 ↛ 507line 506 didn't jump to line 507 because the condition on line 506 was never true
507 counters.append(
508 _BudgetCounter(
509 counter_key=f"spend:key:{valid_token.token}",
510 source_cache_key=valid_token.token,
511 max_budget=float(valid_token.max_budget),
512 fallback_spend=float(valid_token.spend or 0.0),
513 entity_type="Key",
514 entity_id=valid_token.token,
515 )
516 )
517 counters.extend(
518 _get_budget_limit_counters(
519 entity_prefix=f"spend:key:{valid_token.token}",
520 entity_type="Key",
521 entity_id=valid_token.token,
522 budget_limits=valid_token.budget_limits,
523 fallback_spend=float(valid_token.spend or 0.0),
524 )
525 )
527 if team_object is not None and team_object.team_id is not None: 527 ↛ 528line 527 didn't jump to line 528 because the condition on line 527 was never true
528 team_id: Final = team_object.team_id
529 if team_object.max_budget is not None and team_object.max_budget > 0:
530 counters.append(
531 _BudgetCounter(
532 counter_key=f"spend:team:{team_id}",
533 source_cache_key=f"team_id:{team_id}",
534 max_budget=float(team_object.max_budget),
535 fallback_spend=float(team_object.spend or 0.0),
536 entity_type="Team",
537 entity_id=team_id,
538 )
539 )
540 counters.extend(
541 _get_budget_limit_counters(
542 entity_prefix=f"spend:team:{team_id}",
543 entity_type="Team",
544 entity_id=team_id,
545 budget_limits=team_object.budget_limits,
546 fallback_spend=float(team_object.spend or 0.0),
547 )
548 )
550 is_team_key: Final = team_object is not None and team_object.team_id is not None
551 if ( 551 ↛ 558line 551 didn't jump to line 558 because the condition on line 551 was never true
552 (not is_team_key or apply_user_budget_to_team_keys)
553 and user_object is not None
554 and user_object.user_id is not None
555 and user_object.max_budget is not None
556 and user_object.max_budget > 0
557 ):
558 counters.append(
559 _BudgetCounter(
560 counter_key=f"spend:user:{user_object.user_id}",
561 source_cache_key=user_object.user_id,
562 max_budget=float(user_object.max_budget),
563 fallback_spend=float(user_object.spend or 0.0),
564 entity_type="User",
565 entity_id=user_object.user_id,
566 )
567 )
569 end_user_counter: Final = await _get_end_user_budget_counter(
570 valid_token=valid_token,
571 end_user_id=end_user_id,
572 end_user_object=end_user_object,
573 )
574 if end_user_counter is not None: 574 ↛ 575line 574 didn't jump to line 575 because the condition on line 574 was never true
575 counters.append(end_user_counter)
577 counters.extend(
578 await _get_tag_budget_counters(
579 request_body=request_body,
580 prisma_client=prisma_client,
581 user_api_key_cache=user_api_key_cache,
582 proxy_logging_obj=proxy_logging_obj,
583 )
584 )
586 counters.extend(
587 await _get_model_access_group_budget_counters(
588 valid_token=valid_token,
589 prisma_client=prisma_client,
590 user_api_key_cache=user_api_key_cache,
591 )
592 )
594 team_member_counter: Final = await _get_team_member_budget_counter(
595 valid_token=valid_token,
596 team_object=team_object,
597 user_object=user_object,
598 user_api_key_cache=user_api_key_cache,
599 )
600 if team_member_counter is not None: 600 ↛ 601line 600 didn't jump to line 601 because the condition on line 600 was never true
601 counters.append(team_member_counter)
603 org_counter: Final = await _get_org_budget_counter(
604 valid_token=valid_token,
605 team_object=team_object,
606 user_api_key_cache=user_api_key_cache,
607 )
608 if org_counter is not None: 608 ↛ 609line 608 didn't jump to line 609 because the condition on line 608 was never true
609 counters.append(org_counter)
611 project_counter: Final = await _get_project_budget_counter(
612 valid_token=valid_token,
613 user_api_key_cache=user_api_key_cache,
614 )
615 if project_counter is not None: 615 ↛ 616line 615 didn't jump to line 616 because the condition on line 615 was never true
616 counters.append(project_counter)
618 return counters
621async def _get_end_user_budget_counter(
622 valid_token: UserAPIKeyAuth,
623 end_user_id: str | None,
624 end_user_object: object,
625) -> _BudgetCounter | None:
626 end_user_id = end_user_id or valid_token.end_user_id
627 if end_user_id is None:
628 return None
630 source_cache_key: Final = end_user_cache_key(end_user_id)
631 max_budget = _to_float(valid_token.end_user_max_budget)
632 fallback_spend = 0.0
633 if end_user_object is not None: 633 ↛ 634line 633 didn't jump to line 634 because the condition on line 633 was never true
634 fallback_spend = _to_float(_get_value(end_user_object, "spend")) or 0.0
635 if max_budget is None:
636 budget_table: Final = _get_value(end_user_object, "litellm_budget_table")
637 max_budget = _to_float(_get_value(budget_table, "max_budget"))
639 if max_budget is None or max_budget <= 0: 639 ↛ 642line 639 didn't jump to line 642 because the condition on line 639 was always true
640 return None
642 return _BudgetCounter(
643 counter_key=f"spend:end_user:{end_user_id}",
644 source_cache_key=source_cache_key,
645 max_budget=max_budget,
646 fallback_spend=fallback_spend,
647 entity_type="EndUser",
648 entity_id=end_user_id,
649 )
652async def _get_tag_budget_counters(
653 request_body: dict,
654 prisma_client: PrismaClient | None,
655 user_api_key_cache: UserApiKeyCache,
656 proxy_logging_obj: ProxyLogging,
657) -> list[_BudgetCounter]:
658 from litellm.proxy.auth.auth_checks import get_tag_objects_batch
659 from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
661 tag_names: Final = _dedupe_tags(get_tags_from_request_body(request_body=request_body))
662 if not tag_names:
663 return []
665 tag_objects: Final = await get_tag_objects_batch(
666 tag_names=tag_names,
667 prisma_client=prisma_client,
668 user_api_key_cache=user_api_key_cache,
669 proxy_logging_obj=proxy_logging_obj,
670 )
672 counters: Final[list[_BudgetCounter]] = []
673 for tag_name in tag_names:
674 tag_object = tag_objects.get(tag_name)
675 if tag_object is None:
676 continue
677 budget_table = _get_value(tag_object, "litellm_budget_table")
678 max_budget = _to_float(_get_value(budget_table, "max_budget"))
679 if max_budget is None or max_budget <= 0: 679 ↛ 681line 679 didn't jump to line 681 because the condition on line 679 was always true
680 continue
681 counters.append(
682 _BudgetCounter(
683 counter_key=f"spend:tag:{tag_name}",
684 source_cache_key=tag_cache_key(tag_name),
685 max_budget=max_budget,
686 fallback_spend=_to_float(_get_value(tag_object, "spend")) or 0.0,
687 entity_type="Tag",
688 entity_id=tag_name,
689 )
690 )
691 return counters
694async def _get_model_access_group_budget_counters(
695 valid_token: UserAPIKeyAuth,
696 prisma_client: PrismaClient | None,
697 user_api_key_cache: UserApiKeyCache,
698) -> list[_BudgetCounter]:
699 """Reservation counters for the model access groups that authorized this request.
701 The names come off the auth object rather than the request body: ``common_checks`` already
702 resolved which granted groups serve the requested model, and re-deriving that here would both
703 duplicate the walk and risk disagreeing with what the spend writer attributes.
704 """
705 from litellm.proxy.auth.auth_checks import get_model_access_group_budgets_batch
707 group_names: Final = tuple(dict.fromkeys(valid_token.matched_model_access_groups or ()))
708 if not group_names: 708 ↛ 711line 708 didn't jump to line 711 because the condition on line 708 was always true
709 return []
711 budgets: Final = await get_model_access_group_budgets_batch(
712 access_group_names=group_names,
713 prisma_client=prisma_client,
714 user_api_key_cache=user_api_key_cache,
715 )
716 candidates: Final = (_model_access_group_counter(group, budgets.get(group)) for group in group_names)
717 return [counter for counter in candidates if counter is not None]
720def _model_access_group_counter(group: str, budget: ModelAccessGroupBudget | None) -> _BudgetCounter | None:
721 """A counter for one group, or nothing when the group carries no budget to reserve against."""
722 if budget is None or budget.max_budget is None or budget.max_budget <= 0:
723 return None
724 return _BudgetCounter(
725 counter_key=model_access_group_spend_counter_key(group),
726 source_cache_key=model_access_group_cache_key(group),
727 max_budget=budget.max_budget,
728 fallback_spend=budget.spend,
729 entity_type="Model access group",
730 entity_id=group,
731 )
734def _dedupe_tags(tags: list[str]) -> list[str]:
735 seen: Final = set()
736 deduped_tags: Final = []
737 for tag in tags:
738 if tag in seen:
739 continue
740 seen.add(tag)
741 deduped_tags.append(tag)
742 return deduped_tags
745async def _get_team_member_budget_counter(
746 valid_token: UserAPIKeyAuth,
747 team_object: LiteLLM_TeamTable | None,
748 user_object: LiteLLM_UserTable | None,
749 user_api_key_cache: UserApiKeyCache,
750) -> _BudgetCounter | None:
751 if team_object is None or team_object.team_id is None or user_object is None or valid_token.user_id is None: 751 ↛ 754line 751 didn't jump to line 754 because the condition on line 751 was always true
752 return None
754 membership_cache_key: Final = team_membership_reservation_cache_key(
755 user_id=valid_token.user_id, team_id=team_object.team_id
756 )
757 cached_team_membership: Final = await user_api_key_cache.async_get_cache(key=membership_cache_key)
758 team_membership: LiteLLM_TeamMembership | None = None
759 if isinstance(cached_team_membership, LiteLLM_TeamMembership):
760 team_membership = cached_team_membership
761 elif isinstance(cached_team_membership, dict):
762 team_membership = LiteLLM_TeamMembership(**cached_team_membership)
764 member_budget_row: Final = team_membership.litellm_budget_table if team_membership is not None else None
765 now: Final = datetime.now(timezone.utc)
766 team_member_budget: float | None = None
767 if member_budget_row is not None and member_budget_row.max_budget is not None:
768 team_member_budget = member_budget_row.effective_max_budget(now=now)
769 else:
770 default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id")
771 if isinstance(default_budget_id, str):
772 default_budget: Final = await user_api_key_cache.async_get_cache(
773 key=f"team_member_default_budget:{default_budget_id}",
774 )
775 default_cap: Final = _to_float(_get_value(default_budget, "max_budget"))
776 if default_cap is not None and default_cap > 0:
777 team_member_budget = default_cap + (
778 member_budget_row.active_temp_budget_increase(now=now) if member_budget_row is not None else 0.0
779 )
781 if team_member_budget is None or team_member_budget <= 0:
782 return None
784 team_member_spend = cast(LiteLLM_TeamMembership, team_membership).spend if team_membership is not None else 0.0
785 return _BudgetCounter(
786 counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}",
787 source_cache_key=membership_cache_key,
788 max_budget=float(team_member_budget),
789 fallback_spend=float(team_member_spend or 0.0),
790 entity_type="TeamMember",
791 entity_id=f"{valid_token.user_id}:{team_object.team_id}",
792 )
795async def _get_org_budget_counter(
796 valid_token: UserAPIKeyAuth,
797 team_object: LiteLLM_TeamTable | None,
798 user_api_key_cache: UserApiKeyCache,
799) -> _BudgetCounter | None:
800 org_id: str | None = None
801 if valid_token.org_id is not None: 801 ↛ 802line 801 didn't jump to line 802 because the condition on line 801 was never true
802 org_id = valid_token.org_id
803 elif team_object is not None and team_object.organization_id is not None: 803 ↛ 804line 803 didn't jump to line 804 because the condition on line 803 was never true
804 org_id = team_object.organization_id
805 if org_id is None: 805 ↛ 808line 805 didn't jump to line 808 because the condition on line 805 was always true
806 return None
808 org_table: Final = await user_api_key_cache.async_get_cache(
809 key=f"org_id:{org_id}:with_budget",
810 )
811 if org_table is None:
812 return None
814 org_budget_table: Final = _get_value(org_table, "litellm_budget_table")
815 if org_budget_table is None:
816 return None
818 org_max_budget: Final = _to_float(_get_value(org_budget_table, "max_budget"))
819 if org_max_budget is None or org_max_budget <= 0:
820 return None
822 org_spend: Final = _to_float(_get_value(org_table, "spend")) or 0.0
823 return _BudgetCounter(
824 counter_key=f"spend:org:{org_id}",
825 source_cache_key=f"org_id:{org_id}:with_budget",
826 max_budget=org_max_budget,
827 fallback_spend=org_spend,
828 entity_type="Organization",
829 entity_id=org_id,
830 )
833async def _get_project_budget_counter(
834 valid_token: UserAPIKeyAuth,
835 user_api_key_cache: UserApiKeyCache,
836) -> _BudgetCounter | None:
837 if valid_token.project_id is None: 837 ↛ 840line 837 didn't jump to line 840 because the condition on line 837 was always true
838 return None
840 source_cache_key: Final = project_cache_key(valid_token.project_id)
841 project_object: Final = await user_api_key_cache.async_get_cache(key=source_cache_key)
842 if project_object is None:
843 return None
845 project_budget_table: Final = _get_value(project_object, "litellm_budget_table")
846 if project_budget_table is None:
847 return None
849 project_max_budget: Final = _to_float(_get_value(project_budget_table, "max_budget"))
850 if project_max_budget is None or project_max_budget <= 0 or not math.isfinite(project_max_budget):
851 return None
853 return _BudgetCounter(
854 counter_key=project_spend_counter_key(valid_token.project_id),
855 source_cache_key=source_cache_key,
856 max_budget=project_max_budget,
857 fallback_spend=_to_float(_get_value(project_object, "spend")) or 0.0,
858 entity_type="Project",
859 entity_id=valid_token.project_id,
860 )
863def _get_budget_limit_counters(
864 entity_prefix: str,
865 entity_type: str,
866 entity_id: str,
867 budget_limits: Sequence[object] | None,
868 fallback_spend: float,
869) -> list[_BudgetCounter]:
870 counters: Final[list[_BudgetCounter]] = []
871 if not budget_limits: 871 ↛ 874line 871 didn't jump to line 874 because the condition on line 871 was always true
872 return counters
874 for window in budget_limits:
875 window_dict = _coerce_window(window)
876 budget_duration = window_dict.get("budget_duration")
877 max_budget = _to_float(window_dict.get("max_budget"))
878 if not budget_duration or max_budget is None or max_budget <= 0:
879 continue
880 window_start = get_budget_window_start(window_dict)
881 if window_start is None:
882 verbose_proxy_logger.warning(
883 "Skipping budget window with invalid duration for %s=%s: %s",
884 entity_type,
885 entity_id,
886 budget_duration,
887 )
888 continue
889 counters.append(
890 _BudgetCounter(
891 counter_key=f"{entity_prefix}:window:{budget_duration}",
892 max_budget=float(max_budget),
893 fallback_spend=0.0,
894 entity_type=entity_type,
895 entity_id=f"{entity_id}:{budget_duration}",
896 spend_log_entity_id=entity_id,
897 window_duration=str(budget_duration),
898 window_start=window_start,
899 )
900 )
901 return counters
904def _coerce_window(window: object) -> Mapping[str, object]:
905 if isinstance(window, Mapping):
906 return window
907 if isinstance(window, str):
908 try:
909 parsed: Final[object] = json.loads(window)
910 except Exception:
911 return {}
912 return parsed if isinstance(parsed, Mapping) else {}
913 model_dump: Final = getattr(window, "model_dump", None)
914 if not callable(model_dump):
915 return {}
916 dumped: Final[object] = model_dump()
917 return dumped if isinstance(dumped, Mapping) else {}
920async def _reserve_counter(
921 counter: _BudgetCounter,
922 reservation_cost: float,
923) -> float | None:
924 from litellm.proxy.proxy_server import (
925 _ensure_spend_counter_initialized,
926 _ensure_window_spend_counter_initialized,
927 _increment_spend_counter_cache,
928 _invalidate_spend_counter,
929 )
931 attempted_increment = False
932 try:
933 if counter.source_cache_key is not None:
934 await _ensure_spend_counter_initialized(
935 counter_key=counter.counter_key,
936 source_cache_key=counter.source_cache_key,
937 )
938 elif counter.spend_log_entity_id is not None and counter.window_start is not None:
939 initialized: Final = await _ensure_window_spend_counter_initialized(
940 counter_key=counter.counter_key,
941 entity_type=counter.entity_type,
942 entity_id=counter.spend_log_entity_id,
943 window_duration=counter.window_duration,
944 window_start=counter.window_start,
945 )
946 if initialized is False:
947 verbose_proxy_logger.warning(
948 "Skipping budget reservation for %s because window spend could not be loaded",
949 counter.counter_key,
950 )
951 raise _CounterReservationUnavailable
953 attempted_increment = True
954 reserved_value: Final = await _increment_spend_counter_cache(
955 counter_key=counter.counter_key,
956 increment=reservation_cost,
957 )
958 return float(reserved_value) if reserved_value is not None else None
959 except _CounterReservationUnavailable:
960 raise
961 except Exception:
962 verbose_proxy_logger.warning(
963 "Skipping budget reservation for %s because spend counter reservation failed",
964 counter.counter_key,
965 exc_info=True,
966 )
967 counter_invalidated = False
968 try:
969 await _invalidate_spend_counter(counter_key=counter.counter_key)
970 counter_invalidated = True
971 except Exception:
972 verbose_proxy_logger.warning(
973 "Failed to invalidate spend counter after budget reservation failure for %s",
974 counter.counter_key,
975 exc_info=True,
976 )
977 raise _CounterReservationUnavailable(
978 touched_counter=attempted_increment,
979 counter_invalidated=counter_invalidated,
980 )
983async def _get_current_counter_value(counter: _BudgetCounter) -> float:
984 from litellm.proxy.proxy_server import get_current_spend
986 return await get_current_spend(
987 counter_key=counter.counter_key,
988 fallback_spend=counter.fallback_spend,
989 )
992def _counters_batch_scope(counter_keys: frozenset[str]) -> spend_counter_batch_scope:
993 """Each counter is read once, then written, so one MGET up front serves every read in the loop."""
994 from litellm.proxy.proxy_server import spend_counter_cache
996 return spend_counter_batch_scope(spend_counter_cache.redis_cache, counter_keys=counter_keys)
999@dataclass(frozen=True, slots=True)
1000class _EntryAdjustment:
1001 entry: dict[str, float | str]
1002 counter_key: str
1003 target_adjustment: float
1004 adjustment: float
1007def _entry_adjustment(
1008 entry: dict[str, float | str], actual_cost: float, default_reserved_cost: float
1009) -> _EntryAdjustment | None:
1010 counter_key: Final = entry.get("counter_key")
1011 if counter_key is None:
1012 return None
1013 target_adjustment: Final = actual_cost - _get_entry_reserved_cost(
1014 entry=entry, default_reserved_cost=default_reserved_cost
1015 )
1016 adjustment: Final = target_adjustment - float(entry.get("applied_adjustment") or 0.0)
1017 if adjustment == 0:
1018 return None
1019 return _EntryAdjustment(
1020 entry=entry, counter_key=str(counter_key), target_adjustment=target_adjustment, adjustment=adjustment
1021 )
1024async def _set_reserved_entries_actual_cost(
1025 entries: list[dict],
1026 actual_cost: float,
1027 default_reserved_cost: float,
1028 reseed_on_inconsistent: bool = True,
1029) -> None:
1030 """Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline.
1031 A counter that was flushed or reseeded since reservation is settled on its own after the pipeline."""
1032 from litellm.proxy.proxy_server import increment_spend_counters_pipeline
1034 with _counters_batch_scope(frozenset(str(entry["counter_key"]) for entry in entries if "counter_key" in entry)):
1035 adjustments: Final = tuple(
1036 adjustment
1037 for entry in entries
1038 if (adjustment := _entry_adjustment(entry, actual_cost, default_reserved_cost)) is not None
1039 )
1040 consistent: Final = tuple(
1041 await asyncio.gather(
1042 *(
1043 _counter_can_apply_adjustment(counter_key=item.counter_key, adjustment=item.adjustment)
1044 for item in adjustments
1045 )
1046 )
1047 )
1048 inconsistent: Final = tuple(item for item, ok in zip(adjustments, consistent) if not ok)
1049 if inconsistent and not reseed_on_inconsistent:
1050 # Pre-call admission resize: the in-flight reservation cost is not yet
1051 # persisted, so the DB floor would discard it. Keep the original
1052 # fail-closed behavior (raise -> reserve_budget_for_request releases and
1053 # denies) rather than admitting against an inconsistent counter.
1054 raise RuntimeError(
1055 f"Cannot resize budget reservation against inconsistent counter {inconsistent[0].counter_key}"
1056 )
1057 applicable: Final = tuple(item for item, ok in zip(adjustments, consistent) if ok)
1058 await increment_spend_counters_pipeline(
1059 pending=tuple(
1060 PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable
1061 )
1062 )
1063 for item in inconsistent:
1064 await _reseed_reserved_entry(item=item, actual_cost=actual_cost)
1065 for item in adjustments:
1066 item.entry["applied_adjustment"] = item.target_adjustment
1069async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> None:
1070 """Post-call reconcile / release of a counter that was flushed, expired or reseeded between reservation and
1071 reconcile: the optimistic delta no longer applies, so reseed from the DB floor and add the settled cost, since
1072 increment_spend_counters skips reserved keys. The reconcile runs before this request's spend is enqueued to the
1073 DB, so the reseeded floor excludes it."""
1074 from litellm.proxy.proxy_server import _increment_spend_counter_cache, reseed_spend_counter_from_db
1076 reseeded: Final = await reseed_spend_counter_from_db(counter_key=item.counter_key)
1077 if reseeded and actual_cost > 0:
1078 await _increment_spend_counter_cache(counter_key=item.counter_key, increment=actual_cost)
1081async def _counter_can_apply_adjustment(
1082 counter_key: str,
1083 adjustment: float,
1084) -> bool:
1085 from litellm.proxy.proxy_server import read_spend_counter_cache_value
1087 try:
1088 current_value, _ = await read_spend_counter_cache_value(counter_key=counter_key)
1089 except (TypeError, ValueError):
1090 return False
1091 if current_value is None:
1092 return False
1094 return not (adjustment < 0 and current_value + adjustment < -1e-12)
1097async def _release_applied_entries_best_effort(
1098 entries: list[dict],
1099 default_reserved_cost: float,
1100) -> None:
1101 for entry in entries:
1102 try:
1103 await _set_reserved_entries_actual_cost(
1104 entries=[entry], # mutable-ok: the reconcile takes the reservation's list of entries
1105 actual_cost=0.0,
1106 default_reserved_cost=default_reserved_cost,
1107 )
1108 except Exception:
1109 counter_key = entry.get("counter_key")
1110 verbose_proxy_logger.exception("Failed to release partial budget reservation during exception cleanup")
1111 if counter_key is None:
1112 continue
1113 try:
1114 from litellm.proxy.proxy_server import _invalidate_spend_counter
1116 await _invalidate_spend_counter(counter_key=counter_key)
1117 except Exception:
1118 verbose_proxy_logger.exception(
1119 "Failed to invalidate partial budget reservation counter during exception cleanup"
1120 )
1123async def _resize_applied_reservation(
1124 entries: list[dict],
1125 current_reserved_cost: float,
1126 new_reserved_cost: float,
1127) -> None:
1128 await _set_reserved_entries_actual_cost(
1129 entries=entries,
1130 actual_cost=new_reserved_cost,
1131 default_reserved_cost=current_reserved_cost,
1132 reseed_on_inconsistent=False,
1133 )
1134 for entry in entries:
1135 entry["reserved_cost"] = new_reserved_cost
1136 entry["applied_adjustment"] = 0.0
1139def _counter_to_reservation_entry(
1140 counter: _BudgetCounter,
1141 reserved_cost: float,
1142) -> dict[str, float | str]:
1143 return {
1144 "counter_key": counter.counter_key,
1145 "entity_type": counter.entity_type,
1146 "entity_id": counter.entity_id,
1147 "reserved_cost": reserved_cost,
1148 "applied_adjustment": 0.0,
1149 }
1152def _get_entry_reserved_cost(entry: dict, default_reserved_cost: float) -> float:
1153 try:
1154 return float(entry.get("reserved_cost", default_reserved_cost) or 0.0)
1155 except (TypeError, ValueError):
1156 return default_reserved_cost
1159def get_budget_window_start(window: object) -> datetime | None:
1160 window_dict: Final = _coerce_window(window)
1161 budget_duration: Final = window_dict.get("budget_duration")
1162 if budget_duration is None:
1163 return None
1164 try:
1165 duration_seconds: Final = duration_in_seconds(str(budget_duration))
1166 except Exception:
1167 return None
1169 reset_at = _coerce_datetime(window_dict.get("reset_at"))
1170 if reset_at is None:
1171 return datetime.now(timezone.utc) - timedelta(seconds=duration_seconds)
1172 if reset_at.tzinfo is None:
1173 reset_at = reset_at.replace(tzinfo=timezone.utc)
1174 return reset_at - timedelta(seconds=duration_seconds)
1177def _coerce_datetime(value: object) -> datetime | None:
1178 if value is None:
1179 return None
1180 if isinstance(value, datetime):
1181 return value
1182 if isinstance(value, str):
1183 try:
1184 return datetime.fromisoformat(value.replace("Z", "+00:00"))
1185 except ValueError:
1186 return None
1187 return None
1190def estimate_request_max_cost(
1191 request_body: dict,
1192 route: str,
1193 llm_router: Router | None,
1194 input_token_counts: Mapping[str, int] | None = None,
1195) -> float | None:
1196 estimates = [
1197 _estimate_request_max_cost_for_model(
1198 request_body=request_body,
1199 route=route,
1200 model=model_name,
1201 llm_router=llm_router,
1202 input_tokens=(input_token_counts or {}).get(model_name),
1203 )
1204 for model_name in _get_request_models(request_body=request_body, route=route, llm_router=llm_router)
1205 ]
1206 estimates = [estimate for estimate in estimates if estimate is not None]
1207 if not estimates:
1208 return None
1209 return max(cast(list[float], estimates))
1212def estimate_request_input_cost(
1213 request_body: dict,
1214 route: str,
1215 llm_router: Router | None,
1216 input_token_counts: Mapping[str, int] | None = None,
1217) -> float | None:
1218 """Cost of the request's input tokens alone.
1220 Once the provider request is dispatched the input tokens are billed even if
1221 the client disconnects before the first chunk, so this is the cost floor a
1222 cancelled in-flight request has already incurred. A cancelled reservation is
1223 reconciled to this instead of being refunded to zero.
1224 """
1225 estimates = [
1226 _estimate_request_input_cost_for_model(
1227 request_body=request_body,
1228 route=route,
1229 model=model_name,
1230 llm_router=llm_router,
1231 input_tokens=(input_token_counts or {}).get(model_name),
1232 )
1233 for model_name in _get_request_models(request_body=request_body, route=route, llm_router=llm_router)
1234 ]
1235 estimates = [estimate for estimate in estimates if estimate is not None]
1236 if not estimates:
1237 return None
1238 return max(cast("list[float]", estimates))
1241def _estimate_request_input_cost_for_model(
1242 request_body: dict,
1243 route: str,
1244 model: str,
1245 llm_router: Router | None,
1246 input_tokens: int | None = None,
1247) -> float | None:
1248 estimates: Final = [
1249 _input_cost_for_cost_info(
1250 request_body=request_body,
1251 route=route,
1252 model=model,
1253 model_info=model_info,
1254 input_tokens=input_tokens,
1255 )
1256 for model_info in _get_model_cost_infos(model=model, llm_router=llm_router)
1257 ]
1258 valid_estimates: Final = [estimate for estimate in estimates if estimate is not None]
1259 return max(valid_estimates) if valid_estimates else None
1262def _input_cost_for_cost_info(
1263 request_body: dict,
1264 route: str,
1265 model: str,
1266 model_info: Mapping[str, object],
1267 input_tokens: int | None = None,
1268) -> float | None:
1269 estimated_input_tokens: Final = _estimate_input_tokens(
1270 request_body=request_body,
1271 route=route,
1272 model=model,
1273 model_info=model_info,
1274 input_tokens=input_tokens,
1275 )
1276 if estimated_input_tokens is None:
1277 return None
1278 tiered_pricing: Final = model_info.get("tiered_pricing")
1279 if isinstance(tiered_pricing, list) and tiered_pricing:
1280 tier: Final = select_tier_for_input(tiered_pricing=tiered_pricing, input_tokens=estimated_input_tokens)
1281 if tier is not None:
1282 return estimated_input_tokens * tier_rate(tier, "input_cost_per_token")
1283 input_cost_per_token: Final = _to_float(model_info.get("input_cost_per_token"))
1284 if input_cost_per_token is None:
1285 return None
1286 return estimated_input_tokens * input_cost_per_token
1289def _estimate_request_max_cost_for_model(
1290 request_body: dict,
1291 route: str,
1292 model: str,
1293 llm_router: Router | None,
1294 input_tokens: int | None = None,
1295) -> float | None:
1296 estimates: Final = [
1297 _max_cost_for_cost_info(
1298 request_body=request_body,
1299 route=route,
1300 model=model,
1301 model_info=model_info,
1302 input_tokens=input_tokens,
1303 )
1304 for model_info in _get_model_cost_infos(model=model, llm_router=llm_router)
1305 ]
1306 valid_estimates: Final = [estimate for estimate in estimates if estimate is not None]
1307 return max(valid_estimates) if valid_estimates else None
1310def _max_cost_for_cost_info(
1311 request_body: dict,
1312 route: str,
1313 model: str,
1314 model_info: Mapping[str, object],
1315 input_tokens: int | None = None,
1316) -> float | None:
1317 image_cost: Final = _estimate_image_generation_cost(
1318 request_body=request_body,
1319 model_info=model_info,
1320 )
1321 if image_cost is not None:
1322 return image_cost
1324 estimated_input_tokens: Final = _estimate_input_tokens(
1325 request_body=request_body,
1326 route=route,
1327 model=model,
1328 model_info=model_info,
1329 input_tokens=input_tokens,
1330 )
1331 output_tokens: Final = _estimate_output_tokens(
1332 request_body=request_body,
1333 route=route,
1334 model_info=model_info,
1335 )
1336 if estimated_input_tokens is None or output_tokens is None:
1337 return None
1339 output_multiplier: Final = _get_output_multiplier(request_body=request_body)
1340 tiered_pricing: Final = model_info.get("tiered_pricing")
1341 if isinstance(tiered_pricing, list) and tiered_pricing:
1342 tier: Final = select_tier_for_input(tiered_pricing=tiered_pricing, input_tokens=estimated_input_tokens)
1343 if tier is not None:
1344 output_rate = max(
1345 tier_rate(tier, "output_cost_per_token"),
1346 tier_rate(tier, "output_cost_per_reasoning_token"),
1347 )
1348 return (estimated_input_tokens * tier_rate(tier, "input_cost_per_token")) + (
1349 output_tokens * output_multiplier * output_rate
1350 )
1352 input_cost_per_token: Final = _to_float(model_info.get("input_cost_per_token"))
1353 output_cost_per_token: Final = _to_float(model_info.get("output_cost_per_token"))
1354 output_cost_per_reasoning_token: Final = _to_float(model_info.get("output_cost_per_reasoning_token"))
1355 cost = 0.0
1356 if input_cost_per_token is not None:
1357 cost += estimated_input_tokens * input_cost_per_token
1358 elif estimated_input_tokens > 0:
1359 return None
1361 # The reasoning-token share is unknown before the request runs, so reserve every
1362 # output token at the higher of the standard and reasoning rates to avoid
1363 # under-reserving reasoning-heavy requests.
1364 output_rate = max(output_cost_per_token or 0.0, output_cost_per_reasoning_token or 0.0)
1365 if output_cost_per_token is not None or output_cost_per_reasoning_token is not None:
1366 cost += output_tokens * output_multiplier * output_rate
1367 elif output_tokens > 0:
1368 return None
1370 return cost
1373def _estimate_image_generation_cost(
1374 request_body: dict,
1375 model_info: Mapping[str, object],
1376) -> float | None:
1377 """
1378 Reserve `n × per-image cost` for image-generation requests so concurrent
1379 requests against a depleted budget cannot all slip past the admission gate
1380 onto the provider. Token-based pricing (e.g. gpt-image-1) is handled by
1381 the chat-route token path; per-pixel and size/quality-tiered pricing
1382 (DALL-E 2 size variants, premium tiers) are not handled here and fall
1383 through to read-time enforcement.
1385 The "output" vs "input" cost-per-image naming is inconsistent across
1386 providers — OpenAI's dall-e-3 entry uses ``input_cost_per_image`` while
1387 aiml/dall-e-3 uses ``output_cost_per_image`` — so both are summed.
1388 """
1389 # Gate strictly on `mode`. Several chat and embedding models carry
1390 # ``input_cost_per_image`` / ``output_cost_per_image`` to price multimodal
1391 # *vision input* (e.g. ``gemini-3.1-pro-preview``, ``azure/gpt-realtime-*``,
1392 # ``amazon.titan-embed-image-v1``). Falling back to "treat as image-gen if
1393 # an image cost field is present" would short-circuit the token-priced
1394 # path for those models and reserve a fraction of a cent instead of the
1395 # true per-token cost. All real image-generation entries in
1396 # ``model_prices_and_context_window.json`` carry ``mode: image_generation``
1397 # or ``mode: image_edit``, so the field-presence fallback is unnecessary.
1398 if model_info.get("mode") not in ("image_generation", "image_edit"):
1399 return None
1401 output_cost_per_image: Final = _to_float(model_info.get("output_cost_per_image"))
1402 input_cost_per_image: Final = _to_float(model_info.get("input_cost_per_image"))
1403 cost_per_image: Final = (output_cost_per_image or 0.0) + (input_cost_per_image or 0.0)
1404 if cost_per_image <= 0:
1405 return None
1407 n: Final = _to_int(request_body.get("n")) or 1
1408 return cost_per_image * max(n, 1)
1411def _get_model_cost_info(
1412 model: str,
1413 llm_router: Router | None,
1414) -> Mapping[str, object] | None:
1415 if llm_router is not None:
1416 model_group_info: Final = llm_router.cached_model_group_info(model)
1417 if model_group_info is not None:
1418 return model_group_info.model_dump()
1419 return dict(litellm.get_model_info(model=model))
1422def _get_model_cost_infos(
1423 model: str,
1424 llm_router: Router | None,
1425) -> Sequence[Mapping[str, object]]:
1426 """Cost-info candidates to estimate a request against for one model group.
1428 Reservation runs before routing, so the deployment that will serve the request
1429 is unknown. Rather than guess, we estimate the cost against every eligible
1430 pricing shape in the group (the group's flat rates plus each deployment's
1431 tiered table) and let the caller reserve the maximum, so a cheaper sibling
1432 deployment can never leave the request under-reserved.
1433 """
1434 try:
1435 base: Final = _get_model_cost_info(model=model, llm_router=llm_router)
1436 if base is None:
1437 return []
1438 tiered_tables: Final = _get_deployment_tiered_pricing_tables(model=model, llm_router=llm_router)
1439 except Exception:
1440 verbose_proxy_logger.debug(
1441 "Unable to load model cost info for budget reservation",
1442 exc_info=True,
1443 )
1444 return []
1445 if not tiered_tables:
1446 return [base]
1447 return [base, *({**base, "tiered_pricing": table} for table in tiered_tables)]
1450def _deployment_tiered_pricing_table(
1451 deployment: DeploymentTypedDict,
1452 llm_router: Router,
1453) -> Sequence[Mapping[str, object]] | None:
1454 model_id: Final = _get_value(_get_value(deployment, "model_info"), "id")
1455 backend_model: Final = _get_value(_get_value(deployment, "litellm_params"), "model")
1456 if not isinstance(model_id, str) or not isinstance(backend_model, str):
1457 return None
1458 deployment_model_info: Final = llm_router.cached_deployment_model_info(model_id, backend_model)
1459 if deployment_model_info is None:
1460 return None
1461 tiered_pricing: Final = deployment_model_info.get("tiered_pricing")
1462 if isinstance(tiered_pricing, list) and tiered_pricing:
1463 return tiered_pricing
1464 return None
1467def _get_deployment_tiered_pricing_tables(
1468 model: str,
1469 llm_router: Router | None,
1470) -> Sequence[Sequence[Mapping[str, object]]]:
1471 if llm_router is None:
1472 return []
1473 deployments: Final = llm_router.get_model_list(model_name=model) or []
1474 return [
1475 table
1476 for deployment in deployments
1477 if (table := _deployment_tiered_pricing_table(deployment, llm_router)) is not None
1478 ]
1481def _get_request_models(
1482 request_body: dict,
1483 route: str,
1484 llm_router: Router | None,
1485) -> Sequence[str]:
1486 model: Final = get_model_from_request(request_body, route, llm_router=llm_router)
1487 if model is None:
1488 return ()
1489 return (model,) if isinstance(model, str) else tuple(model)
1492async def count_request_input_tokens(
1493 request_body: dict,
1494 route: str,
1495 llm_router: Router | None,
1496 raw_body: bytes | None = None,
1497) -> Mapping[str, int]:
1498 """Input-token count per candidate model, counted once per request.
1500 The counts are reused by both the max-cost and the input-cost estimate."""
1501 models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router)
1502 if not models:
1503 return MappingProxyType({})
1504 return await count_input_tokens(request_body=request_body, raw_body=raw_body, models=models)
1507def _estimate_input_tokens(
1508 request_body: dict,
1509 route: str,
1510 model: str,
1511 model_info: Mapping[str, object],
1512 input_tokens: int | None = None,
1513) -> int | None:
1514 counted: Final = (
1515 input_tokens
1516 if input_tokens is not None
1517 else count_input_tokens_for_model(request_body=request_body, model=model)
1518 )
1519 if counted is not None:
1520 return counted
1522 max_input_tokens: Final = _to_int(model_info.get("max_input_tokens"))
1523 if max_input_tokens is not None:
1524 return max_input_tokens
1526 return None
1529DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK: Final = 16384
1532def _estimate_output_tokens(
1533 request_body: dict,
1534 route: str,
1535 model_info: Mapping[str, object],
1536) -> int | None:
1537 if _is_input_only_route(route=route):
1538 return 0
1540 requested: Final = _requested_output_tokens(request_body)
1542 # Clamp at min(requested-or-default, model_max-or-default). Two purposes:
1543 # (1) Without an explicit cap we still need a finite reservation so the
1544 # atomic admission counter actually bounds concurrent in-flight cost
1545 # (mirrors parallel_request_limiter_v3's DEFAULT_MAX_TOKENS_ESTIMATE).
1546 # (2) An adversarial caller cannot send max_tokens=999999999 to inflate
1547 # the reservation up to remaining team headroom and pin the counter
1548 # at the cap — the model can only physically emit max_output_tokens
1549 # anyway, so reserving more is both wasteful and a DoS surface.
1550 model_ceiling: Final = _to_int(model_info.get("max_output_tokens")) or DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK
1551 return min(DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK if requested is None else requested, model_ceiling)
1554_OUTPUT_TOKEN_FIELDS: Final = ("max_completion_tokens", "max_tokens", "max_output_tokens")
1557def _requested_output_tokens(request_body: Mapping[str, object]) -> int | None:
1558 inference_config: Final = request_body.get("inferenceConfig")
1559 candidates: Final = (
1560 *(request_body.get(field) for field in _OUTPUT_TOKEN_FIELDS),
1561 inference_config.get("maxTokens") if isinstance(inference_config, Mapping) else None,
1562 )
1563 return next((tokens for tokens in map(_to_int, candidates) if tokens is not None), None)
1566def _get_output_multiplier(request_body: dict) -> int:
1567 output_multiplier = 1
1568 for key in ("n", "best_of"):
1569 value = _to_int(request_body.get(key))
1570 if value is not None:
1571 output_multiplier = max(output_multiplier, value)
1572 return output_multiplier
1575def _is_input_only_route(route: str) -> bool:
1576 return any(
1577 route_part in route
1578 for route_part in (
1579 "embeddings",
1580 "rerank",
1581 "moderations",
1582 )
1583 )
1586def _to_float(value: object) -> float | None:
1587 if not isinstance(value, (SupportsFloat, SupportsIndex, str, bytes, bytearray)): 1587 ↛ 1589line 1587 didn't jump to line 1589 because the condition on line 1587 was always true
1588 return None
1589 try:
1590 return float(value)
1591 except (TypeError, ValueError):
1592 return None
1595def _to_int(value: object) -> int | None:
1596 if not isinstance(value, (SupportsInt, SupportsIndex, str, bytes, bytearray)):
1597 return None
1598 try:
1599 return int(value)
1600 except (TypeError, ValueError):
1601 return None
1604def _get_value(obj: object, key: str) -> object:
1605 if isinstance(obj, Mapping): 1605 ↛ 1606line 1605 didn't jump to line 1606 because the condition on line 1605 was never true
1606 return obj.get(key)
1607 return getattr(obj, key, None)