From 5a619cf87965db7196a3065e432ae46e69ff3b1c Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Wed, 29 Apr 2026 19:13:55 -0700 Subject: [PATCH 01/31] tighten budget spend admission --- litellm/proxy/_types.py | 1 + litellm/proxy/auth/user_api_key_auth.py | 65 +- .../proxy/hooks/proxy_track_cost_callback.py | 175 +++-- litellm/proxy/proxy_server.py | 109 ++- .../spend_tracking/budget_reservation.py | 673 ++++++++++++++++++ .../proxy/auth/test_user_api_key_auth.py | 21 +- .../hooks/test_proxy_track_cost_callback.py | 230 +++++- .../proxy/test_budget_reservation.py | 509 +++++++++++++ 8 files changed, 1678 insertions(+), 105 deletions(-) create mode 100644 litellm/proxy/spend_tracking/budget_reservation.py create mode 100644 tests/test_litellm/proxy/test_budget_reservation.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 92c920ca594..8be73f3427e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2567,6 +2567,7 @@ class UserAPIKeyAuth( user_spend: Optional[float] = None user_max_budget: Optional[float] = None request_route: Optional[str] = None + budget_reservation: Optional[Dict[str, Any]] = None user: Optional[Any] = None # Expanded user object when expand=user is used created_by_user: Optional[Any] = ( None # Expanded created_by user when expand=user is used diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index b8db3cd2a7b..7e70e7bb3fa 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1840,10 +1840,11 @@ async def _run_centralized_common_checks( user_api_key_auth_obj.project_metadata = project_object.metadata user_api_key_auth_obj.project_alias = project_object.project_alias - skip_budget_checks = False - model = get_model_from_request(request_data, route) - if model is not None and llm_router is not None: - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = _should_skip_budget_checks( + request_data=request_data, + route=route, + llm_router=llm_router, + ) _ = await common_checks( request=request, @@ -1861,6 +1862,19 @@ async def _run_centralized_common_checks( project_object=project_object, ) + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data=request_data, + route=route, + llm_router=llm_router, + team_object=team_object, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + skip_budget_checks=skip_budget_checks, + ) + async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary @@ -1868,6 +1882,49 @@ async def _noop_none() -> None: return None +async def _reserve_budget_after_common_checks( + user_api_key_auth_obj: UserAPIKeyAuth, + request_data: dict, + route: str, + llm_router: Optional[Any], + team_object: Optional[LiteLLM_TeamTableCachedObj], + user_object: Optional[LiteLLM_UserTable], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, + skip_budget_checks: bool, +) -> None: + if skip_budget_checks: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + reserve_budget_for_request, + ) + + user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request( + request_body=request_data, + route=route, + llm_router=llm_router, + valid_token=user_api_key_auth_obj, + team_object=team_object, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + +def _should_skip_budget_checks( + request_data: dict, + route: str, + llm_router: Optional[Any], +) -> bool: + model = get_model_from_request(request_data, route) + if model is not None and llm_router is not None: + return _is_model_cost_zero(model=model, llm_router=llm_router) + return False + + @tracer.wrap() async def user_api_key_auth( request: Request, diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index c9946f4e26f..71ad962d584 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -30,16 +30,20 @@ class _ProxyDBLogger(CustomLogger): kwargs, response_obj, start_time, end_time ) - async def async_post_call_failure_hook( - self, - request_data: dict, - original_exception: Exception, - user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, - ): - request_route = user_api_key_dict.request_route - if _ProxyDBLogger._should_track_errors_in_db() is False: - return + async def async_post_call_failure_hook( + self, + request_data: dict, + original_exception: Exception, + user_api_key_dict: UserAPIKeyAuth, + traceback_str: Optional[str] = None, + ): + await _release_budget_reservation( + budget_reservation=user_api_key_dict.budget_reservation + ) + + request_route = user_api_key_dict.request_route + if _ProxyDBLogger._should_track_errors_in_db() is False: + return elif request_route is not None and not ( RouteChecks.is_llm_api_route(route=request_route) or RouteChecks.is_info_route(route=request_route) @@ -155,12 +159,15 @@ class _ProxyDBLogger(CustomLogger): f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" ) parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) - litellm_params = kwargs.get("litellm_params", {}) or {} - end_user_id = get_end_user_id_for_cost_tracking(litellm_params) - metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) - user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) - team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) - org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) + litellm_params = kwargs.get("litellm_params", {}) or {} + end_user_id = get_end_user_id_for_cost_tracking(litellm_params) + metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) + budget_reservation = _get_budget_reservation_from_metadata( + metadata=metadata + ) + user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) + team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) + org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None)) end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) sl_object: Optional[StandardLoggingPayload] = kwargs.get( @@ -183,38 +190,31 @@ class _ProxyDBLogger(CustomLogger): f"Cache Hit: response_cost {response_cost}, for user_id {user_id}" ) - verbose_proxy_logger.debug( - f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" - ) - if _should_track_cost_callback( - user_api_key=user_api_key, + verbose_proxy_logger.debug( + f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" + ) + if _should_track_cost_callback( + user_api_key=user_api_key, user_id=user_id, team_id=team_id, - end_user_id=end_user_id, - ): - ## UPDATE DATABASE - await proxy_logging_obj.db_spend_update_writer.update_database( - token=user_api_key, - response_cost=response_cost, - user_id=user_id, - end_user_id=end_user_id, - team_id=team_id, - kwargs=kwargs, - completion_response=completion_response, - start_time=start_time, - end_time=end_time, - org_id=org_id, - ) - - # Atomically update spend counters (in-memory + Redis) - # for cross-pod budget enforcement. - await increment_spend_counters( - token=user_api_key, - team_id=team_id, - user_id=user_id, - response_cost=response_cost, - org_id=org_id, - ) + end_user_id=end_user_id, + ): + ## UPDATE DATABASE + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + org_id=org_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + budget_reservation=budget_reservation, + ) # update cache (fire-and-forget for backward compat: # cached object fields, soft budget alerts, etc.) @@ -234,10 +234,15 @@ class _ProxyDBLogger(CustomLogger): token=user_api_key, key_alias=key_alias, end_user_id=end_user_id, - response_cost=response_cost, - max_budget=end_user_max_budget, - ) + response_cost=response_cost, + max_budget=end_user_max_budget, + ) + elif budget_reservation is not None: + await _release_budget_reservation( + budget_reservation=budget_reservation + ) else: + await _release_budget_reservation(budget_reservation=budget_reservation) # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. # Use .get() for "stream" to avoid KeyError on health checks. if sl_object is None and not kwargs.get("model"): @@ -366,7 +371,7 @@ class _ProxyDBLogger(CustomLogger): return -def _should_track_cost_callback( +def _should_track_cost_callback( user_api_key: Optional[str], user_id: Optional[str], team_id: Optional[str], @@ -387,4 +392,72 @@ def _should_track_cost_callback( or end_user_id is not None ): return True - return False + return False + + +def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: + user_api_key_auth_obj = metadata.get("user_api_key_auth") + if user_api_key_auth_obj is None: + return None + return getattr(user_api_key_auth_obj, "budget_reservation", None) + + +async def _update_database_and_spend_counters( + proxy_logging_obj: Any, + increment_spend_counters: Any, + user_api_key: Optional[str], + user_id: Optional[str], + end_user_id: Optional[str], + team_id: Optional[str], + org_id: Optional[str], + kwargs: dict, + completion_response: Optional[Union[litellm.ModelResponse, Any]], + start_time: Any, + end_time: Any, + response_cost: float, + budget_reservation: Optional[dict], +) -> None: + try: + await proxy_logging_obj.db_spend_update_writer.update_database( + token=user_api_key, + response_cost=response_cost, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + org_id=org_id, + ) + except Exception: + if budget_reservation is not None: + await _release_budget_reservation(budget_reservation=budget_reservation) + raise + + try: + await increment_spend_counters( + token=user_api_key, + team_id=team_id, + user_id=user_id, + response_cost=response_cost, + org_id=org_id, + budget_reservation=budget_reservation, + ) + except Exception: + if budget_reservation is not None: + await _release_budget_reservation(budget_reservation=budget_reservation) + raise + + +async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + release_budget_reservation, + ) + + await release_budget_reservation( + budget_reservation=budget_reservation, + ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 12229955299..811bdfbc303 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1830,6 +1830,7 @@ async def increment_spend_counters( user_id: Optional[str], response_cost: Optional[float], org_id: Optional[str] = None, + budget_reservation: Optional[dict] = None, ): """ Atomically increment spend counters for budget enforcement. @@ -1841,6 +1842,21 @@ async def increment_spend_counters( Awaited (not create_task) in the cost callback, so the counter is updated before the next request's auth check runs. """ + reserved_counter_keys = set() + if budget_reservation is not None: + from litellm.proxy.spend_tracking.budget_reservation import ( + get_reserved_counter_keys, + reconcile_budget_reservation, + ) + + reserved_counter_keys = get_reserved_counter_keys( + budget_reservation=budget_reservation + ) + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=response_cost or 0.0, + ) + if response_cost is None or response_cost == 0: return @@ -1856,11 +1872,13 @@ async def increment_spend_counters( if isinstance(token, str) and token.startswith("sk-") else token ) - await _init_and_increment_spend_counter( - counter_key=f"spend:key:{hashed_token}", - source_cache_key=hashed_token, - increment=response_cost, - ) + key_counter_key = f"spend:key:{hashed_token}" + if key_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=key_counter_key, + source_cache_key=hashed_token, + increment=response_cost, + ) # Increment per-window budget counters for multi-budget keys key_obj = await user_api_key_cache.async_get_cache(key=hashed_token) @@ -1877,17 +1895,21 @@ async def increment_spend_counters( if isinstance(window, dict) else window.budget_duration ) - await spend_counter_cache.async_increment_cache( - key=f"spend:key:{hashed_token}:window:{duration}", - value=response_cost, - ) + key_window_counter = f"spend:key:{hashed_token}:window:{duration}" + if key_window_counter not in reserved_counter_keys: + await spend_counter_cache.async_increment_cache( + key=key_window_counter, + value=response_cost, + ) if team_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:team:{team_id}", - source_cache_key=f"team_id:{team_id}", - increment=response_cost, - ) + team_counter_key = f"spend:team:{team_id}" + if team_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=team_counter_key, + source_cache_key=f"team_id:{team_id}", + increment=response_cost, + ) # Increment per-window budget counters for multi-budget teams team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{team_id}") @@ -1904,31 +1926,39 @@ async def increment_spend_counters( if isinstance(window, dict) else window.budget_duration ) - await spend_counter_cache.async_increment_cache( - key=f"spend:team:{team_id}:window:{duration}", - value=response_cost, - ) + team_window_counter = f"spend:team:{team_id}:window:{duration}" + if team_window_counter not in reserved_counter_keys: + await spend_counter_cache.async_increment_cache( + key=team_window_counter, + value=response_cost, + ) if user_id is not None and team_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:team_member:{user_id}:{team_id}", - source_cache_key=f"team_membership:{user_id}:{team_id}", - increment=response_cost, - ) + team_member_counter_key = f"spend:team_member:{user_id}:{team_id}" + if team_member_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=team_member_counter_key, + source_cache_key=f"team_membership:{user_id}:{team_id}", + increment=response_cost, + ) if user_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:user:{user_id}", - source_cache_key=user_id, - increment=response_cost, - ) + user_counter_key = f"spend:user:{user_id}" + if user_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=user_counter_key, + source_cache_key=user_id, + increment=response_cost, + ) if org_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:org:{org_id}", - source_cache_key=f"org_id:{org_id}", - increment=response_cost, - ) + org_counter_key = f"spend:org:{org_id}" + if org_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=org_counter_key, + source_cache_key=f"org_id:{org_id}:with_budget", + increment=response_cost, + ) async def _init_and_increment_spend_counter( @@ -1952,6 +1982,17 @@ async def _init_and_increment_spend_counter( under-counting (would allow overspend). 4. Increment atomically (both in-memory + Redis) """ + await _ensure_spend_counter_initialized( + counter_key=counter_key, + source_cache_key=source_cache_key, + ) + await spend_counter_cache.async_increment_cache(key=counter_key, value=increment) + + +async def _ensure_spend_counter_initialized( + counter_key: str, + source_cache_key: str, +): current = await spend_counter_cache.async_get_cache(key=counter_key) if current is None: # Shares the per-counter lock with get_current_spend. @@ -1974,8 +2015,6 @@ async def _init_and_increment_spend_counter( key=counter_key, value=base_spend ) - await spend_counter_cache.async_increment_cache(key=counter_key, value=increment) - async def update_cache( # noqa: PLR0915 token: Optional[str], diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py new file mode 100644 index 00000000000..cccf1745e80 --- /dev/null +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -0,0 +1,673 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Sequence, cast + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.proxy._types import ( + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + LiteLLM_UserTable, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_utils import get_model_from_request +from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.router import Router + + +@dataclass +class _BudgetCounter: + counter_key: str + max_budget: float + fallback_spend: float + entity_type: str + entity_id: str + source_cache_key: Optional[str] = None + + +def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: + if not budget_reservation: + return set() + entries = budget_reservation.get("entries") or [] + return { + entry["counter_key"] + for entry in entries + if isinstance(entry, dict) and entry.get("counter_key") is not None + } + + +async def reserve_budget_for_request( + request_body: dict, + route: str, + llm_router: Optional[Router], + valid_token: Optional[UserAPIKeyAuth], + team_object: Optional[LiteLLM_TeamTable], + user_object: Optional[LiteLLM_UserTable], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, +) -> Optional[dict]: + if valid_token is None or not RouteChecks.is_llm_api_route(route=route): + return None + if route in {"/models", "/v1/models", "/utils/token_counter"}: + return None + if get_model_from_request(request_body, route) is None: + return None + + counters = await _get_budget_counters( + valid_token=valid_token, + team_object=team_object, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if not counters: + return None + + current_spend_by_counter_key: Dict[str, float] = {} + reservation_cost = estimate_request_max_cost( + request_body=request_body, + route=route, + llm_router=llm_router, + ) + if reservation_cost is None: + reservation_cost = await _get_smallest_remaining_budget( + counters=counters, + current_spend_by_counter_key=current_spend_by_counter_key, + ) + if reservation_cost is None or reservation_cost <= 0: + return None + + applied_entries: List[Dict[str, Any]] = [] + try: + for counter in counters: + entry = _counter_to_reservation_entry(counter) + reserved_value = await _reserve_counter( + counter=counter, + reservation_cost=reservation_cost, + ) + applied_entries.append(entry) + + if reserved_value is not None: + current_spend = reserved_value + else: + cached_spend = current_spend_by_counter_key.get(counter.counter_key) + if cached_spend is None: + cached_spend = await _get_current_counter_value(counter=counter) + current_spend = cached_spend + reservation_cost + if current_spend > counter.max_budget: + raise litellm.BudgetExceededError( + current_cost=current_spend, + max_budget=counter.max_budget, + message=( + "Budget has been exceeded! " + f"{counter.entity_type}={counter.entity_id} " + f"Current cost: {current_spend}, " + f"Max budget: {counter.max_budget}" + ), + ) + except Exception: + await _set_reserved_entries_adjustment( + entries=applied_entries, + target_adjustment=-reservation_cost, + ) + raise + + return { + "reserved_cost": reservation_cost, + "entries": applied_entries, + "finalized": False, + } + + +async def reconcile_budget_reservation( + budget_reservation: Optional[dict], + actual_cost: Optional[float], +) -> None: + if not budget_reservation or budget_reservation.get("finalized") is True: + return + + reserved_cost = float(budget_reservation.get("reserved_cost") or 0.0) + actual = float(actual_cost or 0.0) + adjustment = actual - reserved_cost + await _set_reserved_entries_adjustment( + entries=budget_reservation.get("entries") or [], + target_adjustment=adjustment, + ) + budget_reservation["finalized"] = True + + +async def release_budget_reservation(budget_reservation: Optional[dict]) -> None: + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=0.0, + ) + + +async def _get_budget_counters( + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], + user_object: Optional[LiteLLM_UserTable], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, +) -> List[_BudgetCounter]: + counters: List[_BudgetCounter] = [] + + if valid_token.token is not None: + if valid_token.max_budget is not None and valid_token.max_budget > 0: + counters.append( + _BudgetCounter( + counter_key=f"spend:key:{valid_token.token}", + source_cache_key=valid_token.token, + max_budget=float(valid_token.max_budget), + fallback_spend=float(valid_token.spend or 0.0), + entity_type="Key", + entity_id=valid_token.token, + ) + ) + counters.extend( + _get_budget_limit_counters( + entity_prefix=f"spend:key:{valid_token.token}", + entity_type="Key", + entity_id=valid_token.token, + budget_limits=valid_token.budget_limits, + ) + ) + + if team_object is not None and team_object.team_id is not None: + team_id = team_object.team_id + if team_object.max_budget is not None and team_object.max_budget > 0: + counters.append( + _BudgetCounter( + counter_key=f"spend:team:{team_id}", + source_cache_key=f"team_id:{team_id}", + max_budget=float(team_object.max_budget), + fallback_spend=float(team_object.spend or 0.0), + entity_type="Team", + entity_id=team_id, + ) + ) + counters.extend( + _get_budget_limit_counters( + entity_prefix=f"spend:team:{team_id}", + entity_type="Team", + entity_id=team_id, + budget_limits=team_object.budget_limits, + ) + ) + + if ( + (team_object is None or team_object.team_id is None) + and user_object is not None + and user_object.user_id is not None + and user_object.max_budget is not None + and user_object.max_budget > 0 + ): + counters.append( + _BudgetCounter( + counter_key=f"spend:user:{user_object.user_id}", + source_cache_key=user_object.user_id, + max_budget=float(user_object.max_budget), + fallback_spend=float(user_object.spend or 0.0), + entity_type="User", + entity_id=user_object.user_id, + ) + ) + + team_member_counter = await _get_team_member_budget_counter( + valid_token=valid_token, + team_object=team_object, + user_object=user_object, + user_api_key_cache=user_api_key_cache, + ) + if team_member_counter is not None: + counters.append(team_member_counter) + + org_counter = await _get_org_budget_counter( + valid_token=valid_token, + team_object=team_object, + user_api_key_cache=user_api_key_cache, + ) + if org_counter is not None: + counters.append(org_counter) + + return counters + + +async def _get_team_member_budget_counter( + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], + user_object: Optional[LiteLLM_UserTable], + user_api_key_cache: DualCache, +) -> Optional[_BudgetCounter]: + if ( + team_object is None + or team_object.team_id is None + or user_object is None + or valid_token.user_id is None + ): + return None + + membership_cache_key = ( + f"team_membership:{valid_token.user_id}:{team_object.team_id}" + ) + cached_team_membership = await user_api_key_cache.async_get_cache( + key=membership_cache_key + ) + team_membership: Optional[LiteLLM_TeamMembership] = None + if isinstance(cached_team_membership, LiteLLM_TeamMembership): + team_membership = cached_team_membership + elif isinstance(cached_team_membership, dict): + team_membership = LiteLLM_TeamMembership(**cached_team_membership) + + team_member_budget: Optional[float] = None + if team_membership is not None and team_membership.litellm_budget_table is not None: + team_member_budget = team_membership.litellm_budget_table.max_budget + else: + default_budget_id = (team_object.metadata or {}).get("team_member_budget_id") + if isinstance(default_budget_id, str): + default_budget = await user_api_key_cache.async_get_cache( + key=f"team_member_default_budget:{default_budget_id}", + ) + team_member_budget = _to_float(_get_value(default_budget, "max_budget")) + + if team_member_budget is None or team_member_budget <= 0: + return None + + team_member_spend = ( + cast(LiteLLM_TeamMembership, team_membership).spend + if team_membership is not None + else 0.0 + ) + return _BudgetCounter( + counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}", + source_cache_key=membership_cache_key, + max_budget=float(team_member_budget), + fallback_spend=float(team_member_spend or 0.0), + entity_type="TeamMember", + entity_id=f"{valid_token.user_id}:{team_object.team_id}", + ) + + +async def _get_org_budget_counter( + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], + user_api_key_cache: DualCache, +) -> Optional[_BudgetCounter]: + org_id: Optional[str] = None + if valid_token.org_id is not None: + org_id = valid_token.org_id + elif team_object is not None and team_object.organization_id is not None: + org_id = team_object.organization_id + if org_id is None: + return None + + org_table = await user_api_key_cache.async_get_cache( + key=f"org_id:{org_id}:with_budget", + ) + if org_table is None: + return None + + org_budget_table = _get_value(org_table, "litellm_budget_table") + if org_budget_table is None: + return None + + org_max_budget = _to_float(_get_value(org_budget_table, "max_budget")) + if org_max_budget is None or org_max_budget <= 0: + return None + + org_spend = _to_float(_get_value(org_table, "spend")) or 0.0 + return _BudgetCounter( + counter_key=f"spend:org:{org_id}", + source_cache_key=f"org_id:{org_id}:with_budget", + max_budget=org_max_budget, + fallback_spend=org_spend, + entity_type="Organization", + entity_id=org_id, + ) + + +def _get_budget_limit_counters( + entity_prefix: str, + entity_type: str, + entity_id: str, + budget_limits: Optional[Sequence[Any]], +) -> List[_BudgetCounter]: + counters: List[_BudgetCounter] = [] + if not budget_limits: + return counters + + for window in budget_limits: + window_dict = _coerce_window(window) + budget_duration = window_dict.get("budget_duration") + max_budget = window_dict.get("max_budget") + if not budget_duration or max_budget is None or max_budget <= 0: + continue + # Window counters intentionally have no source_cache_key: the DB stores + # window definitions/reset times, but not accumulated per-window spend. + counters.append( + _BudgetCounter( + counter_key=f"{entity_prefix}:window:{budget_duration}", + max_budget=float(max_budget), + fallback_spend=0.0, + entity_type=entity_type, + entity_id=f"{entity_id}:{budget_duration}", + ) + ) + return counters + + +def _coerce_window(window: Any) -> dict: + if isinstance(window, dict): + return window + if isinstance(window, str): + try: + parsed = json.loads(window) + return parsed if isinstance(parsed, dict) else {} + except Exception: + return {} + if hasattr(window, "model_dump"): + return window.model_dump() + return {} + + +async def _get_smallest_remaining_budget( + counters: List[_BudgetCounter], + current_spend_by_counter_key: Dict[str, float], +) -> Optional[float]: + remaining_budget: Optional[float] = None + for counter in counters: + current_spend = await _get_current_counter_value(counter=counter) + current_spend_by_counter_key[counter.counter_key] = current_spend + remaining = counter.max_budget - current_spend + if remaining <= 0: + raise litellm.BudgetExceededError( + current_cost=current_spend, + max_budget=counter.max_budget, + message=( + "Budget has been exceeded! " + f"{counter.entity_type}={counter.entity_id} " + f"Current cost: {current_spend}, " + f"Max budget: {counter.max_budget}" + ), + ) + remaining_budget = ( + remaining if remaining_budget is None else min(remaining_budget, remaining) + ) + return remaining_budget + + +async def _reserve_counter( + counter: _BudgetCounter, + reservation_cost: float, +) -> Optional[float]: + from litellm.proxy.proxy_server import ( + _ensure_spend_counter_initialized, + spend_counter_cache, + ) + + if counter.source_cache_key is not None: + await _ensure_spend_counter_initialized( + counter_key=counter.counter_key, + source_cache_key=counter.source_cache_key, + ) + + reserved_value = await spend_counter_cache.async_increment_cache( + key=counter.counter_key, + value=reservation_cost, + ) + return float(reserved_value) if reserved_value is not None else None + + +async def _get_current_counter_value(counter: _BudgetCounter) -> float: + from litellm.proxy.proxy_server import get_current_spend + + return await get_current_spend( + counter_key=counter.counter_key, + fallback_spend=counter.fallback_spend, + ) + + +async def _set_reserved_entries_adjustment( + entries: List[dict], + target_adjustment: float, +) -> None: + from litellm.proxy.proxy_server import spend_counter_cache + + for entry in entries: + counter_key = entry.get("counter_key") + if counter_key is None: + continue + applied_adjustment = float(entry.get("applied_adjustment") or 0.0) + adjustment = target_adjustment - applied_adjustment + if adjustment == 0: + continue + await spend_counter_cache.async_increment_cache( + key=counter_key, + value=adjustment, + ) + entry["applied_adjustment"] = target_adjustment + + +def _counter_to_reservation_entry(counter: _BudgetCounter) -> Dict[str, Any]: + return { + "counter_key": counter.counter_key, + "entity_type": counter.entity_type, + "entity_id": counter.entity_id, + "applied_adjustment": 0.0, + } + + +def estimate_request_max_cost( + request_body: dict, + route: str, + llm_router: Optional[Router], +) -> Optional[float]: + model = get_model_from_request(request_body, route) + if model is None: + return None + + models = [model] if isinstance(model, str) else model + estimates = [ + _estimate_request_max_cost_for_model( + request_body=request_body, + route=route, + model=model_name, + llm_router=llm_router, + ) + for model_name in models + ] + estimates = [estimate for estimate in estimates if estimate is not None] + if not estimates: + return None + return max(cast(List[float], estimates)) + + +def _estimate_request_max_cost_for_model( + request_body: dict, + route: str, + model: str, + llm_router: Optional[Router], +) -> Optional[float]: + model_info = _get_model_cost_info(model=model, llm_router=llm_router) + if model_info is None: + return None + + input_cost_per_token = _to_float(model_info.get("input_cost_per_token")) + output_cost_per_token = _to_float(model_info.get("output_cost_per_token")) + input_tokens = _estimate_input_tokens( + request_body=request_body, + route=route, + model=model, + model_info=model_info, + ) + output_tokens = _estimate_output_tokens( + request_body=request_body, + route=route, + model_info=model_info, + ) + if input_tokens is None or output_tokens is None: + return None + + cost = 0.0 + if input_cost_per_token is not None: + cost += input_tokens * input_cost_per_token + elif input_tokens > 0: + return None + + output_multiplier = _get_output_multiplier(request_body=request_body) + if output_cost_per_token is not None: + cost += output_tokens * output_multiplier * output_cost_per_token + elif output_tokens > 0: + return None + + return cost + + +def _get_model_cost_info( + model: str, + llm_router: Optional[Router], +) -> Optional[Dict[str, Any]]: + if llm_router is not None: + try: + model_group_info = llm_router.get_model_group_info(model_group=model) + if model_group_info is not None: + return model_group_info.model_dump() + except Exception: + verbose_proxy_logger.debug( + "Unable to load router model group info for budget reservation", + exc_info=True, + ) + + try: + return dict(litellm.get_model_info(model=model)) + except Exception: + return None + + +def _estimate_input_tokens( + request_body: dict, + route: str, + model: str, + model_info: Dict[str, Any], +) -> Optional[int]: + try: + if "messages" in request_body: + return litellm.token_counter( + model=model, + messages=request_body.get("messages") or [], + tools=request_body.get("tools"), + tool_choice=request_body.get("tool_choice"), + ) + if "prompt" in request_body: + return _count_text_tokens(model=model, text=request_body.get("prompt")) + if "input" in request_body: + return _count_text_tokens(model=model, text=request_body.get("input")) + if "query" in request_body or "documents" in request_body: + query_tokens = _count_text_tokens( + model=model, text=request_body.get("query") + ) + document_tokens = _count_text_tokens( + model=model, + text=request_body.get("documents"), + ) + return query_tokens + document_tokens + except Exception: + verbose_proxy_logger.debug( + "Unable to count input tokens for budget reservation", exc_info=True + ) + + max_input_tokens = _to_int(model_info.get("max_input_tokens")) + if max_input_tokens is not None: + return max_input_tokens + + return None + + +def _estimate_output_tokens( + request_body: dict, + route: str, + model_info: Dict[str, Any], +) -> Optional[int]: + if _is_input_only_route(route=route): + return 0 + + for key in ("max_completion_tokens", "max_tokens", "max_output_tokens"): + max_tokens = _to_int(request_body.get(key)) + if max_tokens is not None: + return max_tokens + + # If the caller did not cap output tokens, avoid reserving a model's + # theoretical maximum context. The caller can still admit one request by + # reserving the smallest remaining budget in reserve_budget_for_request(). + return None + + +def _count_text_tokens(model: str, text: Any) -> int: + if text is None: + return 0 + + token_count = 0 + stack = [text] + while stack: + item = stack.pop() + if item is None: + continue + if isinstance(item, list): + stack.extend(item) + continue + if isinstance(item, dict): + token_count += litellm.token_counter(model=model, text=json.dumps(item)) + continue + token_count += litellm.token_counter(model=model, text=str(item)) + return token_count + + +def _get_output_multiplier(request_body: dict) -> int: + output_multiplier = 1 + for key in ("n", "best_of"): + value = _to_int(request_body.get(key)) + if value is not None: + output_multiplier = max(output_multiplier, value) + return output_multiplier + + +def _is_input_only_route(route: str) -> bool: + return any( + route_part in route + for route_part in ( + "embeddings", + "rerank", + "moderations", + ) + ) + + +def _to_float(value: Any) -> Optional[float]: + if value is None: + return None + try: + return float(value) + except (TypeError, ValueError): + return None + + +def _to_int(value: Any) -> Optional[int]: + if value is None: + return None + try: + return int(value) + except (TypeError, ValueError): + return None + + +def _get_value(obj: Any, key: str) -> Any: + if isinstance(obj, dict): + return obj.get(key) + return getattr(obj, key, None) diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 9c43ebcbe79..23b1de881f7 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1,8 +1,6 @@ -import asyncio import json import os import sys -from typing import Tuple from unittest.mock import ANY, AsyncMock, MagicMock, patch sys.path.insert( @@ -1752,7 +1750,11 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): from starlette.datastructures import URL from starlette.requests import Request - from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy._types import ( + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + UserAPIKeyAuth, + ) from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder api_key = "sk-test-team-metadata-refresh" @@ -1833,16 +1835,17 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): request_data={}, ) - assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, ( - f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" - ) + assert result.team_metadata == { + "guardrails": ["test-guardrail-333"] + }, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" finally: for k, v in _originals.items(): setattr(_proxy_server_mod, k, v) - + + # --------------------------------------------------------------------------- - + # _run_centralized_common_checks — centralized authz gate # --------------------------------------------------------------------------- @@ -1859,7 +1862,7 @@ def _proxy_attrs_for_centralized_checks( """ return { "prisma_client": None, - "user_api_key_cache": MagicMock(), + "user_api_key_cache": DualCache(), "proxy_logging_obj": MagicMock(), "general_settings": ({"custom_auth_run_common_checks": True} if flag else {}), "llm_router": None, diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 65e7f744c85..dc052d8f012 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1,9 +1,7 @@ -import json import os import sys import pytest -from fastapi.testclient import TestClient sys.path.insert( 0, os.path.abspath("../../../..") @@ -13,8 +11,10 @@ from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger -from litellm.types.utils import StandardLoggingPayload +from litellm.proxy.hooks.proxy_track_cost_callback import ( + _ProxyDBLogger, + _update_database_and_spend_counters, +) @pytest.mark.asyncio @@ -62,7 +62,6 @@ async def test_async_post_call_failure_hook(): # Check the arguments passed to update_database call_args = mock_update_database.call_args[1] - print("call_args", json.dumps(call_args, indent=4, default=str)) assert call_args["token"] == "test_api_key" assert call_args["response_cost"] == 0.0 assert call_args["user_id"] == "test_user_id" @@ -128,6 +127,225 @@ async def test_async_post_call_failure_hook_non_llm_route(): mock_update_database.assert_not_called() +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_releases_budget_reservation_before_route_skip(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_api_key", + request_route="/custom/route", + budget_reservation=budget_reservation, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation, + patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + ): + await logger.async_post_call_failure_hook( + request_data={}, + original_exception=Exception("Test exception"), + user_api_key_dict=user_api_key_dict, + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + mock_update_database.assert_not_called() + + +@pytest.mark.asyncio +async def test_track_cost_callback_releases_budget_reservation_when_spend_tracking_skips(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "model": "gpt-4", + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "standard_logging_object": { + "response_cost": 0.1, + "request_tags": None, + }, + "stream": False, + } + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation: + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + +@pytest.mark.asyncio +async def test_track_cost_callback_releases_budget_reservation_when_response_cost_missing(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "model": "gpt-4", + "call_type": "acompletion", + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "standard_logging_object": { + "response_cost": None, + "response_cost_failure_debug_info": "missing custom price", + "request_tags": None, + }, + "stream": False, + } + + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + ) as mock_proxy_logging, + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( + side_effect=Exception("db unavailable") + ) + increment_spend_counters = AsyncMock() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation: + with pytest.raises(Exception, match="db unavailable"): + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + increment_spend_counters.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_updates_counters_after_db_update(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() + increment_spend_counters = AsyncMock() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + increment_spend_counters.assert_awaited_once_with( + token="test_api_key", + team_id="test_team_id", + user_id="test_user_id", + response_cost=0.2, + org_id="test_org_id", + budget_reservation=budget_reservation, + ) + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_releases_reservation_when_counter_update_fails(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() + increment_spend_counters = AsyncMock(side_effect=Exception("counter unavailable")) + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation: + with pytest.raises(Exception, match="counter unavailable"): + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + + @pytest.mark.asyncio async def test_track_cost_callback_skips_when_no_standard_logging_object(): """ @@ -344,7 +562,7 @@ async def test_enrich_failure_metadata_skips_when_no_api_key(): "user_api_key_team_id": None, "user_api_key_team_alias": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) mock_get_key.assert_not_called() diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py new file mode 100644 index 00000000000..9e154dcb842 --- /dev/null +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -0,0 +1,509 @@ +from unittest.mock import patch + +import pytest + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_OrganizationTable, + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + LiteLLM_UserTable, + UserAPIKeyAuth, +) +from litellm.proxy.spend_tracking.budget_reservation import ( + estimate_request_max_cost, + release_budget_reservation, + reserve_budget_for_request, +) +from litellm.proxy.utils import ProxyLogging + + +@pytest.fixture() +def spend_counter_state(): + import litellm.proxy.proxy_server as ps + + original_counter_cache = ps.spend_counter_cache + original_key_cache = ps.user_api_key_cache + original_prisma_client = ps.prisma_client + + counter_cache = DualCache() + key_cache = DualCache() + ps.spend_counter_cache = counter_cache + ps.user_api_key_cache = key_cache + ps.prisma_client = None + + try: + yield counter_cache, key_cache + finally: + ps.spend_counter_cache = original_counter_cache + ps.user_api_key_cache = original_key_cache + ps.prisma_client = original_prisma_client + + +def _request_body() -> dict: + return { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 10, + } + + +@pytest.mark.asyncio +async def test_should_prevent_second_key_reservation_over_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-race", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert reservation is not None + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-race") + == 0.6 + ) + + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-race") == 0.6 + ) + + await release_budget_reservation(reservation) + + +@pytest.mark.asyncio +async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_state): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-shared", + spend=0.0, + max_budget=1.0, + user_id="user-budget-shared", + team_id="team-budget-shared", + org_id="org-budget-shared", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-shared", + spend=0.0, + max_budget=1.0, + ) + user_object = LiteLLM_UserTable( + user_id="user-budget-shared", + spend=0.0, + ) + await key_cache.async_set_cache( + key="team_membership:user-budget-shared:team-budget-shared", + value=LiteLLM_TeamMembership( + user_id="user-budget-shared", + team_id="team-budget-shared", + spend=0.1, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="org_id:org-budget-shared:with_budget", + value=LiteLLM_OrganizationTable( + organization_id="org-budget-shared", + organization_alias="shared-org", + budget_id="org-budget-id", + spend=0.1, + models=[], + created_by="test", + updated_by="test", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.3, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=user_object, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:team_member:user-budget-shared:team-budget-shared" + ) == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:org:org-budget-shared" + ) == pytest.approx(0.4) + + await release_budget_reservation(reservation) + + +@pytest.mark.asyncio +async def test_should_seed_org_counter_from_with_budget_cache(spend_counter_state): + counter_cache, key_cache = spend_counter_state + await key_cache.async_set_cache( + key="org_id:org-counter-with-budget:with_budget", + value=LiteLLM_OrganizationTable( + organization_id="org-counter-with-budget", + organization_alias="shared-org", + budget_id="org-budget-id", + spend=2.0, + models=[], + created_by="test", + updated_by="test", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=10.0), + ).model_dump(), + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + org_id="org-counter-with-budget", + response_cost=0.25, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:org:org-counter-with-budget" + ) == pytest.approx(2.25) + + +@pytest.mark.asyncio +async def test_should_reserve_remaining_budget_when_output_cap_missing( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-uncapped", + spend=0.2, + max_budget=1.0, + ) + await key_cache.async_set_cache( + key="key-budget-uncapped", + value=valid_token, + ) + request_body = _request_body() + request_body.pop("max_tokens") + + with patch( + "litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info", + return_value={ + "input_cost_per_token": 0.0, + "output_cost_per_token": 100.0, + "max_output_tokens": 200000, + }, + ): + assert ( + estimate_request_max_cost( + request_body=request_body, + route="/chat/completions", + llm_router=None, + ) + is None + ) + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.8) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-uncapped" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + + +@pytest.mark.asyncio +async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( + spend_counter_state, + monkeypatch, +): + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-uncapped-read-once", + spend=0.2, + max_budget=1.0, + ) + + from litellm.proxy.spend_tracking import budget_reservation + + current_counter_reads = [] + + async def mock_get_current_counter_value(counter): + current_counter_reads.append(counter.counter_key) + return counter.fallback_spend + + async def mock_reserve_counter(counter, reservation_cost): + return None + + monkeypatch.setattr( + budget_reservation, + "_get_current_counter_value", + mock_get_current_counter_value, + ) + monkeypatch.setattr( + budget_reservation, + "_reserve_counter", + mock_reserve_counter, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=None, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.8) + assert current_counter_reads == ["spend:key:key-budget-uncapped-read-once"] + + +@pytest.mark.asyncio +async def test_should_reconcile_reserved_counter_to_actual_spend( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reconcile", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token="key-budget-reconcile", + team_id="team-without-budget", + user_id=None, + response_cost=0.2, + budget_reservation=reservation, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reconcile" + ) == pytest.approx(0.2) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-without-budget" + ) == pytest.approx(0.2) + + +@pytest.mark.asyncio +async def test_should_release_reservation_on_failure(spend_counter_state): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-release", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.4, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + await release_budget_reservation(reservation) + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-release" + ) == pytest.approx(0.0) + + +@pytest.mark.asyncio +async def test_should_retry_partial_release_without_double_decrement( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-partial-release", + spend=0.0, + max_budget=1.0, + team_id="team-budget-partial-release", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-partial-release", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.4, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + original_increment_cache = counter_cache.async_increment_cache + fail_next_team_release = True + + async def flaky_increment_cache(key, value, *args, **kwargs): + nonlocal fail_next_team_release + if ( + key == "spend:team:team-budget-partial-release" + and value < 0 + and fail_next_team_release + ): + fail_next_team_release = False + raise RuntimeError("simulated counter failure") + return await original_increment_cache(key=key, value=value, *args, **kwargs) + + monkeypatch.setattr(counter_cache, "async_increment_cache", flaky_increment_cache) + + with pytest.raises(RuntimeError): + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-partial-release" + ) == pytest.approx(0.0) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-partial-release" + ) == pytest.approx(0.4) + + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-partial-release" + ) == pytest.approx(0.0) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-partial-release" + ) == pytest.approx(0.0) + + +@pytest.mark.asyncio +async def test_should_reserve_all_budgeted_counters(spend_counter_state): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-all", + spend=0.0, + max_budget=1.0, + team_id="team-budget-all", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-all", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.3, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-all") == 0.3 + ) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:team-budget-all") == 0.3 + ) + + await release_budget_reservation(reservation) From 926de696a11bf60fce682c3e68933f7e81418855 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Wed, 29 Apr 2026 20:51:07 -0700 Subject: [PATCH 02/31] tighten budget counter cache recovery --- litellm/proxy/db/spend_counter_reseed.py | 132 +++++++++++++++++- litellm/proxy/proxy_server.py | 105 ++++++++++++-- .../spend_tracking/budget_reservation.py | 64 +++++++-- tests/test_litellm/proxy/test_proxy_server.py | 74 +++++++++- 4 files changed, 352 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index bf60a087c65..b4dc4a4d0a2 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -14,6 +14,7 @@ memory in long-lived deployments. import asyncio from collections import OrderedDict +from datetime import datetime from typing import TYPE_CHECKING, ClassVar, Optional from litellm._logging import verbose_proxy_logger @@ -118,6 +119,7 @@ class SpendCounterReseed: prisma_client: Optional["PrismaClient"], spend_counter_cache: "DualCache", counter_key: str, + require_cache_warm: bool = False, ) -> Optional[float]: """ Reseed a cold spend counter from the DB and warm the cache, @@ -148,12 +150,136 @@ class SpendCounterReseed: return None # Warm even when 0 so subsequent reads hit cache, not DB. try: - await spend_counter_cache.async_increment_cache( - key=counter_key, value=db_spend - ) + if require_cache_warm and spend_counter_cache.redis_cache is not None: + current_value = ( + await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=db_spend, + ) + ) + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + else: + await spend_counter_cache.async_increment_cache( + key=counter_key, value=db_spend + ) except Exception: verbose_proxy_logger.exception( "SpendCounterReseed.coalesced: failed to warm counter %s", counter_key, ) + if require_cache_warm: + raise return db_spend + + @staticmethod + async def window_from_spend_logs( + prisma_client: Optional["PrismaClient"], + entity_type: str, + entity_id: str, + window_start: datetime, + ) -> Optional[float]: + if prisma_client is None: + return None + + if entity_type == "Key": + group_field = "api_key" + where = { + "api_key": entity_id, + "startTime": {"gte": window_start}, + } + elif entity_type == "Team": + group_field = "team_id" + where = { + "team_id": entity_id, + "startTime": {"gte": window_start}, + } + else: + return None + + try: + response = await prisma_client.db.litellm_spendlogs.group_by( + by=[group_field], + where=where, # type: ignore[arg-type] + sum={"spend": True}, + ) + except Exception: + verbose_proxy_logger.exception( + "SpendCounterReseed.window_from_spend_logs: failed for %s=%s", + entity_type, + entity_id, + ) + return None + + if not response: + return 0.0 + first_row = response[0] + sum_row = ( + first_row.get("_sum") + if isinstance(first_row, dict) + else getattr(first_row, "_sum", None) + ) + spend = ( + sum_row.get("spend") + if isinstance(sum_row, dict) + else getattr(sum_row, "spend", None) + ) + return float(spend or 0.0) + + @staticmethod + async def coalesced_window( + prisma_client: Optional["PrismaClient"], + spend_counter_cache: "DualCache", + counter_key: str, + entity_type: str, + entity_id: str, + window_start: datetime, + ) -> Optional[float]: + lock = await SpendCounterReseed._get_lock(counter_key) + async with lock: + if spend_counter_cache.redis_cache is not None: + try: + val = await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key + ) + if val is not None: + return float(val) + except Exception: + pass + val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) + if val is not None: + return float(val) + + window_spend = await SpendCounterReseed.window_from_spend_logs( + prisma_client=prisma_client, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + if window_spend is None: + return None + try: + if spend_counter_cache.redis_cache is not None: + current_value = ( + await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=window_spend, + ) + ) + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + else: + await spend_counter_cache.async_increment_cache( + key=counter_key, value=window_spend + ) + except Exception: + verbose_proxy_logger.exception( + "SpendCounterReseed.coalesced_window: failed to warm counter %s", + counter_key, + ) + raise + return window_spend diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 811bdfbc303..f6655ba9285 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1897,9 +1897,16 @@ async def increment_spend_counters( ) key_window_counter = f"spend:key:{hashed_token}:window:{duration}" if key_window_counter not in reserved_counter_keys: - await spend_counter_cache.async_increment_cache( - key=key_window_counter, - value=response_cost, + from litellm.proxy.spend_tracking.budget_reservation import ( + get_budget_window_start, + ) + + await _init_and_increment_window_spend_counter( + counter_key=key_window_counter, + entity_type="Key", + entity_id=hashed_token, + window_start=get_budget_window_start(window), + increment=response_cost, ) if team_id is not None: @@ -1928,9 +1935,16 @@ async def increment_spend_counters( ) team_window_counter = f"spend:team:{team_id}:window:{duration}" if team_window_counter not in reserved_counter_keys: - await spend_counter_cache.async_increment_cache( - key=team_window_counter, - value=response_cost, + from litellm.proxy.spend_tracking.budget_reservation import ( + get_budget_window_start, + ) + + await _init_and_increment_window_spend_counter( + counter_key=team_window_counter, + entity_type="Team", + entity_id=team_id, + window_start=get_budget_window_start(window), + increment=response_cost, ) if user_id is not None and team_id is not None: @@ -1986,7 +2000,24 @@ async def _init_and_increment_spend_counter( counter_key=counter_key, source_cache_key=source_cache_key, ) - await spend_counter_cache.async_increment_cache(key=counter_key, value=increment) + await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + + +async def _init_and_increment_window_spend_counter( + counter_key: str, + entity_type: str, + entity_id: str, + window_start: Optional[datetime], + increment: float, +): + if window_start is not None: + await _ensure_window_spend_counter_initialized( + counter_key=counter_key, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) async def _ensure_spend_counter_initialized( @@ -2000,6 +2031,7 @@ async def _ensure_spend_counter_initialized( prisma_client=prisma_client, spend_counter_cache=spend_counter_cache, counter_key=counter_key, + require_cache_warm=True, ) if db_spend is None: # DB unavailable - fall back to in-process cache (may be stale). @@ -2011,11 +2043,66 @@ async def _ensure_spend_counter_initialized( else: base_spend = getattr(source, "spend", 0.0) or 0.0 if base_spend > 0: - await spend_counter_cache.async_increment_cache( - key=counter_key, value=base_spend + await _increment_spend_counter_cache( + counter_key=counter_key, increment=base_spend ) +async def _ensure_window_spend_counter_initialized( + counter_key: str, + entity_type: str, + entity_id: str, + window_start: datetime, +): + current = await spend_counter_cache.async_get_cache(key=counter_key) + if current is None: + window_spend = await SpendCounterReseed.coalesced_window( + prisma_client=prisma_client, + spend_counter_cache=spend_counter_cache, + counter_key=counter_key, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + if window_spend is None: + await _increment_spend_counter_cache(counter_key=counter_key, increment=0.0) + + +async def _increment_spend_counter_cache(counter_key: str, increment: float): + if spend_counter_cache.redis_cache is not None: + try: + current_value = await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=increment, + ) + except Exception: + await _invalidate_spend_counter(counter_key=counter_key) + raise + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + return current_value + + return await spend_counter_cache.async_increment_cache( + key=counter_key, + value=increment, + ) + + +async def _invalidate_spend_counter(counter_key: str): + spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) + if spend_counter_cache.redis_cache is not None: + try: + await spend_counter_cache.redis_cache.async_delete_cache(key=counter_key) + except Exception: + verbose_proxy_logger.debug( + "Unable to delete stale spend counter %s after increment failure", + counter_key, + exc_info=True, + ) + + async def update_cache( # noqa: PLR0915 token: Optional[str], user_id: Optional[str], diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index cccf1745e80..130a3e20275 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -2,11 +2,13 @@ from __future__ import annotations import json from dataclasses import dataclass +from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional, Sequence, cast import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache +from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import ( LiteLLM_TeamMembership, LiteLLM_TeamTable, @@ -27,6 +29,8 @@ class _BudgetCounter: entity_type: str entity_id: str source_cache_key: Optional[str] = None + spend_log_entity_id: Optional[str] = None + window_start: Optional[datetime] = None def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: @@ -349,8 +353,7 @@ def _get_budget_limit_counters( max_budget = window_dict.get("max_budget") if not budget_duration or max_budget is None or max_budget <= 0: continue - # Window counters intentionally have no source_cache_key: the DB stores - # window definitions/reset times, but not accumulated per-window spend. + window_start = get_budget_window_start(window_dict) counters.append( _BudgetCounter( counter_key=f"{entity_prefix}:window:{budget_duration}", @@ -358,6 +361,8 @@ def _get_budget_limit_counters( fallback_spend=0.0, entity_type=entity_type, entity_id=f"{entity_id}:{budget_duration}", + spend_log_entity_id=entity_id, + window_start=window_start, ) ) return counters @@ -409,7 +414,8 @@ async def _reserve_counter( ) -> Optional[float]: from litellm.proxy.proxy_server import ( _ensure_spend_counter_initialized, - spend_counter_cache, + _ensure_window_spend_counter_initialized, + _increment_spend_counter_cache, ) if counter.source_cache_key is not None: @@ -417,10 +423,17 @@ async def _reserve_counter( counter_key=counter.counter_key, source_cache_key=counter.source_cache_key, ) + elif counter.spend_log_entity_id is not None and counter.window_start is not None: + await _ensure_window_spend_counter_initialized( + counter_key=counter.counter_key, + entity_type=counter.entity_type, + entity_id=counter.spend_log_entity_id, + window_start=counter.window_start, + ) - reserved_value = await spend_counter_cache.async_increment_cache( - key=counter.counter_key, - value=reservation_cost, + reserved_value = await _increment_spend_counter_cache( + counter_key=counter.counter_key, + increment=reservation_cost, ) return float(reserved_value) if reserved_value is not None else None @@ -438,7 +451,7 @@ async def _set_reserved_entries_adjustment( entries: List[dict], target_adjustment: float, ) -> None: - from litellm.proxy.proxy_server import spend_counter_cache + from litellm.proxy.proxy_server import _increment_spend_counter_cache for entry in entries: counter_key = entry.get("counter_key") @@ -448,9 +461,9 @@ async def _set_reserved_entries_adjustment( adjustment = target_adjustment - applied_adjustment if adjustment == 0: continue - await spend_counter_cache.async_increment_cache( - key=counter_key, - value=adjustment, + await _increment_spend_counter_cache( + counter_key=counter_key, + increment=adjustment, ) entry["applied_adjustment"] = target_adjustment @@ -464,6 +477,37 @@ def _counter_to_reservation_entry(counter: _BudgetCounter) -> Dict[str, Any]: } +def get_budget_window_start(window: Any) -> Optional[datetime]: + window_dict = _coerce_window(window) + budget_duration = window_dict.get("budget_duration") + if budget_duration is None: + return None + try: + duration_seconds = duration_in_seconds(str(budget_duration)) + except Exception: + return None + + reset_at = _coerce_datetime(window_dict.get("reset_at")) + if reset_at is None: + reset_at = datetime.now(timezone.utc) + timedelta(seconds=duration_seconds) + if reset_at.tzinfo is None: + reset_at = reset_at.replace(tzinfo=timezone.utc) + return reset_at - timedelta(seconds=duration_seconds) + + +def _coerce_datetime(value: Any) -> Optional[datetime]: + if value is None: + return None + if isinstance(value, datetime): + return value + if isinstance(value, str): + try: + return datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError: + return None + return None + + def estimate_request_max_cost( request_body: dict, route: str, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 1f4f82a64ef..8f3297780e6 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5,7 +5,7 @@ import os import socket import subprocess import sys -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from pathlib import Path from unittest import mock from unittest.mock import AsyncMock, MagicMock, mock_open, patch @@ -5101,6 +5101,78 @@ async def test_reseed_spend_from_db_skips_window_variant_keys(): fake_prisma.db.litellm_teamtable.find_unique.assert_not_awaited() +@pytest.mark.asyncio +async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + window_start = datetime.now(timezone.utc) - timedelta(hours=1) + fake_prisma = MagicMock() + fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( + return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}] + ) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_window_spend_counter( + counter_key="spend:key:key-window:window:1h", + entity_type="Key", + entity_id="key-window", + window_start=window_start, + increment=0.5, + ) + + fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with( + by=["api_key"], + where={"api_key": "key-window", "startTime": {"gte": window_start}}, + sum={"spend": True}, + ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-window:window:1h" + ) == pytest.approx(2.75) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + +@pytest.mark.asyncio +async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _increment_spend_counter_cache + + counter_cache = DualCache() + counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0) + fake_redis = AsyncMock() + fake_redis.async_increment = AsyncMock(side_effect=RuntimeError("redis down")) + fake_redis.async_delete_cache = AsyncMock() + counter_cache.redis_cache = fake_redis + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + with pytest.raises(RuntimeError): + await _increment_spend_counter_cache( + counter_key="spend:team:redis-fail", + increment=0.5, + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None + ) + fake_redis.async_delete_cache.assert_awaited_once_with( + key="spend:team:redis-fail" + ) + finally: + ps.spend_counter_cache = orig_counter + + @pytest.mark.asyncio async def test_get_current_spend_reseeds_from_db_when_counter_missing(): """ From 09503ebb8fb41db6e0f36784d7b5686993808bc2 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Wed, 29 Apr 2026 21:06:30 -0700 Subject: [PATCH 03/31] harden budget reservation recovery --- litellm/proxy/_types.py | 2 +- litellm/proxy/auth/user_api_key_auth.py | 2 + .../proxy/hooks/proxy_track_cost_callback.py | 19 ++++++- .../spend_tracking/budget_reservation.py | 12 +++++ .../proxy/auth/test_user_api_key_auth.py | 27 ++++++++++ .../hooks/test_proxy_track_cost_callback.py | 13 +++-- .../proxy/test_budget_reservation.py | 49 +++++++++++++++++++ 7 files changed, 117 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8be73f3427e..05e6135c16d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2567,7 +2567,7 @@ class UserAPIKeyAuth( user_spend: Optional[float] = None user_max_budget: Optional[float] = None request_route: Optional[str] = None - budget_reservation: Optional[Dict[str, Any]] = None + budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) user: Optional[Any] = None # Expanded user object when expand=user is used created_by_user: Optional[Any] = ( None # Expanded created_by user when expand=user is used diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 7e70e7bb3fa..0d60b36d35e 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1894,6 +1894,7 @@ async def _reserve_budget_after_common_checks( proxy_logging_obj: ProxyLogging, skip_budget_checks: bool, ) -> None: + user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: return @@ -1962,6 +1963,7 @@ async def user_api_key_auth( request_data=request_data, custom_litellm_key_header=custom_litellm_key_header, ) + user_api_key_auth_obj.budget_reservation = None ## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ## RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 71ad962d584..b5e659757ad 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -446,7 +446,9 @@ async def _update_database_and_spend_counters( ) except Exception: if budget_reservation is not None: - await _release_budget_reservation(budget_reservation=budget_reservation) + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) raise @@ -461,3 +463,18 @@ async def _release_budget_reservation(budget_reservation: Optional[dict]) -> Non await release_budget_reservation( budget_reservation=budget_reservation, ) + + +async def _invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + ) + + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation, + ) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 130a3e20275..18ef6dff0e6 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -153,6 +153,18 @@ async def release_budget_reservation(budget_reservation: Optional[dict]) -> None ) +async def invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.proxy_server import _invalidate_spend_counter + + for counter_key in get_reserved_counter_keys(budget_reservation=budget_reservation): + await _invalidate_spend_counter(counter_key=counter_key) + + async def _get_budget_counters( valid_token: UserAPIKeyAuth, team_object: Optional[LiteLLM_TeamTable], diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 23b1de881f7..bcd56cf872b 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -23,6 +23,7 @@ from litellm.proxy._types import ( from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( + _reserve_budget_after_common_checks, _run_centralized_common_checks, _run_post_custom_auth_checks, get_api_key, @@ -47,6 +48,32 @@ def test_get_api_key(): ) == (api_key, passed_in_key) +@pytest.mark.asyncio +async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): + user_api_key_auth_obj = UserAPIKeyAuth( + token="test_token", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_token"}], + }, + ) + + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "free-model"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=True, + ) + + assert user_api_key_auth_obj.budget_reservation is None + + @pytest.mark.asyncio async def test_custom_auth_does_not_enforce_key_model_access_by_default(): valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"]) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index dc052d8f012..ec7f06ac099 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -312,16 +312,19 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda @pytest.mark.asyncio -async def test_update_database_and_spend_counters_releases_reservation_when_counter_update_fails(): +async def test_update_database_and_spend_counters_invalidates_reservation_when_counter_update_fails(): proxy_logging_obj = MagicMock() proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() increment_spend_counters = AsyncMock(side_effect=Exception("counter unavailable")) - budget_reservation = {"reserved_cost": 0.5, "entries": []} + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } with patch( - "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", new_callable=AsyncMock, - ) as mock_release_budget_reservation: + ) as mock_invalidate_budget_reservation_counters: with pytest.raises(Exception, match="counter unavailable"): await _update_database_and_spend_counters( proxy_logging_obj=proxy_logging_obj, @@ -339,7 +342,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_coun budget_reservation=budget_reservation, ) - mock_release_budget_reservation.assert_awaited_once_with( + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( budget_reservation=budget_reservation, ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 9e154dcb842..c14582638c4 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -14,6 +14,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.spend_tracking.budget_reservation import ( estimate_request_max_cost, + invalidate_budget_reservation_counters, release_budget_reservation, reserve_budget_for_request, ) @@ -50,6 +51,20 @@ def _request_body() -> dict: } +def test_should_not_serialize_budget_reservation_on_user_api_key_auth(): + auth = UserAPIKeyAuth( + token="key-budget-runtime-state", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:key-budget-runtime-state"}], + }, + ) + + assert "budget_reservation" not in auth.model_dump() + assert "budget_reservation" not in auth.model_dump(exclude_none=True) + assert "budget_reservation" not in auth.model_dump_json() + + @pytest.mark.asyncio async def test_should_prevent_second_key_reservation_over_budget( spend_counter_state, @@ -467,6 +482,40 @@ async def test_should_retry_partial_release_without_double_decrement( ) == pytest.approx(0.0) +@pytest.mark.asyncio +async def test_should_invalidate_reserved_counters_after_persisted_spend_failure( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + await counter_cache.async_increment_cache( + key="spend:key:key-budget-invalidate", + value=0.4, + ) + await counter_cache.async_increment_cache( + key="spend:team:team-budget-invalidate", + value=0.4, + ) + + await invalidate_budget_reservation_counters( + { + "reserved_cost": 0.4, + "entries": [ + {"counter_key": "spend:key:key-budget-invalidate"}, + {"counter_key": "spend:team:team-budget-invalidate"}, + ], + } + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-invalidate") + is None + ) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:team-budget-invalidate") + is None + ) + + @pytest.mark.asyncio async def test_should_reserve_all_budgeted_counters(spend_counter_state): counter_cache, key_cache = spend_counter_state From ca50868b752ec023a0b2707e5fa4e35bbe8392af Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 13:37:10 -0700 Subject: [PATCH 04/31] harden end-user and tag budget reservations --- litellm/proxy/auth/auth_checks.py | 43 +-- litellm/proxy/auth/user_api_key_auth.py | 6 + litellm/proxy/db/spend_counter_reseed.py | 12 + .../spend_tracking/budget_reservation.py | 250 ++++++++++++++++++ .../proxy/auth/test_auth_checks.py | 66 +++++ .../proxy/auth/test_user_api_key_auth.py | 76 ++++++ .../proxy/test_budget_reservation.py | 167 ++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 21 ++ 8 files changed, 624 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 840f64cfede..a798a9fe9f7 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -655,13 +655,7 @@ async def common_checks( # noqa: PLR0915 end_user_object is not None and end_user_object.litellm_budget_table is not None ): - end_user_budget = end_user_object.litellm_budget_table.max_budget - if end_user_budget is not None and end_user_object.spend > end_user_budget: - raise litellm.BudgetExceededError( - current_cost=end_user_object.spend, - max_budget=end_user_budget, - message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}", - ) + await _check_end_user_budget(end_user_obj=end_user_object, route=route) _enforce_user_param_check(general_settings, request, request_body, route) _reject_clientside_metadata_tags_check(general_settings, request_body, route) @@ -1005,7 +999,7 @@ async def _apply_default_budget_to_end_user( return end_user_obj -def _check_end_user_budget( +async def _check_end_user_budget( end_user_obj: LiteLLM_EndUserTable, route: str, ) -> None: @@ -1026,11 +1020,20 @@ def _check_end_user_budget( return end_user_budget = end_user_obj.litellm_budget_table.max_budget - if end_user_budget is not None and end_user_obj.spend > end_user_budget: + if end_user_budget is None: + return + + from litellm.proxy.proxy_server import get_current_spend + + end_user_spend = await get_current_spend( + counter_key=f"spend:end_user:{end_user_obj.user_id}", + fallback_spend=end_user_obj.spend or 0.0, + ) + if end_user_spend > end_user_budget: raise litellm.BudgetExceededError( - current_cost=end_user_obj.spend, + current_cost=end_user_spend, max_budget=end_user_budget, - message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_obj.spend}, Budget={end_user_budget}", + message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_spend}, Budget={end_user_budget}", ) @@ -1082,7 +1085,7 @@ async def get_end_user_object( ) # Check budget limits - _check_end_user_budget(end_user_obj=return_obj, route=route) + await _check_end_user_budget(end_user_obj=return_obj, route=route) return return_obj @@ -1113,7 +1116,7 @@ async def get_end_user_object( ) # Check budget limits - _check_end_user_budget(end_user_obj=_response, route=route) + await _check_end_user_budget(end_user_obj=_response, route=route) return _response @@ -3823,13 +3826,19 @@ async def _tag_max_budget_check( if ( tag_object.litellm_budget_table is not None and tag_object.litellm_budget_table.max_budget is not None - and tag_object.spend is not None - and tag_object.spend > tag_object.litellm_budget_table.max_budget ): + from litellm.proxy.proxy_server import get_current_spend + + tag_spend = await get_current_spend( + counter_key=f"spend:tag:{tag_name}", + fallback_spend=tag_object.spend or 0.0, + ) + if tag_spend <= tag_object.litellm_budget_table.max_budget: + continue raise litellm.BudgetExceededError( - current_cost=tag_object.spend, + current_cost=tag_spend, max_budget=tag_object.litellm_budget_table.max_budget, - message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_object.spend}, Max budget: {tag_object.litellm_budget_table.max_budget}", + message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_spend}, Max budget: {tag_object.litellm_budget_table.max_budget}", ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 0d60b36d35e..b21c6483981 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1869,6 +1869,8 @@ async def _run_centralized_common_checks( llm_router=llm_router, team_object=team_object, user_object=user_object, + end_user_id=end_user_id, + end_user_object=end_user_object, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -1893,6 +1895,8 @@ async def _reserve_budget_after_common_checks( user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, skip_budget_checks: bool, + end_user_id: Optional[str] = None, + end_user_object: Optional[LiteLLM_EndUserTable] = None, ) -> None: user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: @@ -1912,6 +1916,8 @@ async def _reserve_budget_after_common_checks( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + end_user_id=end_user_id, + end_user_object=end_user_object, ) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index b4dc4a4d0a2..1a2f64d988d 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -35,6 +35,8 @@ class SpendCounterReseed: spend:team:{team_id} -> LiteLLM_TeamTable.spend spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend spend:user:{user_id} -> LiteLLM_UserTable.spend + spend:end_user:{end_user_id} -> LiteLLM_EndUserTable.spend + spend:tag:{tag_name} -> LiteLLM_TagTable.spend spend:org:{org_id} -> LiteLLM_OrganizationTable.spend """ @@ -98,6 +100,16 @@ class SpendCounterReseed: row = await prisma_client.db.litellm_usertable.find_unique( where={"user_id": user_id} ) + elif counter_key.startswith("spend:end_user:"): + end_user_id = counter_key[len("spend:end_user:") :] + row = await prisma_client.db.litellm_endusertable.find_unique( + where={"user_id": end_user_id} + ) + elif counter_key.startswith("spend:tag:"): + tag_name = counter_key[len("spend:tag:") :] + row = await prisma_client.db.litellm_tagtable.find_unique( + where={"tag_name": tag_name} + ) elif counter_key.startswith("spend:org:"): org_id = counter_key[len("spend:org:") :] row = await prisma_client.db.litellm_organizationtable.find_unique( diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 18ef6dff0e6..7bb9c285225 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -54,6 +54,8 @@ async def reserve_budget_for_request( prisma_client: Optional[PrismaClient], user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, + end_user_id: Optional[str] = None, + end_user_object: Optional[Any] = None, ) -> Optional[dict]: if valid_token is None or not RouteChecks.is_llm_api_route(route=route): return None @@ -63,12 +65,15 @@ async def reserve_budget_for_request( return None counters = await _get_budget_counters( + request_body=request_body, valid_token=valid_token, team_object=team_object, user_object=user_object, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + end_user_id=end_user_id, + end_user_object=end_user_object, ) if not counters: return None @@ -166,12 +171,15 @@ async def invalidate_budget_reservation_counters( async def _get_budget_counters( + request_body: dict, valid_token: UserAPIKeyAuth, team_object: Optional[LiteLLM_TeamTable], user_object: Optional[LiteLLM_UserTable], prisma_client: Optional[PrismaClient], user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, + end_user_id: Optional[str] = None, + end_user_object: Optional[Any] = None, ) -> List[_BudgetCounter]: counters: List[_BudgetCounter] = [] @@ -236,6 +244,25 @@ async def _get_budget_counters( ) ) + end_user_counter = await _get_end_user_budget_counter( + valid_token=valid_token, + end_user_id=end_user_id, + end_user_object=end_user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + if end_user_counter is not None: + counters.append(end_user_counter) + + counters.extend( + await _get_tag_budget_counters( + request_body=request_body, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + ) + team_member_counter = await _get_team_member_budget_counter( valid_token=valid_token, team_object=team_object, @@ -256,6 +283,217 @@ async def _get_budget_counters( return counters +async def _get_end_user_budget_counter( + valid_token: UserAPIKeyAuth, + end_user_id: Optional[str], + end_user_object: Optional[Any], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, +) -> Optional[_BudgetCounter]: + end_user_id = end_user_id or valid_token.end_user_id + if end_user_id is None: + return None + + source_cache_key = f"end_user_id:{end_user_id}" + end_user_obj = end_user_object or ( + await _get_end_user_from_cache_or_db( + end_user_id=end_user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + ) + + max_budget = _to_float(valid_token.end_user_max_budget) + fallback_spend = 0.0 + if end_user_obj is not None: + fallback_spend = _to_float(_get_value(end_user_obj, "spend")) or 0.0 + if max_budget is None: + budget_table = _get_value(end_user_obj, "litellm_budget_table") + max_budget = _to_float(_get_value(budget_table, "max_budget")) + + if max_budget is None: + max_budget = await _get_default_end_user_max_budget( + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + + if max_budget is None or max_budget <= 0: + return None + + return _BudgetCounter( + counter_key=f"spend:end_user:{end_user_id}", + source_cache_key=source_cache_key, + max_budget=max_budget, + fallback_spend=fallback_spend, + entity_type="EndUser", + entity_id=end_user_id, + ) + + +async def _get_end_user_from_cache_or_db( + end_user_id: str, + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, +) -> Optional[Any]: + cache_key = f"end_user_id:{end_user_id}" + cached_end_user = await user_api_key_cache.async_get_cache(key=cache_key) + if cached_end_user is not None: + return cached_end_user + + if prisma_client is None: + return None + + try: + row = await prisma_client.db.litellm_endusertable.find_unique( + where={"user_id": end_user_id}, + include={"litellm_budget_table": True, "object_permission": True}, + ) + except Exception: + verbose_proxy_logger.debug( + "Unable to fetch end-user budget for reservation", exc_info=True + ) + return None + + if row is None: + return None + + row_dict = _object_to_dict(row) + if row_dict: + await user_api_key_cache.async_set_cache(key=cache_key, value=row_dict) + return row_dict + return row + + +async def _get_default_end_user_max_budget( + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, +) -> Optional[float]: + if litellm.max_end_user_budget_id is None or prisma_client is None: + return None + + cache_key = f"default_end_user_budget:{litellm.max_end_user_budget_id}" + cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) + max_budget = _to_float(_get_value(cached_budget, "max_budget")) + if max_budget is not None: + return max_budget + + try: + budget_record = await prisma_client.db.litellm_budgettable.find_unique( + where={"budget_id": litellm.max_end_user_budget_id} + ) + except Exception: + verbose_proxy_logger.debug( + "Unable to fetch default end-user budget for reservation", exc_info=True + ) + return None + + if budget_record is None: + return None + + budget_dict = _object_to_dict(budget_record) + if budget_dict: + await user_api_key_cache.async_set_cache(key=cache_key, value=budget_dict) + return _to_float(_get_value(budget_record, "max_budget")) + + +async def _get_tag_budget_counters( + request_body: dict, + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, +) -> List[_BudgetCounter]: + from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + + tag_names = _dedupe_tags(get_tags_from_request_body(request_body=request_body)) + if not tag_names: + return [] + + tag_objects = await _get_tag_objects_for_reservation( + tag_names=tag_names, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + counters: List[_BudgetCounter] = [] + for tag_name in tag_names: + tag_object = tag_objects.get(tag_name) + if tag_object is None: + continue + budget_table = _get_value(tag_object, "litellm_budget_table") + max_budget = _to_float(_get_value(budget_table, "max_budget")) + if max_budget is None or max_budget <= 0: + continue + counters.append( + _BudgetCounter( + counter_key=f"spend:tag:{tag_name}", + source_cache_key=f"tag:{tag_name}", + max_budget=max_budget, + fallback_spend=_to_float(_get_value(tag_object, "spend")) or 0.0, + entity_type="Tag", + entity_id=tag_name, + ) + ) + return counters + + +async def _get_tag_objects_for_reservation( + tag_names: List[str], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, +) -> Dict[str, Any]: + tag_objects: Dict[str, Any] = {} + uncached_tags: List[str] = [] + + for tag_name in tag_names: + cached_tag = await user_api_key_cache.async_get_cache(key=f"tag:{tag_name}") + if cached_tag is not None: + tag_objects[tag_name] = cached_tag + else: + uncached_tags.append(tag_name) + + if not uncached_tags or prisma_client is None: + return tag_objects + + try: + db_tags = await prisma_client.db.litellm_tagtable.find_many( + where={"tag_name": {"in": uncached_tags}}, + include={"litellm_budget_table": True}, + ) + except Exception: + verbose_proxy_logger.debug( + "Unable to fetch tag budgets for reservation", exc_info=True + ) + return tag_objects + + for db_tag in db_tags: + tag_name = _get_value(db_tag, "tag_name") + if not isinstance(tag_name, str): + continue + row_dict = _object_to_dict(db_tag) + if row_dict: + await user_api_key_cache.async_set_cache( + key=f"tag:{tag_name}", value=row_dict + ) + tag_objects[tag_name] = row_dict + else: + tag_objects[tag_name] = db_tag + + return tag_objects + + +def _dedupe_tags(tags: List[str]) -> List[str]: + seen = set() + deduped_tags = [] + for tag in tags: + if tag in seen: + continue + seen.add(tag) + deduped_tags.append(tag) + return deduped_tags + + async def _get_team_member_budget_counter( valid_token: UserAPIKeyAuth, team_object: Optional[LiteLLM_TeamTable], @@ -727,3 +965,15 @@ def _get_value(obj: Any, key: str) -> Any: if isinstance(obj, dict): return obj.get(key) return getattr(obj, key, None) + + +def _object_to_dict(obj: Any) -> dict: + if isinstance(obj, dict): + return obj + if hasattr(obj, "model_dump"): + value = obj.model_dump() + return value if isinstance(value, dict) else {} + if hasattr(obj, "dict"): + value = obj.dict() + return value if isinstance(value, dict) else {} + return {} diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 676a32c2027..c50278a4921 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -17,7 +17,10 @@ import litellm from litellm.proxy._types import ( CallInfo, Litellm_EntityType, + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, LiteLLM_ObjectPermissionTable, + LiteLLM_TagTable, LiteLLM_TeamTable, LiteLLM_UserTable, LitellmUserRoles, @@ -29,10 +32,12 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, _can_object_call_vector_stores, + _check_end_user_budget, _check_team_member_budget, _get_fuzzy_user_object, _get_team_db_check, _log_budget_lookup_failure, + _tag_max_budget_check, _team_max_budget_check, _virtual_key_max_budget_alert_check, _virtual_key_max_budget_check, @@ -1962,6 +1967,67 @@ async def test_team_budget_check_reads_from_spend_counter(): assert exc_info.value.current_cost == 1.5 +@pytest.mark.asyncio +async def test_end_user_budget_check_reads_from_spend_counter(): + """End-user budget check should use get_current_spend when counter exists.""" + end_user_object = LiteLLM_EndUserTable( + user_id="customer-1", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:end_user:customer-1": + return 1.5 + return fallback_spend + + with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_end_user_budget( + end_user_obj=end_user_object, + route="/chat/completions", + ) + assert exc_info.value.current_cost == 1.5 + assert exc_info.value.max_budget == 1.0 + + +@pytest.mark.asyncio +async def test_tag_budget_check_reads_from_spend_counter(): + """Tag budget check should use get_current_spend when counter exists.""" + from litellm.proxy.utils import ProxyLogging + + tag_object = LiteLLM_TagTable( + tag_name="paid-tag", + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:tag:paid-tag": + return 1.5 + return fallback_spend + + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_tag_objects_batch", + new_callable=AsyncMock, + return_value={"paid-tag": tag_object}, + ), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _tag_max_budget_check( + request_body={"metadata": {"tags": ["paid-tag"]}}, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=None), + valid_token=UserAPIKeyAuth(token="test-token"), + ) + assert exc_info.value.current_cost == 1.5 + assert exc_info.value.max_budget == 1.0 + + @pytest.mark.asyncio async def test_team_member_budget_check_reads_from_spend_counter(): """Team member budget check should use get_current_spend when counter exists.""" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index bcd56cf872b..96a900ffbf1 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -13,6 +13,8 @@ import litellm.proxy.proxy_server from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ( LiteLLM_JWTAuth, + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, LiteLLM_UserTable, LitellmUserRoles, ProxyErrorTypes, @@ -2150,6 +2152,80 @@ async def test_centralized_common_checks_propagates_end_user_budget_error(): setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +async def test_centralized_common_checks_reserves_request_end_user_budget(): + """Regression: reservation runs before user_api_key_auth() copies the + request end-user onto the token, so centralized checks must pass the + locally extracted end_user_id/end_user_object into reservation.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import Request + from starlette.datastructures import URL + + token = UserAPIKeyAuth(api_key="sk-test", user_id="u") + request = Request(scope={"type": "http", "headers": []}) + request._url = URL(url="/chat/completions") + request_data = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello"}], + "user": "alice", + } + end_user_object = LiteLLM_EndUserTable( + user_id="alice", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + counter_cache = DualCache() + attrs["spend_counter_cache"] = counter_cache + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=end_user_object, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ), + ): + assert token.end_user_id is None + + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data=request_data, + route="/chat/completions", + ) + + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + assert token.end_user_id is None + assert token.budget_reservation is not None + assert token.budget_reservation["entries"] == [ + { + "counter_key": "spend:end_user:alice", + "entity_type": "EndUser", + "entity_id": "alice", + "applied_adjustment": 0.0, + } + ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:alice" + ) == pytest.approx(0.6) + + @pytest.mark.asyncio async def test_centralized_common_checks_short_circuits_when_master_key_unset(): """master_key=None is no-auth dev mode — admin-only routes and diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index c14582638c4..57dc894989f 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -6,7 +6,9 @@ import litellm from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ( LiteLLM_BudgetTable, + LiteLLM_EndUserTable, LiteLLM_OrganizationTable, + LiteLLM_TagTable, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable, @@ -118,6 +120,171 @@ async def test_should_prevent_second_key_reservation_over_budget( await release_budget_reservation(reservation) +@pytest.mark.asyncio +async def test_should_prevent_second_end_user_reservation_over_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-end-user", + end_user_id="end-user-budget-race", + end_user_max_budget=1.0, + ) + await key_cache.async_set_cache( + key="end_user_id:end-user-budget-race", + value=LiteLLM_EndUserTable( + user_id="end-user-budget-race", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert reservation is not None + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(0.6) + + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(0.6) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + response_cost=0.2, + budget_reservation=reservation, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(0.2) + + +@pytest.mark.asyncio +async def test_should_prevent_second_tag_reservation_over_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth(token="key-budget-tag") + request_body = _request_body() + request_body["metadata"] = { + "tags": ["tag-budget-race", "tag-without-budget", "tag-budget-race"] + } + await key_cache.async_set_cache( + key="tag:tag-budget-race", + value=LiteLLM_TagTable( + tag_name="tag-budget-race", + spend=0.0, + budget_id="tag-budget-id", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="tag:tag-without-budget", + value=LiteLLM_TagTable( + tag_name="tag-without-budget", + spend=0.0, + ).model_dump(), + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert reservation is not None + assert reservation["entries"] == [ + { + "counter_key": "spend:tag:tag-budget-race", + "entity_type": "Tag", + "entity_id": "tag-budget-race", + "applied_adjustment": 0.0, + } + ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(0.6) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:tag:tag-without-budget") + is None + ) + + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(0.6) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + response_cost=0.2, + budget_reservation=reservation, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(0.2) + + @pytest.mark.asyncio async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_state): counter_cache, key_cache = spend_counter_state diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 8f3297780e6..ed03f3431c4 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5058,11 +5058,19 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): user_row = MagicMock() user_row.spend = 17.0 + end_user_row = MagicMock() + end_user_row.spend = 21.0 + tag_row = MagicMock() + tag_row.spend = 8.0 org_row = MagicMock() org_row.spend = 305.0 fake_prisma = MagicMock() fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + fake_prisma.db.litellm_endusertable.find_unique = AsyncMock( + return_value=end_user_row + ) + fake_prisma.db.litellm_tagtable.find_unique = AsyncMock(return_value=tag_row) fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock( return_value=org_row ) @@ -5072,6 +5080,19 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): where={"user_id": "alice"} ) + assert ( + await SpendCounterReseed.from_db(fake_prisma, "spend:end_user:customer-1") + == 21.0 + ) + fake_prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with( + where={"user_id": "customer-1"} + ) + + assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") == 8.0 + fake_prisma.db.litellm_tagtable.find_unique.assert_awaited_once_with( + where={"tag_name": "paid-tag"} + ) + assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0 fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with( where={"organization_id": "acme"} From 0794ae67be93642ea6041222a109e527781409ed Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 13:49:59 -0700 Subject: [PATCH 05/31] avoid direct budget reservation db lookups --- .../spend_tracking/budget_reservation.py | 151 +----------------- .../proxy/test_budget_reservation.py | 24 +-- 2 files changed, 17 insertions(+), 158 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 7bb9c285225..286f76e2b97 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -248,8 +248,6 @@ async def _get_budget_counters( valid_token=valid_token, end_user_id=end_user_id, end_user_object=end_user_object, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, ) if end_user_counter is not None: counters.append(end_user_counter) @@ -287,36 +285,20 @@ async def _get_end_user_budget_counter( valid_token: UserAPIKeyAuth, end_user_id: Optional[str], end_user_object: Optional[Any], - prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, ) -> Optional[_BudgetCounter]: end_user_id = end_user_id or valid_token.end_user_id if end_user_id is None: return None source_cache_key = f"end_user_id:{end_user_id}" - end_user_obj = end_user_object or ( - await _get_end_user_from_cache_or_db( - end_user_id=end_user_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - ) - ) - max_budget = _to_float(valid_token.end_user_max_budget) fallback_spend = 0.0 - if end_user_obj is not None: - fallback_spend = _to_float(_get_value(end_user_obj, "spend")) or 0.0 + if end_user_object is not None: + fallback_spend = _to_float(_get_value(end_user_object, "spend")) or 0.0 if max_budget is None: - budget_table = _get_value(end_user_obj, "litellm_budget_table") + budget_table = _get_value(end_user_object, "litellm_budget_table") max_budget = _to_float(_get_value(budget_table, "max_budget")) - if max_budget is None: - max_budget = await _get_default_end_user_max_budget( - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - ) - if max_budget is None or max_budget <= 0: return None @@ -330,72 +312,6 @@ async def _get_end_user_budget_counter( ) -async def _get_end_user_from_cache_or_db( - end_user_id: str, - prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, -) -> Optional[Any]: - cache_key = f"end_user_id:{end_user_id}" - cached_end_user = await user_api_key_cache.async_get_cache(key=cache_key) - if cached_end_user is not None: - return cached_end_user - - if prisma_client is None: - return None - - try: - row = await prisma_client.db.litellm_endusertable.find_unique( - where={"user_id": end_user_id}, - include={"litellm_budget_table": True, "object_permission": True}, - ) - except Exception: - verbose_proxy_logger.debug( - "Unable to fetch end-user budget for reservation", exc_info=True - ) - return None - - if row is None: - return None - - row_dict = _object_to_dict(row) - if row_dict: - await user_api_key_cache.async_set_cache(key=cache_key, value=row_dict) - return row_dict - return row - - -async def _get_default_end_user_max_budget( - prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, -) -> Optional[float]: - if litellm.max_end_user_budget_id is None or prisma_client is None: - return None - - cache_key = f"default_end_user_budget:{litellm.max_end_user_budget_id}" - cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) - max_budget = _to_float(_get_value(cached_budget, "max_budget")) - if max_budget is not None: - return max_budget - - try: - budget_record = await prisma_client.db.litellm_budgettable.find_unique( - where={"budget_id": litellm.max_end_user_budget_id} - ) - except Exception: - verbose_proxy_logger.debug( - "Unable to fetch default end-user budget for reservation", exc_info=True - ) - return None - - if budget_record is None: - return None - - budget_dict = _object_to_dict(budget_record) - if budget_dict: - await user_api_key_cache.async_set_cache(key=cache_key, value=budget_dict) - return _to_float(_get_value(budget_record, "max_budget")) - - async def _get_tag_budget_counters( request_body: dict, prisma_client: Optional[PrismaClient], @@ -403,12 +319,13 @@ async def _get_tag_budget_counters( proxy_logging_obj: ProxyLogging, ) -> List[_BudgetCounter]: from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.auth.auth_checks import get_tag_objects_batch tag_names = _dedupe_tags(get_tags_from_request_body(request_body=request_body)) if not tag_names: return [] - tag_objects = await _get_tag_objects_for_reservation( + tag_objects = await get_tag_objects_batch( tag_names=tag_names, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -437,52 +354,6 @@ async def _get_tag_budget_counters( return counters -async def _get_tag_objects_for_reservation( - tag_names: List[str], - prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, - proxy_logging_obj: ProxyLogging, -) -> Dict[str, Any]: - tag_objects: Dict[str, Any] = {} - uncached_tags: List[str] = [] - - for tag_name in tag_names: - cached_tag = await user_api_key_cache.async_get_cache(key=f"tag:{tag_name}") - if cached_tag is not None: - tag_objects[tag_name] = cached_tag - else: - uncached_tags.append(tag_name) - - if not uncached_tags or prisma_client is None: - return tag_objects - - try: - db_tags = await prisma_client.db.litellm_tagtable.find_many( - where={"tag_name": {"in": uncached_tags}}, - include={"litellm_budget_table": True}, - ) - except Exception: - verbose_proxy_logger.debug( - "Unable to fetch tag budgets for reservation", exc_info=True - ) - return tag_objects - - for db_tag in db_tags: - tag_name = _get_value(db_tag, "tag_name") - if not isinstance(tag_name, str): - continue - row_dict = _object_to_dict(db_tag) - if row_dict: - await user_api_key_cache.async_set_cache( - key=f"tag:{tag_name}", value=row_dict - ) - tag_objects[tag_name] = row_dict - else: - tag_objects[tag_name] = db_tag - - return tag_objects - - def _dedupe_tags(tags: List[str]) -> List[str]: seen = set() deduped_tags = [] @@ -965,15 +836,3 @@ def _get_value(obj: Any, key: str) -> Any: if isinstance(obj, dict): return obj.get(key) return getattr(obj, key, None) - - -def _object_to_dict(obj: Any) -> dict: - if isinstance(obj, dict): - return obj - if hasattr(obj, "model_dump"): - value = obj.model_dump() - return value if isinstance(value, dict) else {} - if hasattr(obj, "dict"): - value = obj.dict() - return value if isinstance(value, dict) else {} - return {} diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 57dc894989f..350f82e2f7c 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -129,16 +129,12 @@ async def test_should_prevent_second_end_user_reservation_over_budget( valid_token = UserAPIKeyAuth( token="key-budget-end-user", end_user_id="end-user-budget-race", - end_user_max_budget=1.0, ) - await key_cache.async_set_cache( - key="end_user_id:end-user-budget-race", - value=LiteLLM_EndUserTable( - user_id="end-user-budget-race", - blocked=False, - spend=0.0, - litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), - ).model_dump(), + end_user_object = LiteLLM_EndUserTable( + user_id="end-user-budget-race", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), ) with patch( @@ -155,6 +151,7 @@ async def test_should_prevent_second_end_user_reservation_over_budget( prisma_client=None, user_api_key_cache=key_cache, proxy_logging_obj=proxy_logging_obj, + end_user_object=end_user_object, ) assert reservation is not None assert counter_cache.in_memory_cache.get_cache( @@ -172,6 +169,7 @@ async def test_should_prevent_second_end_user_reservation_over_budget( prisma_client=None, user_api_key_cache=key_cache, proxy_logging_obj=proxy_logging_obj, + end_user_object=end_user_object, ) assert counter_cache.in_memory_cache.get_cache( @@ -220,6 +218,8 @@ async def test_should_prevent_second_tag_reservation_over_budget( spend=0.0, ).model_dump(), ) + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[]) with patch( "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", @@ -232,7 +232,7 @@ async def test_should_prevent_second_tag_reservation_over_budget( valid_token=valid_token, team_object=None, user_object=None, - prisma_client=None, + prisma_client=prisma_client, user_api_key_cache=key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -261,7 +261,7 @@ async def test_should_prevent_second_tag_reservation_over_budget( valid_token=valid_token, team_object=None, user_object=None, - prisma_client=None, + prisma_client=prisma_client, user_api_key_cache=key_cache, proxy_logging_obj=proxy_logging_obj, ) From 08705d2b3c0de9d52fb52645b0ed0b919654a930 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 13:58:43 -0700 Subject: [PATCH 06/31] stabilize lazy openapi snapshot --- litellm/proxy/_lazy_openapi_snapshot.json | 34 +++++++++++------------ litellm/proxy/_lazy_openapi_snapshot.py | 15 ++++++++++ 2 files changed, 32 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8331f748c6e..b8e9eb6c261 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3572,7 +3572,7 @@ "/anthropic/{endpoint}": { "delete": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3616,7 +3616,7 @@ }, "get": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3660,7 +3660,7 @@ }, "patch": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3704,7 +3704,7 @@ }, "post": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3748,7 +3748,7 @@ }, "put": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -13260,7 +13260,7 @@ "/langfuse/{endpoint}": { "delete": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13299,7 +13299,7 @@ }, "get": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13338,7 +13338,7 @@ }, "patch": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13377,7 +13377,7 @@ }, "post": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13416,7 +13416,7 @@ }, "put": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -26883,7 +26883,7 @@ "/toolset/{toolset_name}/mcp": { "delete": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -26922,7 +26922,7 @@ }, "get": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -26961,7 +26961,7 @@ }, "head": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27000,7 +27000,7 @@ }, "options": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27039,7 +27039,7 @@ }, "patch": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27078,7 +27078,7 @@ }, "post": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27117,7 +27117,7 @@ }, "put": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 315f6a9742a..a51cb173dd5 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -8,6 +8,7 @@ any drift as a neutral check. """ import json +import re import sys from pathlib import Path from typing import Dict, Optional @@ -15,6 +16,19 @@ from typing import Dict, Optional SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" +def _stabilize_multi_method_route_ids(routes) -> None: + """FastAPI derives route IDs from a set of methods; make snapshots stable.""" + + for route in routes: + methods = sorted(getattr(route, "methods", None) or []) + if len(methods) <= 1 or not getattr(route, "path_format", None): + continue + + operation_id = f"{route.name}{route.path_format}" + operation_id = re.sub(r"\W", "_", operation_id) + route.unique_id = f"{operation_id}_{methods[0].lower()}" + + def load_snapshot() -> Optional[Dict[str, Dict]]: if not SNAPSHOT_FILE.exists(): return None @@ -51,6 +65,7 @@ def generate_snapshot() -> Dict[str, Dict]: ] if not feat_routes: continue + _stabilize_multi_method_route_ids(feat_routes) full = get_openapi(title=app.title, version=app.version, routes=feat_routes) # Group all of a feature's routes under one tag. for path_ops in full.get("paths", {}).values(): From 0b71282985941d4b29a86b872a47ed6f95ed702f Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:06:42 -0700 Subject: [PATCH 07/31] address budget reservation review findings --- .../proxy/hooks/proxy_track_cost_callback.py | 11 ++-- .../spend_tracking/budget_reservation.py | 38 +++++++++++++- .../hooks/test_proxy_track_cost_callback.py | 42 +++++++++++++++ .../proxy/test_budget_reservation.py | 52 +++++++++++++++++++ 4 files changed, 139 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index b5e659757ad..974bbc66843 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -37,9 +37,14 @@ class _ProxyDBLogger(CustomLogger): user_api_key_dict: UserAPIKeyAuth, traceback_str: Optional[str] = None, ): - await _release_budget_reservation( - budget_reservation=user_api_key_dict.budget_reservation - ) + try: + await _release_budget_reservation( + budget_reservation=user_api_key_dict.budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to release budget reservation during failure handling" + ) request_route = user_api_key_dict.request_route if _ProxyDBLogger._should_track_errors_in_db() is False: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 286f76e2b97..a7413e23709 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -31,6 +31,7 @@ class _BudgetCounter: source_cache_key: Optional[str] = None spend_log_entity_id: Optional[str] = None window_start: Optional[datetime] = None + parent_counter_key: Optional[str] = None def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: @@ -201,6 +202,7 @@ async def _get_budget_counters( entity_type="Key", entity_id=valid_token.token, budget_limits=valid_token.budget_limits, + fallback_spend=float(valid_token.spend or 0.0), ) ) @@ -223,6 +225,7 @@ async def _get_budget_counters( entity_type="Team", entity_id=team_id, budget_limits=team_object.budget_limits, + fallback_spend=float(team_object.spend or 0.0), ) ) @@ -463,6 +466,7 @@ def _get_budget_limit_counters( entity_type: str, entity_id: str, budget_limits: Optional[Sequence[Any]], + fallback_spend: float, ) -> List[_BudgetCounter]: counters: List[_BudgetCounter] = [] if not budget_limits: @@ -479,11 +483,12 @@ def _get_budget_limit_counters( _BudgetCounter( counter_key=f"{entity_prefix}:window:{budget_duration}", max_budget=float(max_budget), - fallback_spend=0.0, + fallback_spend=fallback_spend if window_start is None else 0.0, entity_type=entity_type, entity_id=f"{entity_id}:{budget_duration}", spend_log_entity_id=entity_id, window_start=window_start, + parent_counter_key=entity_prefix if window_start is None else None, ) ) return counters @@ -551,6 +556,8 @@ async def _reserve_counter( entity_id=counter.spend_log_entity_id, window_start=counter.window_start, ) + elif counter.parent_counter_key is not None: + await _ensure_malformed_window_counter_initialized(counter=counter) reserved_value = await _increment_spend_counter_cache( counter_key=counter.counter_key, @@ -562,12 +569,41 @@ async def _reserve_counter( async def _get_current_counter_value(counter: _BudgetCounter) -> float: from litellm.proxy.proxy_server import get_current_spend + if counter.parent_counter_key is not None: + await _ensure_malformed_window_counter_initialized(counter=counter) + return await get_current_spend( counter_key=counter.counter_key, fallback_spend=counter.fallback_spend, ) +async def _ensure_malformed_window_counter_initialized( + counter: _BudgetCounter, +) -> None: + if counter.parent_counter_key is None: + return + + from litellm.proxy.proxy_server import ( + _increment_spend_counter_cache, + get_current_spend, + spend_counter_cache, + ) + + current = await spend_counter_cache.async_get_cache(key=counter.counter_key) + if current is not None: + return + + parent_spend = await get_current_spend( + counter_key=counter.parent_counter_key, + fallback_spend=counter.fallback_spend, + ) + await _increment_spend_counter_cache( + counter_key=counter.counter_key, + increment=parent_spend, + ) + + async def _set_reserved_entries_adjustment( entries: List[dict], target_adjustment: float, diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index ec7f06ac099..f7c12306e4a 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -159,6 +159,48 @@ async def test_async_post_call_failure_hook_releases_budget_reservation_before_r mock_update_database.assert_not_called() +@pytest.mark.asyncio +async def test_should_continue_failure_tracking_when_budget_release_fails(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_api_key", + user_id="test_user_id", + team_id="test_team_id", + request_route="/chat/completions", + budget_reservation=budget_reservation, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + side_effect=RuntimeError("redis unavailable"), + ) as mock_release_budget_reservation, + patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", + ) as mock_log_exception, + ): + await logger.async_post_call_failure_hook( + request_data={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + }, + original_exception=Exception("provider failed"), + user_api_key_dict=user_api_key_dict, + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + mock_log_exception.assert_called_once() + mock_update_database.assert_called_once() + + @pytest.mark.asyncio async def test_track_cost_callback_releases_budget_reservation_when_spend_tracking_skips(): logger = _ProxyDBLogger() diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 350f82e2f7c..9873dcdf727 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -442,6 +442,58 @@ async def test_should_reserve_remaining_budget_when_output_cap_missing( await release_budget_reservation(reservation) +@pytest.mark.asyncio +async def test_should_seed_malformed_window_counter_from_parent_authoritative_spend( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-malformed-window", + spend=0.0, + budget_limits=[ + { + "budget_duration": "not-a-duration", + "max_budget": 1.0, + } + ], + ) + + import litellm.proxy.proxy_server as ps + + db_row = MagicMock() + db_row.spend = 0.9 + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=db_row + ) + ps.prisma_client = prisma_client + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.2, + ): + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=prisma_client, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + prisma_client.db.litellm_verificationtoken.find_unique.assert_awaited_once_with( + where={"token": "key-budget-malformed-window"} + ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-malformed-window:window:not-a-duration" + ) == pytest.approx(0.9) + + @pytest.mark.asyncio async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( spend_counter_state, From 5521af096e9b28edb4095a736bd4841eb9ca8045 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:23:57 -0700 Subject: [PATCH 08/31] preserve database failure during reservation cleanup --- .../proxy/hooks/proxy_track_cost_callback.py | 7 ++- .../hooks/test_proxy_track_cost_callback.py | 48 +++++++++++++++++++ 2 files changed, 54 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 974bbc66843..9ea6c65eeb5 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -437,7 +437,12 @@ async def _update_database_and_spend_counters( ) except Exception: if budget_reservation is not None: - await _release_budget_reservation(budget_reservation=budget_reservation) + try: + await _release_budget_reservation(budget_reservation=budget_reservation) + except Exception: + verbose_proxy_logger.exception( + "Failed to release budget reservation after database update failed" + ) raise try: diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index f7c12306e4a..c0ae488d47b 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -319,6 +319,54 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u increment_spend_counters.assert_not_awaited() +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails(): + proxy_logging_obj = MagicMock() + db_exception = RuntimeError("db unavailable") + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( + side_effect=db_exception + ) + increment_spend_counters = AsyncMock() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + side_effect=RuntimeError("release unavailable"), + ) as mock_release_budget_reservation, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", + ) as mock_log_exception, + ): + with pytest.raises(RuntimeError) as exc_info: + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + assert exc_info.value is db_exception + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + mock_log_exception.assert_called_once_with( + "Failed to release budget reservation after database update failed" + ) + + increment_spend_counters.assert_not_awaited() + + @pytest.mark.asyncio async def test_update_database_and_spend_counters_updates_counters_after_db_update(): proxy_logging_obj = MagicMock() From 719e891c3af32fb34df0b1b4a432c0f8661e8a3b Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:38:37 -0700 Subject: [PATCH 09/31] coalesce malformed window reservation seeding --- .../spend_tracking/budget_reservation.py | 23 ++++-- .../proxy/test_budget_reservation.py | 76 +++++++++++++++++++ 2 files changed, 91 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index a7413e23709..c23f444595c 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -584,6 +584,7 @@ async def _ensure_malformed_window_counter_initialized( if counter.parent_counter_key is None: return + from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed from litellm.proxy.proxy_server import ( _increment_spend_counter_cache, get_current_spend, @@ -594,14 +595,20 @@ async def _ensure_malformed_window_counter_initialized( if current is not None: return - parent_spend = await get_current_spend( - counter_key=counter.parent_counter_key, - fallback_spend=counter.fallback_spend, - ) - await _increment_spend_counter_cache( - counter_key=counter.counter_key, - increment=parent_spend, - ) + lock = await SpendCounterReseed._get_lock(counter.counter_key) + async with lock: + current = await spend_counter_cache.async_get_cache(key=counter.counter_key) + if current is not None: + return + + parent_spend = await get_current_spend( + counter_key=counter.parent_counter_key, + fallback_spend=counter.fallback_spend, + ) + await _increment_spend_counter_cache( + counter_key=counter.counter_key, + increment=parent_spend, + ) async def _set_reserved_entries_adjustment( diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 9873dcdf727..7ed62788083 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -494,6 +494,82 @@ async def test_should_seed_malformed_window_counter_from_parent_authoritative_sp ) == pytest.approx(0.9) +@pytest.mark.asyncio +async def test_should_coalesce_malformed_window_counter_initialization( + spend_counter_state, +): + import asyncio as _asyncio + + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + token = "key-budget-malformed-window-concurrent" + malformed_window_counter_key = f"spend:key:{token}:window:not-a-duration" + valid_token = UserAPIKeyAuth( + token=token, + spend=0.0, + budget_limits=[ + { + "budget_duration": "not-a-duration", + "max_budget": 1.0, + } + ], + ) + + import litellm.proxy.proxy_server as ps + + db_call_count = 0 + + async def slow_find_unique(**kwargs): + nonlocal db_call_count + db_call_count += 1 + await _asyncio.sleep(0.05) + db_row = MagicMock() + db_row.spend = 0.35 + return db_row + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + side_effect=slow_find_unique + ) + ps.prisma_client = prisma_client + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.2, + ): + results = await _asyncio.gather( + *[ + reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=prisma_client, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + for _ in range(2) + ], + return_exceptions=True, + ) + + assert not any(isinstance(result, Exception) for result in results), results + assert all(result is not None for result in results) + assert db_call_count == 1 + assert counter_cache.in_memory_cache.get_cache( + key=malformed_window_counter_key + ) == pytest.approx(0.75) + + for reservation in results: + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key=malformed_window_counter_key + ) == pytest.approx(0.35) + + @pytest.mark.asyncio async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( spend_counter_state, From 8311456cfc61d888bffa72ddbced57d1aa2f9e7f Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:48:55 -0700 Subject: [PATCH 10/31] invalidate reservations after release cleanup failure --- litellm/proxy/hooks/proxy_track_cost_callback.py | 8 ++++++++ .../proxy/hooks/test_proxy_track_cost_callback.py | 14 +++++++++++++- 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 9ea6c65eeb5..c1e039aff51 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -443,6 +443,14 @@ async def _update_database_and_spend_counters( verbose_proxy_logger.exception( "Failed to release budget reservation after database update failed" ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after release failed" + ) raise try: diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index c0ae488d47b..614697a489f 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -338,6 +338,11 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re patch( "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", ) as mock_log_exception, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback._invalidate_budget_reservation_counters", + new_callable=AsyncMock, + side_effect=RuntimeError("invalidate unavailable"), + ) as mock_invalidate_budget_reservation_counters, ): with pytest.raises(RuntimeError) as exc_info: await _update_database_and_spend_counters( @@ -360,9 +365,16 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re mock_release_budget_reservation.assert_awaited_once_with( budget_reservation=budget_reservation, ) - mock_log_exception.assert_called_once_with( + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + assert mock_log_exception.call_count == 2 + mock_log_exception.assert_any_call( "Failed to release budget reservation after database update failed" ) + mock_log_exception.assert_any_call( + "Failed to invalidate budget reservation counters after release failed" + ) increment_spend_counters.assert_not_awaited() From 96a283ed0f36543d67bd1ed9481a33fd6df95dd3 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 15:00:35 -0700 Subject: [PATCH 11/31] guard reservation invalidation cleanup --- .../proxy/hooks/proxy_track_cost_callback.py | 11 +++-- .../hooks/test_proxy_track_cost_callback.py | 49 +++++++++++++++++++ 2 files changed, 57 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index c1e039aff51..6e274f16ce1 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -464,9 +464,14 @@ async def _update_database_and_spend_counters( ) except Exception: if budget_reservation is not None: - await _invalidate_budget_reservation_counters( - budget_reservation=budget_reservation - ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after spend counter update failed" + ) raise diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 614697a489f..65f04c290d5 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -451,6 +451,55 @@ async def test_update_database_and_spend_counters_invalidates_reservation_when_c proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_preserves_counter_exception_when_invalidation_fails(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() + counter_exception = RuntimeError("counter unavailable") + increment_spend_counters = AsyncMock(side_effect=counter_exception) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", + new_callable=AsyncMock, + side_effect=RuntimeError("invalidate unavailable"), + ) as mock_invalidate_budget_reservation_counters, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", + ) as mock_log_exception, + ): + with pytest.raises(RuntimeError) as exc_info: + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + assert exc_info.value is counter_exception + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + mock_log_exception.assert_called_once_with( + "Failed to invalidate budget reservation counters after spend counter update failed" + ) + + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + + @pytest.mark.asyncio async def test_track_cost_callback_skips_when_no_standard_logging_object(): """ From 1373ae10218f76aa964dfdcdb47ce0f82b9adb85 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:09:38 -0700 Subject: [PATCH 12/31] fix budget tag spend counter reconciliation --- .../proxy/hooks/proxy_track_cost_callback.py | 43 +++++--- litellm/proxy/proxy_server.py | 97 +++++++++++++++++-- .../hooks/test_proxy_track_cost_callback.py | 5 +- .../proxy/test_budget_reservation.py | 36 +++++++ 4 files changed, 160 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 6e274f16ce1..a61fda91fdf 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -178,17 +178,18 @@ class _ProxyDBLogger(CustomLogger): sl_object: Optional[StandardLoggingPayload] = kwargs.get( "standard_logging_object", None ) - response_cost = ( - sl_object.get("response_cost", None) - if sl_object is not None - else kwargs.get("response_cost", None) - ) - tags: Optional[List[str]] = ( - sl_object.get("request_tags", None) if sl_object is not None else None - ) - - if response_cost is not None: - user_api_key = metadata.get("user_api_key", None) + response_cost = ( + sl_object.get("response_cost", None) + if sl_object is not None + else kwargs.get("response_cost", None) + ) + tags = _get_request_tags_for_cost_tracking( + sl_object=sl_object, + metadata=metadata, + ) + + if response_cost is not None: + user_api_key = metadata.get("user_api_key", None) if kwargs.get("cache_hit", False) is True: response_cost = 0.0 verbose_proxy_logger.debug( @@ -219,6 +220,7 @@ class _ProxyDBLogger(CustomLogger): end_time=end_time, response_cost=response_cost, budget_reservation=budget_reservation, + request_tags=tags, ) # update cache (fire-and-forget for backward compat: @@ -407,6 +409,22 @@ def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: return getattr(user_api_key_auth_obj, "budget_reservation", None) +def _get_request_tags_for_cost_tracking( + sl_object: Optional[StandardLoggingPayload], + metadata: dict, +) -> Optional[List[str]]: + if sl_object is not None: + request_tags = sl_object.get("request_tags", None) + if isinstance(request_tags, list): + return request_tags + + metadata_tags = metadata.get("tags", None) + if isinstance(metadata_tags, list): + return metadata_tags + + return None + + async def _update_database_and_spend_counters( proxy_logging_obj: Any, increment_spend_counters: Any, @@ -421,6 +439,7 @@ async def _update_database_and_spend_counters( end_time: Any, response_cost: float, budget_reservation: Optional[dict], + request_tags: Optional[List[str]] = None, ) -> None: try: await proxy_logging_obj.db_spend_update_writer.update_database( @@ -461,6 +480,8 @@ async def _update_database_and_spend_counters( response_cost=response_cost, org_id=org_id, budget_reservation=budget_reservation, + end_user_id=end_user_id, + tags=request_tags, ) except Exception: if budget_reservation is not None: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 10d4a8fb073..2f2b59db679 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1836,6 +1836,8 @@ async def increment_spend_counters( response_cost: Optional[float], org_id: Optional[str] = None, budget_reservation: Optional[dict] = None, + end_user_id: Optional[str] = None, + tags: Optional[List[str]] = None, ): """ Atomically increment spend counters for budget enforcement. @@ -1847,7 +1849,7 @@ async def increment_spend_counters( Awaited (not create_task) in the cost callback, so the counter is updated before the next request's auth check runs. """ - reserved_counter_keys = set() + reserved_counter_keys: Set[str] = set() if budget_reservation is not None: from litellm.proxy.spend_tracking.budget_reservation import ( get_reserved_counter_keys, @@ -1970,14 +1972,91 @@ async def increment_spend_counters( increment=response_cost, ) - if org_id is not None: - org_counter_key = f"spend:org:{org_id}" - if org_counter_key not in reserved_counter_keys: - await _init_and_increment_spend_counter( - counter_key=org_counter_key, - source_cache_key=f"org_id:{org_id}:with_budget", - increment=response_cost, - ) + await _increment_end_user_and_tag_spend_counters( + end_user_id=end_user_id, + tags=tags, + response_cost=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + await _increment_org_spend_counter( + org_id=org_id, + response_cost=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + +async def _increment_end_user_and_tag_spend_counters( + end_user_id: Optional[str], + tags: Optional[List[str]], + response_cost: float, + reserved_counter_keys: Set[str], +) -> None: + if end_user_id is not None: + await _increment_warm_unreserved_spend_counter( + counter_key=f"spend:end_user:{end_user_id}", + increment=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + if tags is None: + return + + seen_tags: Set[str] = set() + for tag_name in tags: + if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags: + continue + seen_tags.add(tag_name) + await _increment_warm_unreserved_spend_counter( + counter_key=f"spend:tag:{tag_name}", + increment=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + +async def _increment_warm_unreserved_spend_counter( + counter_key: str, + increment: float, + reserved_counter_keys: Set[str], +) -> None: + if counter_key in reserved_counter_keys: + return + if await spend_counter_cache.async_get_cache(key=counter_key) is None: + return + + await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + + +async def _increment_org_spend_counter( + org_id: Optional[str], + response_cost: float, + reserved_counter_keys: Set[str], +) -> None: + if org_id is None: + return + + await _init_and_increment_unreserved_spend_counter( + counter_key=f"spend:org:{org_id}", + source_cache_key=f"org_id:{org_id}:with_budget", + increment=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + +async def _init_and_increment_unreserved_spend_counter( + counter_key: str, + source_cache_key: str, + increment: float, + reserved_counter_keys: Set[str], +) -> None: + if counter_key in reserved_counter_keys: + return + + await _init_and_increment_spend_counter( + counter_key=counter_key, + source_cache_key=source_cache_key, + increment=increment, + ) async def _init_and_increment_spend_counter( diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 65f04c290d5..9f414f560c7 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -391,7 +391,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda increment_spend_counters=increment_spend_counters, user_api_key="test_api_key", user_id="test_user_id", - end_user_id=None, + end_user_id="test_end_user_id", team_id="test_team_id", org_id="test_org_id", kwargs={}, @@ -400,6 +400,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda end_time=datetime.now(), response_cost=0.2, budget_reservation=budget_reservation, + request_tags=["tag-a"], ) proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() @@ -410,6 +411,8 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda response_cost=0.2, org_id="test_org_id", budget_reservation=budget_reservation, + end_user_id="test_end_user_id", + tags=["tag-a"], ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 7ed62788083..481d88b8d04 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -184,6 +184,7 @@ async def test_should_prevent_second_end_user_reservation_over_budget( user_id=None, response_cost=0.2, budget_reservation=reservation, + end_user_id="end-user-budget-race", ) assert counter_cache.in_memory_cache.get_cache( @@ -278,6 +279,7 @@ async def test_should_prevent_second_tag_reservation_over_budget( user_id=None, response_cost=0.2, budget_reservation=reservation, + tags=["tag-budget-race"], ) assert counter_cache.in_memory_cache.get_cache( @@ -285,6 +287,40 @@ async def test_should_prevent_second_tag_reservation_over_budget( ) == pytest.approx(0.2) +@pytest.mark.asyncio +async def test_should_update_warm_end_user_and_tag_counters_without_reservation( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + counter_cache.in_memory_cache.set_cache( + key="spend:end_user:customer-1", + value=4.0, + ) + counter_cache.in_memory_cache.set_cache(key="spend:tag:paid-tag", value=7.0) + counter_cache.in_memory_cache.set_cache(key="spend:tag:other-tag", value=2.0) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + response_cost=0.50, + end_user_id="customer-1", + tags=["paid-tag", "paid-tag", "other-tag", ""], + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:customer-1" + ) == pytest.approx(4.50) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:paid-tag" + ) == pytest.approx(7.50) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:other-tag" + ) == pytest.approx(2.50) + + @pytest.mark.asyncio async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_state): counter_cache, key_cache = spend_counter_state From e034935b538889a27ee424222a3a66fd930cf127 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:26:43 -0700 Subject: [PATCH 13/31] fix budget reservation window fallback races --- .../spend_tracking/budget_reservation.py | 28 +++++++- .../proxy/test_budget_reservation.py | 71 +++++++++++++++++++ 2 files changed, 98 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index c23f444595c..b17251f4b55 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -85,6 +85,7 @@ async def reserve_budget_for_request( route=route, llm_router=llm_router, ) + using_remaining_budget_fallback = reservation_cost is None if reservation_cost is None: reservation_cost = await _get_smallest_remaining_budget( counters=counters, @@ -111,6 +112,18 @@ async def reserve_budget_for_request( cached_spend = await _get_current_counter_value(counter=counter) current_spend = cached_spend + reservation_cost if current_spend > counter.max_budget: + if using_remaining_budget_fallback: + remaining_before_reservation = counter.max_budget - ( + current_spend - reservation_cost + ) + if remaining_before_reservation > 0: + await _resize_applied_reservation( + entries=applied_entries, + current_reserved_cost=reservation_cost, + new_reserved_cost=remaining_before_reservation, + ) + reservation_cost = remaining_before_reservation + continue raise litellm.BudgetExceededError( current_cost=current_spend, max_budget=counter.max_budget, @@ -632,6 +645,19 @@ async def _set_reserved_entries_adjustment( entry["applied_adjustment"] = target_adjustment +async def _resize_applied_reservation( + entries: List[dict], + current_reserved_cost: float, + new_reserved_cost: float, +) -> None: + await _set_reserved_entries_adjustment( + entries=entries, + target_adjustment=new_reserved_cost - current_reserved_cost, + ) + for entry in entries: + entry["applied_adjustment"] = 0.0 + + def _counter_to_reservation_entry(counter: _BudgetCounter) -> Dict[str, Any]: return { "counter_key": counter.counter_key, @@ -653,7 +679,7 @@ def get_budget_window_start(window: Any) -> Optional[datetime]: reset_at = _coerce_datetime(window_dict.get("reset_at")) if reset_at is None: - reset_at = datetime.now(timezone.utc) + timedelta(seconds=duration_seconds) + return datetime.now(timezone.utc) - timedelta(seconds=duration_seconds) if reset_at.tzinfo is None: reset_at = reset_at.replace(tzinfo=timezone.utc) return reset_at - timedelta(seconds=duration_seconds) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 481d88b8d04..51016126d35 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -1,3 +1,4 @@ +from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -16,6 +17,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.spend_tracking.budget_reservation import ( estimate_request_max_cost, + get_budget_window_start, invalidate_budget_reservation_counters, release_budget_reservation, reserve_budget_for_request, @@ -478,6 +480,75 @@ async def test_should_reserve_remaining_budget_when_output_cap_missing( await release_budget_reservation(reservation) +@pytest.mark.asyncio +async def test_should_shrink_uncapped_reservation_when_counter_advances( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-uncapped-race", + spend=0.2, + max_budget=1.0, + ) + request_body = _request_body() + request_body.pop("max_tokens") + + from litellm.proxy.spend_tracking import budget_reservation + + async def stale_counter_read(counter): + await counter_cache.async_increment_cache( + key=counter.counter_key, + value=0.3, + ) + return 0.2 + + monkeypatch.setattr( + budget_reservation, + "_get_current_counter_value", + stale_counter_read, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=None, + ): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.7) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-uncapped-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-uncapped-race" + ) == pytest.approx(0.3) + + +def test_should_start_window_without_reset_at_at_duration_boundary(): + before = datetime.now(timezone.utc) - timedelta(hours=1) + + window_start = get_budget_window_start({"budget_duration": "1h"}) + + after = datetime.now(timezone.utc) - timedelta(hours=1) + assert window_start is not None + assert before <= window_start <= after + + @pytest.mark.asyncio async def test_should_seed_malformed_window_counter_from_parent_authoritative_spend( spend_counter_state, From f8d187785d82deb8d9b74b87aea73bc537e0019a Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:52:09 -0700 Subject: [PATCH 14/31] finalize invalidated budget reservations --- litellm/proxy/hooks/proxy_track_cost_callback.py | 1 + tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py | 1 + 2 files changed, 2 insertions(+) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index a61fda91fdf..86f7849bb15 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -489,6 +489,7 @@ async def _update_database_and_spend_counters( await _invalidate_budget_reservation_counters( budget_reservation=budget_reservation ) + budget_reservation["finalized"] = True except Exception: verbose_proxy_logger.exception( "Failed to invalidate budget reservation counters after spend counter update failed" diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 9f414f560c7..351d3059d4e 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -450,6 +450,7 @@ async def test_update_database_and_spend_counters_invalidates_reservation_when_c mock_invalidate_budget_reservation_counters.assert_awaited_once_with( budget_reservation=budget_reservation, ) + assert budget_reservation["finalized"] is True proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() From fce86d13342dd3111d5a7659a3e720e4859e87f5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 17:08:45 -0700 Subject: [PATCH 15/31] fix budget reservation greptile findings --- .../proxy/hooks/proxy_track_cost_callback.py | 3 +- litellm/proxy/proxy_server.py | 19 +-- .../spend_tracking/budget_reservation.py | 46 +++++--- .../hooks/test_proxy_track_cost_callback.py | 1 + .../proxy/test_budget_reservation.py | 108 ++++++++++++++++-- 5 files changed, 141 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 86f7849bb15..56b471c290b 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -489,11 +489,12 @@ async def _update_database_and_spend_counters( await _invalidate_budget_reservation_counters( budget_reservation=budget_reservation ) - budget_reservation["finalized"] = True except Exception: verbose_proxy_logger.exception( "Failed to invalidate budget reservation counters after spend counter update failed" ) + finally: + budget_reservation["finalized"] = True raise diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2f2b59db679..dd46f09fbca 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1993,8 +1993,9 @@ async def _increment_end_user_and_tag_spend_counters( reserved_counter_keys: Set[str], ) -> None: if end_user_id is not None: - await _increment_warm_unreserved_spend_counter( + await _init_and_increment_unreserved_spend_counter( counter_key=f"spend:end_user:{end_user_id}", + source_cache_key=f"end_user_id:{end_user_id}", increment=response_cost, reserved_counter_keys=reserved_counter_keys, ) @@ -2007,26 +2008,14 @@ async def _increment_end_user_and_tag_spend_counters( if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags: continue seen_tags.add(tag_name) - await _increment_warm_unreserved_spend_counter( + await _init_and_increment_unreserved_spend_counter( counter_key=f"spend:tag:{tag_name}", + source_cache_key=f"tag:{tag_name}", increment=response_cost, reserved_counter_keys=reserved_counter_keys, ) -async def _increment_warm_unreserved_spend_counter( - counter_key: str, - increment: float, - reserved_counter_keys: Set[str], -) -> None: - if counter_key in reserved_counter_keys: - return - if await spend_counter_cache.async_get_cache(key=counter_key) is None: - return - - await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) - - async def _increment_org_spend_counter( org_id: Optional[str], response_cost: float, diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index b17251f4b55..3cc532def63 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -97,7 +97,10 @@ async def reserve_budget_for_request( applied_entries: List[Dict[str, Any]] = [] try: for counter in counters: - entry = _counter_to_reservation_entry(counter) + entry = _counter_to_reservation_entry( + counter=counter, + reserved_cost=reservation_cost, + ) reserved_value = await _reserve_counter( counter=counter, reservation_cost=reservation_cost, @@ -135,9 +138,10 @@ async def reserve_budget_for_request( ), ) except Exception: - await _set_reserved_entries_adjustment( + await _set_reserved_entries_actual_cost( entries=applied_entries, - target_adjustment=-reservation_cost, + actual_cost=0.0, + default_reserved_cost=reservation_cost, ) raise @@ -157,10 +161,10 @@ async def reconcile_budget_reservation( reserved_cost = float(budget_reservation.get("reserved_cost") or 0.0) actual = float(actual_cost or 0.0) - adjustment = actual - reserved_cost - await _set_reserved_entries_adjustment( + await _set_reserved_entries_actual_cost( entries=budget_reservation.get("entries") or [], - target_adjustment=adjustment, + actual_cost=actual, + default_reserved_cost=reserved_cost, ) budget_reservation["finalized"] = True @@ -624,9 +628,10 @@ async def _ensure_malformed_window_counter_initialized( ) -async def _set_reserved_entries_adjustment( +async def _set_reserved_entries_actual_cost( entries: List[dict], - target_adjustment: float, + actual_cost: float, + default_reserved_cost: float, ) -> None: from litellm.proxy.proxy_server import _increment_spend_counter_cache @@ -634,6 +639,11 @@ async def _set_reserved_entries_adjustment( counter_key = entry.get("counter_key") if counter_key is None: continue + reserved_cost = _get_entry_reserved_cost( + entry=entry, + default_reserved_cost=default_reserved_cost, + ) + target_adjustment = actual_cost - reserved_cost applied_adjustment = float(entry.get("applied_adjustment") or 0.0) adjustment = target_adjustment - applied_adjustment if adjustment == 0: @@ -650,23 +660,33 @@ async def _resize_applied_reservation( current_reserved_cost: float, new_reserved_cost: float, ) -> None: - await _set_reserved_entries_adjustment( + await _set_reserved_entries_actual_cost( entries=entries, - target_adjustment=new_reserved_cost - current_reserved_cost, + actual_cost=new_reserved_cost, + default_reserved_cost=current_reserved_cost, ) - for entry in entries: - entry["applied_adjustment"] = 0.0 -def _counter_to_reservation_entry(counter: _BudgetCounter) -> Dict[str, Any]: +def _counter_to_reservation_entry( + counter: _BudgetCounter, + reserved_cost: float, +) -> Dict[str, Any]: return { "counter_key": counter.counter_key, "entity_type": counter.entity_type, "entity_id": counter.entity_id, + "reserved_cost": reserved_cost, "applied_adjustment": 0.0, } +def _get_entry_reserved_cost(entry: dict, default_reserved_cost: float) -> float: + try: + return float(entry.get("reserved_cost", default_reserved_cost) or 0.0) + except (TypeError, ValueError): + return default_reserved_cost + + def get_budget_window_start(window: Any) -> Optional[datetime]: window_dict = _coerce_window(window) budget_duration = window_dict.get("budget_duration") diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 351d3059d4e..01cfd1708aa 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -500,6 +500,7 @@ async def test_update_database_and_spend_counters_preserves_counter_exception_wh mock_log_exception.assert_called_once_with( "Failed to invalidate budget reservation counters after spend counter update failed" ) + assert budget_reservation["finalized"] is True proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 51016126d35..1ae47d268f8 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -245,6 +245,7 @@ async def test_should_prevent_second_tag_reservation_over_budget( "counter_key": "spend:tag:tag-budget-race", "entity_type": "Tag", "entity_id": "tag-budget-race", + "reserved_cost": 0.6, "applied_adjustment": 0.0, } ] @@ -290,16 +291,33 @@ async def test_should_prevent_second_tag_reservation_over_budget( @pytest.mark.asyncio -async def test_should_update_warm_end_user_and_tag_counters_without_reservation( +async def test_should_seed_and_update_end_user_and_tag_counters_without_reservation( spend_counter_state, ): - counter_cache, _ = spend_counter_state - counter_cache.in_memory_cache.set_cache( - key="spend:end_user:customer-1", - value=4.0, + counter_cache, key_cache = spend_counter_state + await key_cache.async_set_cache( + key="end_user_id:customer-1", + value=LiteLLM_EndUserTable( + user_id="customer-1", + blocked=False, + spend=4.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=10.0), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="tag:paid-tag", + value=LiteLLM_TagTable( + tag_name="paid-tag", + spend=7.0, + ).model_dump(), + ) + await key_cache.async_set_cache( + key="tag:other-tag", + value=LiteLLM_TagTable( + tag_name="other-tag", + spend=2.0, + ).model_dump(), ) - counter_cache.in_memory_cache.set_cache(key="spend:tag:paid-tag", value=7.0) - counter_cache.in_memory_cache.set_cache(key="spend:tag:other-tag", value=2.0) from litellm.proxy.proxy_server import increment_spend_counters @@ -539,6 +557,82 @@ async def test_should_shrink_uncapped_reservation_when_counter_advances( ) == pytest.approx(0.3) +@pytest.mark.asyncio +async def test_should_shrink_uncapped_reservation_multiple_times( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-double-resize", + spend=0.2, + max_budget=1.0, + team_id="team-budget-double-resize", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-double-resize", + spend=0.2, + max_budget=1.0, + ) + request_body = _request_body() + request_body.pop("max_tokens") + + from litellm.proxy.spend_tracking import budget_reservation + + stale_spend_by_counter_key = { + "spend:key:key-budget-double-resize": 0.3, + "spend:team:team-budget-double-resize": 0.4, + } + + async def stale_counter_read(counter): + await counter_cache.async_increment_cache( + key=counter.counter_key, + value=stale_spend_by_counter_key[counter.counter_key], + ) + return 0.2 + + monkeypatch.setattr( + budget_reservation, + "_get_current_counter_value", + stale_counter_read, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=None, + ): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.6) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-double-resize" + ) == pytest.approx(0.9) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-double-resize" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-double-resize" + ) == pytest.approx(0.3) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-double-resize" + ) == pytest.approx(0.4) + + def test_should_start_window_without_reset_at_at_duration_boundary(): before = datetime.now(timezone.utc) - timedelta(hours=1) From 9db8ecac12fb736c26870eeb83a352d8a6f4bc5e Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 17:13:49 -0700 Subject: [PATCH 16/31] update budget reservation auth test expectation --- tests/test_litellm/proxy/auth/test_user_api_key_auth.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 28a131baa1e..8fa7dc67a6c 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -2218,6 +2218,7 @@ async def test_centralized_common_checks_reserves_request_end_user_budget(): "counter_key": "spend:end_user:alice", "entity_type": "EndUser", "entity_id": "alice", + "reserved_cost": 0.6, "applied_adjustment": 0.0, } ] From 694fadd175480ac6c93e1e7810ac28e2aa64835b Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 17:38:18 -0700 Subject: [PATCH 17/31] fix budget reservation review findings --- litellm/proxy/auth/auth_checks.py | 21 +++++++-- litellm/proxy/db/spend_counter_reseed.py | 23 ++++++++-- .../proxy/hooks/proxy_track_cost_callback.py | 10 +++++ .../proxy/auth/test_user_api_key_auth.py | 43 +++++++++++++++++++ .../hooks/test_proxy_track_cost_callback.py | 24 +++++++++-- tests/test_litellm/proxy/test_proxy_server.py | 19 ++++++++ 6 files changed, 130 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index d9d9da5c93d..06b28142a7d 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1592,9 +1592,12 @@ async def _cache_key_object( ## CACHE REFRESH TIME user_api_key_obj.last_refreshed_at = time.time() + cached_key_obj = _copy_user_api_key_auth_for_cache( + user_api_key_obj=user_api_key_obj + ) await _cache_management_object( key=key, - value=user_api_key_obj, + value=cached_key_obj, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -2318,9 +2321,11 @@ async def get_key_object( if cached_key_obj is not None: if isinstance(cached_key_obj, dict): - return UserAPIKeyAuth(**cached_key_obj) + return _copy_user_api_key_auth_for_cache( + user_api_key_obj=UserAPIKeyAuth(**cached_key_obj) + ) elif isinstance(cached_key_obj, UserAPIKeyAuth): - return cached_key_obj + return _copy_user_api_key_auth_for_cache(user_api_key_obj=cached_key_obj) if check_cache_only: raise Exception( @@ -2373,6 +2378,16 @@ async def get_key_object( return _response +def _copy_user_api_key_auth_for_cache( + user_api_key_obj: UserAPIKeyAuth, +) -> UserAPIKeyAuth: + copied_key_obj = user_api_key_obj.model_copy() + copied_key_obj.budget_reservation = None + copied_key_obj.parent_otel_span = None + copied_key_obj.request_route = None + return copied_key_obj + + @log_db_metrics async def get_object_permission( object_permission_id: str, diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 1a2f64d988d..0137acf5ed9 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, ClassVar, Optional from litellm._logging import verbose_proxy_logger from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE +from litellm.litellm_core_utils.duration_parser import duration_in_seconds if TYPE_CHECKING: from litellm.caching.dual_cache import DualCache @@ -72,9 +73,10 @@ class SpendCounterReseed: """ if prisma_client is None: return None - # Per-window counters share prefixes with primary counters but - # don't correspond to a DB row. - if ":window:" in counter_key: + # Per-window key/team counters share prefixes with primary counters + # but don't correspond to a DB row. Do not reject arbitrary entity IDs + # or tag names that merely contain ":window:". + if SpendCounterReseed._is_key_or_team_window_counter(counter_key): return None try: if counter_key.startswith("spend:key:"): @@ -126,6 +128,21 @@ class SpendCounterReseed: return None return float(getattr(row, "spend", 0.0) or 0.0) + @staticmethod + def _is_key_or_team_window_counter(counter_key: str) -> bool: + for prefix in ("spend:key:", "spend:team:"): + if not counter_key.startswith(prefix): + continue + _, separator, duration = counter_key.rpartition(":window:") + if not separator or not duration: + return False + try: + duration_in_seconds(duration) + except Exception: + return False + return True + return False + @staticmethod async def coalesced( prisma_client: Optional["PrismaClient"], diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 56b471c290b..82f5a554ca3 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -45,6 +45,16 @@ class _ProxyDBLogger(CustomLogger): verbose_proxy_logger.exception( "Failed to release budget reservation during failure handling" ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=user_api_key_dict.budget_reservation + ) + if user_api_key_dict.budget_reservation is not None: + user_api_key_dict.budget_reservation["finalized"] = True + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after failure release failed" + ) request_route = user_api_key_dict.request_route if _ProxyDBLogger._should_track_errors_in_db() is False: diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 8fa7dc67a6c..83e4788fa16 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -23,6 +23,7 @@ from litellm.proxy._types import ( JWTRoutingOverride, ) from litellm.proxy.auth.handle_jwt import JWTHandler +from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( _reserve_budget_after_common_checks, @@ -76,6 +77,48 @@ async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): assert user_api_key_auth_obj.budget_reservation is None +@pytest.mark.asyncio +async def test_should_not_reuse_cached_key_object_for_request_state(): + key_cache = DualCache() + cached_key = UserAPIKeyAuth( + token="cached-token", + request_route="/old-route", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:cached-token"}], + }, + ) + + await _cache_key_object( + hashed_token="cached-token", + user_api_key_obj=cached_key, + user_api_key_cache=key_cache, + proxy_logging_obj=None, + ) + + first_request_key = await get_key_object( + hashed_token="cached-token", + prisma_client=MagicMock(), + user_api_key_cache=key_cache, + ) + first_request_key.budget_reservation = { + "reserved_cost": 0.9, + "entries": [{"counter_key": "spend:key:cached-token"}], + } + first_request_key.request_route = "/chat/completions" + + second_request_key = await get_key_object( + hashed_token="cached-token", + prisma_client=MagicMock(), + user_api_key_cache=key_cache, + ) + + assert first_request_key is not cached_key + assert second_request_key is not first_request_key + assert second_request_key.budget_reservation is None + assert second_request_key.request_route is None + + @pytest.mark.asyncio async def test_custom_auth_does_not_enforce_key_model_access_by_default(): valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"]) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 01cfd1708aa..482ddf74757 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -153,8 +153,10 @@ async def test_async_post_call_failure_hook_releases_budget_reservation_before_r user_api_key_dict=user_api_key_dict, ) - mock_release_budget_reservation.assert_awaited_once_with( - budget_reservation=budget_reservation, + assert mock_release_budget_reservation.await_count == 1 + assert ( + mock_release_budget_reservation.await_args.kwargs["budget_reservation"] + is user_api_key_dict.budget_reservation ) mock_update_database.assert_not_called() @@ -177,6 +179,10 @@ async def test_should_continue_failure_tracking_when_budget_release_fails(): new_callable=AsyncMock, side_effect=RuntimeError("redis unavailable"), ) as mock_release_budget_reservation, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback._invalidate_budget_reservation_counters", + new_callable=AsyncMock, + ) as mock_invalidate_budget_reservation_counters, patch( "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", new_callable=AsyncMock, @@ -194,9 +200,19 @@ async def test_should_continue_failure_tracking_when_budget_release_fails(): user_api_key_dict=user_api_key_dict, ) - mock_release_budget_reservation.assert_awaited_once_with( - budget_reservation=budget_reservation, + assert mock_release_budget_reservation.await_count == 1 + assert ( + mock_release_budget_reservation.await_args.kwargs["budget_reservation"] + is user_api_key_dict.budget_reservation ) + assert mock_invalidate_budget_reservation_counters.await_count == 1 + assert ( + mock_invalidate_budget_reservation_counters.await_args.kwargs[ + "budget_reservation" + ] + is user_api_key_dict.budget_reservation + ) + assert user_api_key_dict.budget_reservation["finalized"] is True mock_log_exception.assert_called_once() mock_update_database.assert_called_once() diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 9748963f90d..3dfd68a5fff 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5088,11 +5088,30 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): where={"user_id": "customer-1"} ) + fake_prisma.db.litellm_endusertable.find_unique.reset_mock() + assert ( + await SpendCounterReseed.from_db( + fake_prisma, "spend:end_user:customer:window:1h" + ) + == 21.0 + ) + fake_prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with( + where={"user_id": "customer:window:1h"} + ) + assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") == 8.0 fake_prisma.db.litellm_tagtable.find_unique.assert_awaited_once_with( where={"tag_name": "paid-tag"} ) + fake_prisma.db.litellm_tagtable.find_unique.reset_mock() + assert ( + await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid:window:1h") == 8.0 + ) + fake_prisma.db.litellm_tagtable.find_unique.assert_awaited_once_with( + where={"tag_name": "paid:window:1h"} + ) + assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0 fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with( where={"organization_id": "acme"} From 38ebd4de3debfe948c8240a5eb42af22c132f348 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 18:07:54 -0700 Subject: [PATCH 18/31] harden partial budget reservation cleanup --- .../spend_tracking/budget_reservation.py | 108 ++++++++++-- .../proxy/test_budget_reservation.py | 161 ++++++++++++++++++ 2 files changed, 252 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 3cc532def63..b5bafb1cc35 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -138,9 +138,8 @@ async def reserve_budget_for_request( ), ) except Exception: - await _set_reserved_entries_actual_cost( + await _release_applied_entries_best_effort( entries=applied_entries, - actual_cost=0.0, default_reserved_cost=reservation_cost, ) raise @@ -633,26 +632,101 @@ async def _set_reserved_entries_actual_cost( actual_cost: float, default_reserved_cost: float, ) -> None: - from litellm.proxy.proxy_server import _increment_spend_counter_cache - for entry in entries: - counter_key = entry.get("counter_key") - if counter_key is None: - continue - reserved_cost = _get_entry_reserved_cost( + await _set_reserved_entry_actual_cost( entry=entry, + actual_cost=actual_cost, default_reserved_cost=default_reserved_cost, ) - target_adjustment = actual_cost - reserved_cost - applied_adjustment = float(entry.get("applied_adjustment") or 0.0) - adjustment = target_adjustment - applied_adjustment - if adjustment == 0: - continue - await _increment_spend_counter_cache( - counter_key=counter_key, - increment=adjustment, + + +async def _set_reserved_entry_actual_cost( + entry: dict, + actual_cost: float, + default_reserved_cost: float, +) -> None: + from litellm.proxy.proxy_server import _increment_spend_counter_cache + + counter_key = entry.get("counter_key") + if counter_key is None: + return + reserved_cost = _get_entry_reserved_cost( + entry=entry, + default_reserved_cost=default_reserved_cost, + ) + target_adjustment = actual_cost - reserved_cost + applied_adjustment = float(entry.get("applied_adjustment") or 0.0) + adjustment = target_adjustment - applied_adjustment + if adjustment == 0: + return + await _ensure_counter_can_apply_adjustment( + counter_key=counter_key, + adjustment=adjustment, + ) + await _increment_spend_counter_cache( + counter_key=counter_key, + increment=adjustment, + ) + entry["applied_adjustment"] = target_adjustment + + +async def _ensure_counter_can_apply_adjustment( + counter_key: str, + adjustment: float, +) -> None: + from litellm.proxy.proxy_server import ( + _invalidate_spend_counter, + spend_counter_cache, + ) + + current_value = await spend_counter_cache.async_get_cache(key=counter_key) + if current_value is None: + await _invalidate_spend_counter(counter_key=counter_key) + raise RuntimeError( + f"Cannot apply budget reservation adjustment to missing counter {counter_key}" ) - entry["applied_adjustment"] = target_adjustment + + try: + current_float = float(current_value) + except (TypeError, ValueError): + await _invalidate_spend_counter(counter_key=counter_key) + raise RuntimeError( + f"Cannot apply budget reservation adjustment to non-numeric counter {counter_key}" + ) + + if adjustment < 0 and current_float + adjustment < -1e-12: + await _invalidate_spend_counter(counter_key=counter_key) + raise RuntimeError( + f"Budget reservation adjustment would make counter negative {counter_key}" + ) + + +async def _release_applied_entries_best_effort( + entries: List[dict], + default_reserved_cost: float, +) -> None: + for entry in entries: + try: + await _set_reserved_entry_actual_cost( + entry=entry, + actual_cost=0.0, + default_reserved_cost=default_reserved_cost, + ) + except Exception: + counter_key = entry.get("counter_key") + verbose_proxy_logger.exception( + "Failed to release partial budget reservation during exception cleanup" + ) + if counter_key is None: + continue + try: + from litellm.proxy.proxy_server import _invalidate_spend_counter + + await _invalidate_spend_counter(counter_key=counter_key) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate partial budget reservation counter during exception cleanup" + ) async def _resize_applied_reservation( diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 1ae47d268f8..3b9454d5d80 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -978,6 +978,167 @@ async def test_should_retry_partial_release_without_double_decrement( ) == pytest.approx(0.0) +@pytest.mark.asyncio +async def test_should_preserve_budget_error_and_continue_partial_cleanup( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-cleanup-failure", + spend=0.0, + max_budget=1.0, + team_id="team-budget-cleanup-failure", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-cleanup-failure", + spend=0.0, + max_budget=0.3, + ) + + original_increment_cache = counter_cache.async_increment_cache + fail_key_cleanup = True + + async def flaky_increment_cache(key, value, *args, **kwargs): + nonlocal fail_key_cleanup + if key == "spend:key:key-budget-cleanup-failure" and value < 0: + if fail_key_cleanup: + fail_key_cleanup = False + raise RuntimeError("simulated cleanup failure") + return await original_increment_cache(key=key, value=value, *args, **kwargs) + + monkeypatch.setattr(counter_cache, "async_increment_cache", flaky_increment_cache) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.4, + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.verbose_proxy_logger.exception" + ) as mock_log_exception, + ): + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-cleanup-failure" + ) + is None + ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-cleanup-failure" + ) == pytest.approx(0.0) + mock_log_exception.assert_called() + + +@pytest.mark.asyncio +async def test_should_not_create_negative_counter_when_release_counter_is_missing( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + reservation = { + "reserved_cost": 0.4, + "entries": [ + { + "counter_key": "spend:key:key-budget-missing-release", + "reserved_cost": 0.4, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with pytest.raises(RuntimeError, match="missing counter"): + await release_budget_reservation(reservation) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-missing-release" + ) + is None + ) + assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_should_invalidate_counter_when_release_would_underflow( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + await counter_cache.async_increment_cache( + key="spend:key:key-budget-underflow-release", + value=0.1, + ) + reservation = { + "reserved_cost": 0.4, + "entries": [ + { + "counter_key": "spend:key:key-budget-underflow-release", + "reserved_cost": 0.4, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with pytest.raises(RuntimeError, match="negative"): + await release_budget_reservation(reservation) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-underflow-release" + ) + is None + ) + assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_should_invalidate_non_numeric_counter_during_release( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-budget-nonnumeric-release", + value="stale", + ) + reservation = { + "reserved_cost": 0.4, + "entries": [ + { + "counter_key": "spend:key:key-budget-nonnumeric-release", + "reserved_cost": 0.4, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with pytest.raises(RuntimeError, match="non-numeric"): + await release_budget_reservation(reservation) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-nonnumeric-release" + ) + is None + ) + assert reservation["finalized"] is False + + @pytest.mark.asyncio async def test_should_invalidate_reserved_counters_after_persisted_spend_failure( spend_counter_state, From 46068be6f65ed928104b8811cf22050646e3a59b Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 18:29:14 -0700 Subject: [PATCH 19/31] skip invalid budget window reservations --- .../spend_tracking/budget_reservation.py | 50 ++----- .../proxy/test_budget_reservation.py | 132 ++++-------------- 2 files changed, 40 insertions(+), 142 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index b5bafb1cc35..956d2e80818 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -31,7 +31,6 @@ class _BudgetCounter: source_cache_key: Optional[str] = None spend_log_entity_id: Optional[str] = None window_start: Optional[datetime] = None - parent_counter_key: Optional[str] = None def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: @@ -495,16 +494,23 @@ def _get_budget_limit_counters( if not budget_duration or max_budget is None or max_budget <= 0: continue window_start = get_budget_window_start(window_dict) + if window_start is None: + verbose_proxy_logger.warning( + "Skipping budget window with invalid duration for %s=%s: %s", + entity_type, + entity_id, + budget_duration, + ) + continue counters.append( _BudgetCounter( counter_key=f"{entity_prefix}:window:{budget_duration}", max_budget=float(max_budget), - fallback_spend=fallback_spend if window_start is None else 0.0, + fallback_spend=0.0, entity_type=entity_type, entity_id=f"{entity_id}:{budget_duration}", spend_log_entity_id=entity_id, window_start=window_start, - parent_counter_key=entity_prefix if window_start is None else None, ) ) return counters @@ -572,8 +578,6 @@ async def _reserve_counter( entity_id=counter.spend_log_entity_id, window_start=counter.window_start, ) - elif counter.parent_counter_key is not None: - await _ensure_malformed_window_counter_initialized(counter=counter) reserved_value = await _increment_spend_counter_cache( counter_key=counter.counter_key, @@ -585,48 +589,12 @@ async def _reserve_counter( async def _get_current_counter_value(counter: _BudgetCounter) -> float: from litellm.proxy.proxy_server import get_current_spend - if counter.parent_counter_key is not None: - await _ensure_malformed_window_counter_initialized(counter=counter) - return await get_current_spend( counter_key=counter.counter_key, fallback_spend=counter.fallback_spend, ) -async def _ensure_malformed_window_counter_initialized( - counter: _BudgetCounter, -) -> None: - if counter.parent_counter_key is None: - return - - from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed - from litellm.proxy.proxy_server import ( - _increment_spend_counter_cache, - get_current_spend, - spend_counter_cache, - ) - - current = await spend_counter_cache.async_get_cache(key=counter.counter_key) - if current is not None: - return - - lock = await SpendCounterReseed._get_lock(counter.counter_key) - async with lock: - current = await spend_counter_cache.async_get_cache(key=counter.counter_key) - if current is not None: - return - - parent_spend = await get_current_spend( - counter_key=counter.parent_counter_key, - fallback_spend=counter.fallback_spend, - ) - await _increment_spend_counter_cache( - counter_key=counter.counter_key, - increment=parent_spend, - ) - - async def _set_reserved_entries_actual_cost( entries: List[dict], actual_cost: float, diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 3b9454d5d80..dde1056e416 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -644,14 +644,15 @@ def test_should_start_window_without_reset_at_at_duration_boundary(): @pytest.mark.asyncio -async def test_should_seed_malformed_window_counter_from_parent_authoritative_spend( +async def test_should_skip_budget_window_with_unparseable_duration( spend_counter_state, ): counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) valid_token = UserAPIKeyAuth( token="key-budget-malformed-window", - spend=0.0, + spend=0.9, + max_budget=10.0, budget_limits=[ { "budget_duration": "not-a-duration", @@ -659,116 +660,45 @@ async def test_should_seed_malformed_window_counter_from_parent_authoritative_sp } ], ) - - import litellm.proxy.proxy_server as ps - - db_row = MagicMock() - db_row.spend = 0.9 - prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( - return_value=db_row + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-budget-malformed-window", + value=0.9, ) - ps.prisma_client = prisma_client with patch( "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", return_value=0.2, ): - with pytest.raises(litellm.BudgetExceededError): - await reserve_budget_for_request( - request_body=_request_body(), - route="/chat/completions", - llm_router=None, - valid_token=valid_token, - team_object=None, - user_object=None, - prisma_client=prisma_client, - user_api_key_cache=key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - - prisma_client.db.litellm_verificationtoken.find_unique.assert_awaited_once_with( - where={"token": "key-budget-malformed-window"} - ) - assert counter_cache.in_memory_cache.get_cache( - key="spend:key:key-budget-malformed-window:window:not-a-duration" - ) == pytest.approx(0.9) - - -@pytest.mark.asyncio -async def test_should_coalesce_malformed_window_counter_initialization( - spend_counter_state, -): - import asyncio as _asyncio - - counter_cache, key_cache = spend_counter_state - proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) - token = "key-budget-malformed-window-concurrent" - malformed_window_counter_key = f"spend:key:{token}:window:not-a-duration" - valid_token = UserAPIKeyAuth( - token=token, - spend=0.0, - budget_limits=[ - { - "budget_duration": "not-a-duration", - "max_budget": 1.0, - } - ], - ) - - import litellm.proxy.proxy_server as ps - - db_call_count = 0 - - async def slow_find_unique(**kwargs): - nonlocal db_call_count - db_call_count += 1 - await _asyncio.sleep(0.05) - db_row = MagicMock() - db_row.spend = 0.35 - return db_row - - prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( - side_effect=slow_find_unique - ) - ps.prisma_client = prisma_client - - with patch( - "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", - return_value=0.2, - ): - results = await _asyncio.gather( - *[ - reserve_budget_for_request( - request_body=_request_body(), - route="/chat/completions", - llm_router=None, - valid_token=valid_token, - team_object=None, - user_object=None, - prisma_client=prisma_client, - user_api_key_cache=key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - for _ in range(2) - ], - return_exceptions=True, + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, ) - assert not any(isinstance(result, Exception) for result in results), results - assert all(result is not None for result in results) - assert db_call_count == 1 + assert reservation is not None + assert [entry["counter_key"] for entry in reservation["entries"]] == [ + "spend:key:key-budget-malformed-window" + ] assert counter_cache.in_memory_cache.get_cache( - key=malformed_window_counter_key - ) == pytest.approx(0.75) - - for reservation in results: - await release_budget_reservation(reservation) + key="spend:key:key-budget-malformed-window" + ) == pytest.approx(1.1) + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-malformed-window:window:not-a-duration" + ) + is None + ) + await release_budget_reservation(reservation) assert counter_cache.in_memory_cache.get_cache( - key=malformed_window_counter_key - ) == pytest.approx(0.35) + key="spend:key:key-budget-malformed-window" + ) == pytest.approx(0.9) @pytest.mark.asyncio From 405de4632918f13a610166c1cf1f1a9a30a6bc50 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 19:24:48 -0700 Subject: [PATCH 20/31] cap budget reservations to remaining headroom --- .../spend_tracking/budget_reservation.py | 22 ++- .../proxy/test_budget_reservation.py | 130 ++++++++++++++++-- 2 files changed, 132 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 956d2e80818..25aa95e1cb6 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -84,7 +84,6 @@ async def reserve_budget_for_request( route=route, llm_router=llm_router, ) - using_remaining_budget_fallback = reservation_cost is None if reservation_cost is None: reservation_cost = await _get_smallest_remaining_budget( counters=counters, @@ -114,18 +113,17 @@ async def reserve_budget_for_request( cached_spend = await _get_current_counter_value(counter=counter) current_spend = cached_spend + reservation_cost if current_spend > counter.max_budget: - if using_remaining_budget_fallback: - remaining_before_reservation = counter.max_budget - ( - current_spend - reservation_cost + remaining_before_reservation = counter.max_budget - ( + current_spend - reservation_cost + ) + if remaining_before_reservation > 1e-12: + await _resize_applied_reservation( + entries=applied_entries, + current_reserved_cost=reservation_cost, + new_reserved_cost=remaining_before_reservation, ) - if remaining_before_reservation > 0: - await _resize_applied_reservation( - entries=applied_entries, - current_reserved_cost=reservation_cost, - new_reserved_cost=remaining_before_reservation, - ) - reservation_cost = remaining_before_reservation - continue + reservation_cost = remaining_before_reservation + continue raise litellm.BudgetExceededError( current_cost=current_spend, max_budget=counter.max_budget, diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index dde1056e416..72f434ede46 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -70,7 +70,7 @@ def test_should_not_serialize_budget_reservation_on_user_api_key_auth(): @pytest.mark.asyncio -async def test_should_prevent_second_key_reservation_over_budget( +async def test_should_shrink_second_key_reservation_to_remaining_budget( spend_counter_state, ): counter_cache, key_cache = spend_counter_state @@ -102,6 +102,23 @@ async def test_should_prevent_second_key_reservation_over_budget( == 0.6 ) + second_reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert second_reservation is not None + assert second_reservation["reserved_cost"] == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-race" + ) == pytest.approx(1.0) + with pytest.raises(litellm.BudgetExceededError): await reserve_budget_for_request( request_body=_request_body(), @@ -115,15 +132,19 @@ async def test_should_prevent_second_key_reservation_over_budget( proxy_logging_obj=proxy_logging_obj, ) - assert ( - counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-race") == 0.6 - ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-race" + ) == pytest.approx(1.0) + await release_budget_reservation(second_reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-race" + ) == pytest.approx(0.6) await release_budget_reservation(reservation) @pytest.mark.asyncio -async def test_should_prevent_second_end_user_reservation_over_budget( +async def test_should_shrink_second_end_user_reservation_to_remaining_budget( spend_counter_state, ): counter_cache, key_cache = spend_counter_state @@ -160,6 +181,24 @@ async def test_should_prevent_second_end_user_reservation_over_budget( key="spend:end_user:end-user-budget-race" ) == pytest.approx(0.6) + second_reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + end_user_object=end_user_object, + ) + assert second_reservation is not None + assert second_reservation["reserved_cost"] == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(1.0) + with pytest.raises(litellm.BudgetExceededError): await reserve_budget_for_request( request_body=_request_body(), @@ -174,6 +213,11 @@ async def test_should_prevent_second_end_user_reservation_over_budget( end_user_object=end_user_object, ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(second_reservation) assert counter_cache.in_memory_cache.get_cache( key="spend:end_user:end-user-budget-race" ) == pytest.approx(0.6) @@ -195,7 +239,7 @@ async def test_should_prevent_second_end_user_reservation_over_budget( @pytest.mark.asyncio -async def test_should_prevent_second_tag_reservation_over_budget( +async def test_should_shrink_second_tag_reservation_to_remaining_budget( spend_counter_state, ): counter_cache, key_cache = spend_counter_state @@ -257,6 +301,23 @@ async def test_should_prevent_second_tag_reservation_over_budget( is None ) + second_reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=prisma_client, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert second_reservation is not None + assert second_reservation["reserved_cost"] == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(1.0) + with pytest.raises(litellm.BudgetExceededError): await reserve_budget_for_request( request_body=request_body, @@ -270,6 +331,11 @@ async def test_should_prevent_second_tag_reservation_over_budget( proxy_logging_obj=proxy_logging_obj, ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(second_reservation) assert counter_cache.in_memory_cache.get_cache( key="spend:tag:tag-budget-race" ) == pytest.approx(0.6) @@ -443,6 +509,50 @@ async def test_should_seed_org_counter_from_with_budget_cache(spend_counter_stat ) == pytest.approx(2.25) +@pytest.mark.asyncio +async def test_should_cap_known_estimate_to_remaining_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-known-estimate-cap", + spend=0.9, + max_budget=1.0, + ) + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-budget-known-estimate-cap", + value=0.9, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.1) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-known-estimate-cap" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-known-estimate-cap" + ) == pytest.approx(0.9) + + @pytest.mark.asyncio async def test_should_reserve_remaining_budget_when_output_cap_missing( spend_counter_state, @@ -923,9 +1033,13 @@ async def test_should_preserve_budget_error_and_continue_partial_cleanup( ) team_object = LiteLLM_TeamTable( team_id="team-budget-cleanup-failure", - spend=0.0, + spend=0.3, max_budget=0.3, ) + await key_cache.async_set_cache( + key="team_id:team-budget-cleanup-failure", + value=team_object, + ) original_increment_cache = counter_cache.async_increment_cache fail_key_cleanup = True @@ -970,7 +1084,7 @@ async def test_should_preserve_budget_error_and_continue_partial_cleanup( ) assert counter_cache.in_memory_cache.get_cache( key="spend:team:team-budget-cleanup-failure" - ) == pytest.approx(0.0) + ) == pytest.approx(0.3) mock_log_exception.assert_called() From 15d845c3211cc08b9214777a58e3188cf7a148b1 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 19:33:55 -0700 Subject: [PATCH 21/31] avoid stale local spend counters after redis misses --- litellm/proxy/db/spend_counter_reseed.py | 11 +- litellm/proxy/proxy_server.py | 31 ++++- tests/test_litellm/proxy/test_proxy_server.py | 115 ++++++++++++++++++ 3 files changed, 149 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index d4acb39062c..b3c9f4e7c35 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -183,7 +183,7 @@ class SpendCounterReseed: return None # Warm even when 0 so subsequent reads hit cache, not DB. try: - if require_cache_warm and spend_counter_cache.redis_cache is not None: + if spend_counter_cache.redis_cache is not None: current_value = ( await spend_counter_cache.redis_cache.async_increment( key=counter_key, @@ -273,6 +273,7 @@ class SpendCounterReseed: ) -> Optional[float]: lock = await SpendCounterReseed._get_lock(counter_key) async with lock: + redis_clean_miss = False if spend_counter_cache.redis_cache is not None: try: val = await spend_counter_cache.redis_cache.async_get_cache( @@ -280,11 +281,13 @@ class SpendCounterReseed: ) if val is not None: return float(val) + redis_clean_miss = True except Exception: pass - val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) - if val is not None: - return float(val) + if not redis_clean_miss: + val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) + if val is not None: + return float(val) window_spend = await SpendCounterReseed.window_from_spend_logs( prisma_client=prisma_client, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 889f85c40af..e67702a7976 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2102,8 +2102,8 @@ async def _ensure_spend_counter_initialized( counter_key: str, source_cache_key: str, ): - current = await spend_counter_cache.async_get_cache(key=counter_key) - if current is None: + is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) + if is_warm is False: # Shares the per-counter lock with get_current_spend. db_spend = await SpendCounterReseed.coalesced( prisma_client=prisma_client, @@ -2132,8 +2132,8 @@ async def _ensure_window_spend_counter_initialized( entity_id: str, window_start: datetime, ): - current = await spend_counter_cache.async_get_cache(key=counter_key) - if current is None: + is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) + if is_warm is False: window_spend = await SpendCounterReseed.coalesced_window( prisma_client=prisma_client, spend_counter_cache=spend_counter_cache, @@ -2146,6 +2146,29 @@ async def _ensure_window_spend_counter_initialized( await _increment_spend_counter_cache(counter_key=counter_key, increment=0.0) +async def _is_spend_counter_cache_warm(counter_key: str) -> bool: + if spend_counter_cache.redis_cache is not None: + try: + current_value = await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key, + ) + if current_value is None: + return False + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + return True + except Exception as e: + verbose_proxy_logger.debug( + "Unable to read Redis spend counter %s before initialization, falling back to in-memory: %s", + counter_key, + e, + ) + + return spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is not None + + async def _increment_spend_counter_cache(counter_key: str, increment: float): if spend_counter_cache.redis_cache is not None: try: diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 36a5895dbe7..9dfa6bc1936 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5212,6 +5212,118 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss(): ps.prisma_client = orig_prisma +@pytest.mark.asyncio +async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_spend_counter + + counter_cache = DualCache() + counter_key = "spend:team:team-stale-local" + counter_cache.in_memory_cache.set_cache(key=counter_key, value=10.0) + + redis_store: dict = {} + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + counter_cache.redis_cache = fake_redis + + db_row = MagicMock() + db_row.spend = 42.0 + fake_prisma = MagicMock() + fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma, orig_user = ( + ps.spend_counter_cache, + ps.prisma_client, + ps.user_api_key_cache, + ) + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + ps.user_api_key_cache = DualCache() + try: + await _init_and_increment_spend_counter( + counter_key=counter_key, + source_cache_key="team_id:team-stale-local", + increment=1.5, + ) + + fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( + where={"team_id": "team-stale-local"} + ) + assert redis_store[counter_key] == pytest.approx(43.5) + assert counter_cache.in_memory_cache.get_cache( + key=counter_key + ) == pytest.approx(43.5) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + ps.user_api_key_cache = orig_user + + +@pytest.mark.asyncio +async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + counter_key = "spend:key:key-window-stale-local:window:1h" + counter_cache.in_memory_cache.set_cache(key=counter_key, value=100.0) + window_start = datetime.now(timezone.utc) - timedelta(hours=1) + + redis_store: dict = {} + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + counter_cache.redis_cache = fake_redis + + fake_prisma = MagicMock() + fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( + return_value=[{"api_key": "key-window-stale-local", "_sum": {"spend": 2.25}}] + ) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_window_spend_counter( + counter_key=counter_key, + entity_type="Key", + entity_id="key-window-stale-local", + window_start=window_start, + increment=0.5, + ) + + fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with( + by=["api_key"], + where={ + "api_key": "key-window-stale-local", + "startTime": {"gte": window_start}, + }, + sum={"spend": True}, + ) + assert redis_store[counter_key] == pytest.approx(2.75) + assert counter_cache.in_memory_cache.get_cache( + key=counter_key + ) == pytest.approx(2.75) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + @pytest.mark.asyncio async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): from litellm.caching.dual_cache import DualCache @@ -5293,6 +5405,9 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing(): assert ("spend:team_member:user-1:team-1", 362.0) in [ (w["key"], w["value"]) for w in recorded_warms ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:team_member:user-1:team-1" + ) == pytest.approx(362.0) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma From c28e093f41588a6a5bd856a39f0be7ccbecfb1f2 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 19:50:36 -0700 Subject: [PATCH 22/31] finalize budget reservations after counter updates --- litellm/proxy/proxy_server.py | 45 +++++--- .../spend_tracking/budget_reservation.py | 4 +- tests/test_litellm/proxy/test_proxy_server.py | 102 ++++++++++++++++++ 3 files changed, 136 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e67702a7976..3eb536d0850 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1854,22 +1854,14 @@ async def increment_spend_counters( Awaited (not create_task) in the cost callback, so the counter is updated before the next request's auth check runs. """ - reserved_counter_keys: Set[str] = set() - if budget_reservation is not None: - from litellm.proxy.spend_tracking.budget_reservation import ( - get_reserved_counter_keys, - reconcile_budget_reservation, - ) - - reserved_counter_keys = get_reserved_counter_keys( - budget_reservation=budget_reservation - ) - await reconcile_budget_reservation( - budget_reservation=budget_reservation, - actual_cost=response_cost or 0.0, - ) + reserved_counter_keys = await _reconcile_budget_reservation_for_counter_update( + budget_reservation=budget_reservation, + response_cost=response_cost, + ) if response_cost is None or response_cost == 0: + if budget_reservation is not None: + budget_reservation["finalized"] = True return if token is not None: @@ -1989,6 +1981,31 @@ async def increment_spend_counters( response_cost=response_cost, reserved_counter_keys=reserved_counter_keys, ) + if budget_reservation is not None: + budget_reservation["finalized"] = True + + +async def _reconcile_budget_reservation_for_counter_update( + budget_reservation: Optional[dict], + response_cost: Optional[float], +) -> Set[str]: + if budget_reservation is None: + return set() + + from litellm.proxy.spend_tracking.budget_reservation import ( + get_reserved_counter_keys, + reconcile_budget_reservation, + ) + + reserved_counter_keys = get_reserved_counter_keys( + budget_reservation=budget_reservation + ) + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=response_cost or 0.0, + finalize=False, + ) + return reserved_counter_keys async def _increment_end_user_and_tag_spend_counters( diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 25aa95e1cb6..77c6f9ae63e 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -151,6 +151,7 @@ async def reserve_budget_for_request( async def reconcile_budget_reservation( budget_reservation: Optional[dict], actual_cost: Optional[float], + finalize: bool = True, ) -> None: if not budget_reservation or budget_reservation.get("finalized") is True: return @@ -162,7 +163,8 @@ async def reconcile_budget_reservation( actual_cost=actual, default_reserved_cost=reserved_cost, ) - budget_reservation["finalized"] = True + if finalize: + budget_reservation["finalized"] = True async def release_budget_reservation(budget_reservation: Optional[dict]) -> None: diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 9dfa6bc1936..dd6ef04bc93 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5324,6 +5324,108 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): ps.prisma_client = orig_prisma +@pytest.mark.asyncio +async def test_increment_spend_counters_finalizes_after_unreserved_increments(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import increment_spend_counters + + counter_cache = DualCache() + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-finalize-after-increments", + value=0.5, + ) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:key-finalize-after-increments", + "entity_type": "Key", + "entity_id": "key-finalize-after-increments", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + incremented_counters = [] + + async def assert_reservation_not_finalized_yet(**kwargs): + assert budget_reservation["finalized"] is False + incremented_counters.append(kwargs["counter_key"]) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_user = ps.spend_counter_cache, ps.user_api_key_cache + ps.spend_counter_cache = counter_cache + ps.user_api_key_cache = DualCache() + try: + with patch( + "litellm.proxy.proxy_server._init_and_increment_spend_counter", + new=AsyncMock(side_effect=assert_reservation_not_finalized_yet), + ): + await increment_spend_counters( + token="key-finalize-after-increments", + team_id="team-finalize-after-increments", + user_id=None, + response_cost=0.25, + budget_reservation=budget_reservation, + ) + + assert incremented_counters == ["spend:team:team-finalize-after-increments"] + assert budget_reservation["finalized"] is True + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-finalize-after-increments" + ) == pytest.approx(0.25) + finally: + ps.spend_counter_cache = orig_counter + ps.user_api_key_cache = orig_user + + +@pytest.mark.asyncio +async def test_increment_spend_counters_finalizes_none_cost_reservation(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import increment_spend_counters + + counter_cache = DualCache() + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-finalize-none-cost", + value=0.5, + ) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:key-finalize-none-cost", + "entity_type": "Key", + "entity_id": "key-finalize-none-cost", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + await increment_spend_counters( + token="key-finalize-none-cost", + team_id=None, + user_id=None, + response_cost=None, + budget_reservation=budget_reservation, + ) + + assert budget_reservation["finalized"] is True + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-finalize-none-cost" + ) == pytest.approx(0.0) + finally: + ps.spend_counter_cache = orig_counter + + @pytest.mark.asyncio async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): from litellm.caching.dual_cache import DualCache From dcfde1b89952d5417d8f9a4af7f7a8e1bac8d20d Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 20:04:16 -0700 Subject: [PATCH 23/31] fallback to plain org cache for spend counters --- litellm/proxy/proxy_server.py | 34 +++++++++++++------ .../proxy/test_budget_reservation.py | 31 +++++++++++++++++ 2 files changed, 54 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3eb536d0850..927d0ae8e8a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2048,7 +2048,7 @@ async def _increment_org_spend_counter( await _init_and_increment_unreserved_spend_counter( counter_key=f"spend:org:{org_id}", - source_cache_key=f"org_id:{org_id}:with_budget", + source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"], increment=response_cost, reserved_counter_keys=reserved_counter_keys, ) @@ -2056,7 +2056,7 @@ async def _increment_org_spend_counter( async def _init_and_increment_unreserved_spend_counter( counter_key: str, - source_cache_key: str, + source_cache_key: Union[str, List[str]], increment: float, reserved_counter_keys: Set[str], ) -> None: @@ -2072,7 +2072,7 @@ async def _init_and_increment_unreserved_spend_counter( async def _init_and_increment_spend_counter( counter_key: str, - source_cache_key: str, + source_cache_key: Union[str, List[str]], increment: float, ): """ @@ -2117,7 +2117,7 @@ async def _init_and_increment_window_spend_counter( async def _ensure_spend_counter_initialized( counter_key: str, - source_cache_key: str, + source_cache_key: Union[str, List[str]], ): is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) if is_warm is False: @@ -2130,19 +2130,31 @@ async def _ensure_spend_counter_initialized( ) if db_spend is None: # DB unavailable - fall back to in-process cache (may be stale). - source = await user_api_key_cache.async_get_cache(key=source_cache_key) - base_spend: float = 0.0 - if source is not None: - if isinstance(source, dict): - base_spend = source.get("spend", 0.0) or 0.0 - else: - base_spend = getattr(source, "spend", 0.0) or 0.0 + base_spend = await _get_source_cache_base_spend( + source_cache_key=source_cache_key + ) if base_spend > 0: await _increment_spend_counter_cache( counter_key=counter_key, increment=base_spend ) +async def _get_source_cache_base_spend( + source_cache_key: Union[str, List[str]], +) -> float: + source_cache_keys = ( + [source_cache_key] if isinstance(source_cache_key, str) else source_cache_key + ) + for cache_key in source_cache_keys: + source = await user_api_key_cache.async_get_cache(key=cache_key) + if source is None: + continue + if isinstance(source, dict): + return float(source.get("spend", 0.0) or 0.0) + return float(getattr(source, "spend", 0.0) or 0.0) + return 0.0 + + async def _ensure_window_spend_counter_initialized( counter_key: str, entity_type: str, diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 72f434ede46..d1684b7ef1f 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -509,6 +509,37 @@ async def test_should_seed_org_counter_from_with_budget_cache(spend_counter_stat ) == pytest.approx(2.25) +@pytest.mark.asyncio +async def test_should_seed_org_counter_from_plain_org_cache(spend_counter_state): + counter_cache, key_cache = spend_counter_state + await key_cache.async_set_cache( + key="org_id:org-counter-plain", + value=LiteLLM_OrganizationTable( + organization_id="org-counter-plain", + organization_alias="shared-org", + budget_id="org-budget-id", + spend=2.0, + models=[], + created_by="test", + updated_by="test", + ).model_dump(), + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + org_id="org-counter-plain", + response_cost=0.25, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:org:org-counter-plain" + ) == pytest.approx(2.25) + + @pytest.mark.asyncio async def test_should_cap_known_estimate_to_remaining_budget( spend_counter_state, From 4f8769943b6e9b0f4ef2472eca41b67dba2e43ff Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 20:18:07 -0700 Subject: [PATCH 24/31] skip invalid budget window counter increments --- litellm/proxy/proxy_server.py | 18 +++++++---- tests/test_litellm/proxy/test_proxy_server.py | 30 +++++++++++++++++++ 2 files changed, 42 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 927d0ae8e8a..2672ec252d5 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2105,13 +2105,19 @@ async def _init_and_increment_window_spend_counter( window_start: Optional[datetime], increment: float, ): - if window_start is not None: - await _ensure_window_spend_counter_initialized( - counter_key=counter_key, - entity_type=entity_type, - entity_id=entity_id, - window_start=window_start, + if window_start is None: + verbose_proxy_logger.warning( + "Skipping spend counter increment for invalid budget window %s", + counter_key, ) + return + + await _ensure_window_spend_counter_initialized( + counter_key=counter_key, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index dd6ef04bc93..5edff6d903a 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5324,6 +5324,36 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): ps.prisma_client = orig_prisma +@pytest.mark.asyncio +async def test_window_spend_counter_skips_invalid_window_start(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + await _init_and_increment_window_spend_counter( + counter_key="spend:key:key-invalid-window:window:not-a-duration", + entity_type="Key", + entity_id="key-invalid-window", + window_start=None, + increment=0.5, + ) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-invalid-window:window:not-a-duration" + ) + is None + ) + finally: + ps.spend_counter_cache = orig_counter + + @pytest.mark.asyncio async def test_increment_spend_counters_finalizes_after_unreserved_increments(): from litellm.caching.dual_cache import DualCache From f30bfcf36a6621e4b31cbd21042ab22fd07ea8b5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 20:36:14 -0700 Subject: [PATCH 25/31] add budget reservation disable flag --- litellm/proxy/auth/user_api_key_auth.py | 6 +++- .../proxy/auth/test_user_api_key_auth.py | 29 +++++++++++++++++++ 2 files changed, 34 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 2495d33a5cf..ea7bad2bf47 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1909,7 +1909,7 @@ async def _reserve_budget_after_common_checks( end_user_object: Optional[LiteLLM_EndUserTable] = None, ) -> None: user_api_key_auth_obj.budget_reservation = None - if skip_budget_checks: + if skip_budget_checks or _is_budget_reservation_disabled(): return from litellm.proxy.spend_tracking.budget_reservation import ( @@ -1931,6 +1931,10 @@ async def _reserve_budget_after_common_checks( ) +def _is_budget_reservation_disabled() -> bool: + return get_secret_bool("LITELLM_DISABLE_BUDGET_RESERVATION", False) is True + + def _should_skip_budget_checks( request_data: dict, route: str, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 83e4788fa16..8489a7b0f60 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -77,6 +77,35 @@ async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): assert user_api_key_auth_obj.budget_reservation is None +@pytest.mark.asyncio +async def test_should_skip_budget_reservation_when_disabled(): + user_api_key_auth_obj = UserAPIKeyAuth( + token="test_token", + spend=0.0, + max_budget=1.0, + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_token"}], + }, + ) + + with patch.dict(os.environ, {"LITELLM_DISABLE_BUDGET_RESERVATION": "true"}): + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "gpt-4"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=False, + ) + + assert user_api_key_auth_obj.budget_reservation is None + + @pytest.mark.asyncio async def test_should_not_reuse_cached_key_object_for_request_state(): key_cache = DualCache() From ce17639cf7489c16fae33fca06d6957d8d5d487a Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 20:57:20 -0700 Subject: [PATCH 26/31] remove budget reservation disable flag --- litellm/proxy/auth/user_api_key_auth.py | 6 +--- .../proxy/auth/test_user_api_key_auth.py | 29 ------------------- 2 files changed, 1 insertion(+), 34 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ea7bad2bf47..2495d33a5cf 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1909,7 +1909,7 @@ async def _reserve_budget_after_common_checks( end_user_object: Optional[LiteLLM_EndUserTable] = None, ) -> None: user_api_key_auth_obj.budget_reservation = None - if skip_budget_checks or _is_budget_reservation_disabled(): + if skip_budget_checks: return from litellm.proxy.spend_tracking.budget_reservation import ( @@ -1931,10 +1931,6 @@ async def _reserve_budget_after_common_checks( ) -def _is_budget_reservation_disabled() -> bool: - return get_secret_bool("LITELLM_DISABLE_BUDGET_RESERVATION", False) is True - - def _should_skip_budget_checks( request_data: dict, route: str, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 8489a7b0f60..83e4788fa16 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -77,35 +77,6 @@ async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): assert user_api_key_auth_obj.budget_reservation is None -@pytest.mark.asyncio -async def test_should_skip_budget_reservation_when_disabled(): - user_api_key_auth_obj = UserAPIKeyAuth( - token="test_token", - spend=0.0, - max_budget=1.0, - budget_reservation={ - "reserved_cost": 0.5, - "entries": [{"counter_key": "spend:key:test_token"}], - }, - ) - - with patch.dict(os.environ, {"LITELLM_DISABLE_BUDGET_RESERVATION": "true"}): - await _reserve_budget_after_common_checks( - user_api_key_auth_obj=user_api_key_auth_obj, - request_data={"model": "gpt-4"}, - route="/v1/chat/completions", - llm_router=None, - team_object=None, - user_object=None, - prisma_client=None, - user_api_key_cache=DualCache(), - proxy_logging_obj=MagicMock(), - skip_budget_checks=False, - ) - - assert user_api_key_auth_obj.budget_reservation is None - - @pytest.mark.asyncio async def test_should_not_reuse_cached_key_object_for_request_state(): key_cache = DualCache() From b53adf7cff4dbbbc6814a92c101cb9ec73adea66 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 21:21:26 -0700 Subject: [PATCH 27/31] address budget reservation review edges --- litellm/litellm_core_utils/litellm_logging.py | 2 +- litellm/proxy/db/spend_counter_reseed.py | 16 ++-- .../proxy/hooks/proxy_track_cost_callback.py | 7 ++ litellm/proxy/litellm_pre_call_utils.py | 4 + litellm/proxy/proxy_server.py | 5 +- .../hooks/test_proxy_track_cost_callback.py | 41 ++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 77 ++++++++++--------- 7 files changed, 104 insertions(+), 48 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 829c1c9ca07..2e0ddb846c7 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4720,7 +4720,7 @@ class StandardLoggingPayloadSetup: ): for key, value in litellm_params["metadata"].items(): # Skip non-serializable objects like UserAPIKeyAuth - if key == "user_api_key_auth": + if key in {"user_api_key_auth", "user_api_key_budget_reservation"}: continue merged_metadata[key] = value diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index b3c9f4e7c35..f4f193300b2 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -36,9 +36,11 @@ class SpendCounterReseed: spend:team:{team_id} -> LiteLLM_TeamTable.spend spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend spend:user:{user_id} -> LiteLLM_UserTable.spend - spend:end_user:{end_user_id} -> LiteLLM_EndUserTable.spend - spend:tag:{tag_name} -> LiteLLM_TagTable.spend spend:org:{org_id} -> LiteLLM_OrganizationTable.spend + + End-user and tag spend counters intentionally do not reseed here. Their + auth paths already load the corresponding objects via get_end_user_object() + and get_tag_objects_batch(); callers pass those values as fallback_spend. """ _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict() @@ -103,15 +105,9 @@ class SpendCounterReseed: where={"user_id": user_id} ) elif counter_key.startswith("spend:end_user:"): - end_user_id = counter_key[len("spend:end_user:") :] - row = await prisma_client.db.litellm_endusertable.find_unique( - where={"user_id": end_user_id} - ) + return None elif counter_key.startswith("spend:tag:"): - tag_name = counter_key[len("spend:tag:") :] - row = await prisma_client.db.litellm_tagtable.find_unique( - where={"tag_name": tag_name} - ) + return None elif counter_key.startswith("spend:org:"): org_id = counter_key[len("spend:org:") :] row = await prisma_client.db.litellm_organizationtable.find_unique( diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 82f5a554ca3..bd1b8ea79c5 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -413,9 +413,16 @@ def _should_track_cost_callback( def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: + metadata_budget_reservation = metadata.get("user_api_key_budget_reservation") + if isinstance(metadata_budget_reservation, dict): + return metadata_budget_reservation + user_api_key_auth_obj = metadata.get("user_api_key_auth") if user_api_key_auth_obj is None: return None + if isinstance(user_api_key_auth_obj, dict): + budget_reservation = user_api_key_auth_obj.get("budget_reservation") + return budget_reservation if isinstance(budget_reservation, dict) else None return getattr(user_api_key_auth_obj, "budget_reservation", None) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 3077efe1167..853c56856fc 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -893,6 +893,10 @@ class LiteLLMProxyRequestSetup: data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr( user_api_key_dict, "end_user_max_budget", None ) + if user_api_key_dict.budget_reservation is not None: + data[_metadata_variable_name][ + "user_api_key_budget_reservation" + ] = user_api_key_dict.budget_reservation # Add the full UserAPIKeyAuth object for MCP server access control data[_metadata_variable_name]["user_api_key_auth"] = user_api_key_dict return data diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2672ec252d5..675a7d471bb 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2178,7 +2178,10 @@ async def _ensure_window_spend_counter_initialized( window_start=window_start, ) if window_spend is None: - await _increment_spend_counter_cache(counter_key=counter_key, increment=0.0) + verbose_proxy_logger.warning( + "Skipping cold spend counter seed for %s because window spend could not be loaded", + counter_key, + ) async def _is_spend_counter_cache_warm(counter_key: str) -> bool: diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 482ddf74757..8b5835139b4 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -13,6 +13,7 @@ from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.proxy_track_cost_callback import ( _ProxyDBLogger, + _get_budget_reservation_from_metadata, _update_database_and_spend_counters, ) @@ -298,6 +299,46 @@ async def test_track_cost_callback_releases_budget_reservation_when_response_cos ) +def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } + + assert ( + _get_budget_reservation_from_metadata( + metadata={"user_api_key_auth": dict(UserAPIKeyAuth())} + ) + is None + ) + assert ( + _get_budget_reservation_from_metadata( + metadata={ + "user_api_key_auth": UserAPIKeyAuth( + budget_reservation=budget_reservation + ) + } + ) + == budget_reservation + ) + assert ( + _get_budget_reservation_from_metadata( + metadata={ + "user_api_key_auth": dict( + UserAPIKeyAuth(budget_reservation=budget_reservation) + ) + } + ) + == budget_reservation + ) + assert ( + _get_budget_reservation_from_metadata( + metadata={"user_api_key_budget_reservation": budget_reservation} + ) + is budget_reservation + ) + + @pytest.mark.asyncio async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails(): proxy_logging_obj = MagicMock() diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5edff6d903a..53690907712 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5084,25 +5084,23 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( @pytest.mark.asyncio async def test_reseed_spend_from_db_user_and_org_prefixes(): - """User and org counters must reseed from their own DB tables, not - fall through to 0.0 like the other counters do today.""" + """User and org counters reseed from their own DB tables. + + End-user and tag counters use the already fetched auth objects passed as + fallback_spend, so this reseed helper must not add extra per-request DB + reads for them. + """ from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed user_row = MagicMock() user_row.spend = 17.0 - end_user_row = MagicMock() - end_user_row.spend = 21.0 - tag_row = MagicMock() - tag_row.spend = 8.0 org_row = MagicMock() org_row.spend = 305.0 fake_prisma = MagicMock() fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) - fake_prisma.db.litellm_endusertable.find_unique = AsyncMock( - return_value=end_user_row - ) - fake_prisma.db.litellm_tagtable.find_unique = AsyncMock(return_value=tag_row) + fake_prisma.db.litellm_endusertable.find_unique = AsyncMock() + fake_prisma.db.litellm_tagtable.find_unique = AsyncMock() fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock( return_value=org_row ) @@ -5112,37 +5110,17 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): where={"user_id": "alice"} ) - assert ( - await SpendCounterReseed.from_db(fake_prisma, "spend:end_user:customer-1") - == 21.0 - ) - fake_prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with( - where={"user_id": "customer-1"} - ) - - fake_prisma.db.litellm_endusertable.find_unique.reset_mock() assert ( await SpendCounterReseed.from_db( - fake_prisma, "spend:end_user:customer:window:1h" + fake_prisma, + "spend:end_user:customer-1", ) - == 21.0 - ) - fake_prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with( - where={"user_id": "customer:window:1h"} + is None ) + fake_prisma.db.litellm_endusertable.find_unique.assert_not_awaited() - assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") == 8.0 - fake_prisma.db.litellm_tagtable.find_unique.assert_awaited_once_with( - where={"tag_name": "paid-tag"} - ) - - fake_prisma.db.litellm_tagtable.find_unique.reset_mock() - assert ( - await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid:window:1h") == 8.0 - ) - fake_prisma.db.litellm_tagtable.find_unique.assert_awaited_once_with( - where={"tag_name": "paid:window:1h"} - ) + assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") is None + fake_prisma.db.litellm_tagtable.find_unique.assert_not_awaited() assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0 fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with( @@ -5354,6 +5332,33 @@ async def test_window_spend_counter_skips_invalid_window_start(): ps.spend_counter_cache = orig_counter +@pytest.mark.asyncio +async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _ensure_window_spend_counter_initialized + + counter_cache = DualCache() + counter_key = "spend:key:key-window-db-unavailable:window:1h" + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = None + try: + await _ensure_window_spend_counter_initialized( + counter_key=counter_key, + entity_type="Key", + entity_id="key-window-db-unavailable", + window_start=datetime.now(timezone.utc) - timedelta(hours=1), + ) + + assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + @pytest.mark.asyncio async def test_increment_spend_counters_finalizes_after_unreserved_increments(): from litellm.caching.dual_cache import DualCache From 0b1ea9eb8f21e283e2f1a333a3713ba8e8792b6c Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 21:49:31 -0700 Subject: [PATCH 28/31] harden budget reservation edge cases --- litellm/proxy/db/spend_counter_reseed.py | 26 +++++-- litellm/proxy/proxy_server.py | 36 +++++---- .../spend_tracking/budget_reservation.py | 29 ++++++-- .../proxy/test_budget_reservation.py | 49 +++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 73 ++++++++++++++++++- 5 files changed, 187 insertions(+), 26 deletions(-) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index f4f193300b2..19ec6699390 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -295,12 +295,28 @@ class SpendCounterReseed: return None try: if spend_counter_cache.redis_cache is not None: - current_value = ( - await spend_counter_cache.redis_cache.async_increment( - key=counter_key, - value=window_spend, - ) + seeded = await spend_counter_cache.redis_cache.async_set_cache( + key=counter_key, + value=window_spend, + nx=True, ) + if seeded: + current_value = window_spend + else: + current_cached_value = ( + await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key + ) + ) + if current_cached_value is None: + current_value = ( + await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=window_spend, + ) + ) + else: + current_value = float(current_cached_value) spend_counter_cache.in_memory_cache.set_cache( key=counter_key, value=current_value, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 675a7d471bb..e77561a5e59 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2112,12 +2112,14 @@ async def _init_and_increment_window_spend_counter( ) return - await _ensure_window_spend_counter_initialized( + initialized = await _ensure_window_spend_counter_initialized( counter_key=counter_key, entity_type=entity_type, entity_id=entity_id, window_start=window_start, ) + if initialized is False: + return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) @@ -2166,22 +2168,26 @@ async def _ensure_window_spend_counter_initialized( entity_type: str, entity_id: str, window_start: datetime, -): +) -> bool: is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) - if is_warm is False: - window_spend = await SpendCounterReseed.coalesced_window( - prisma_client=prisma_client, - spend_counter_cache=spend_counter_cache, - counter_key=counter_key, - entity_type=entity_type, - entity_id=entity_id, - window_start=window_start, + if is_warm is True: + return True + + window_spend = await SpendCounterReseed.coalesced_window( + prisma_client=prisma_client, + spend_counter_cache=spend_counter_cache, + counter_key=counter_key, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + if window_spend is None: + verbose_proxy_logger.warning( + "Skipping cold spend counter seed for %s because window spend could not be loaded", + counter_key, ) - if window_spend is None: - verbose_proxy_logger.warning( - "Skipping cold spend counter seed for %s because window spend could not be loaded", - counter_key, - ) + return False + return True async def _is_spend_counter_cache_warm(counter_key: str) -> bool: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 77c6f9ae63e..3631f94a60a 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -33,6 +33,10 @@ class _BudgetCounter: window_start: Optional[datetime] = None +class _CounterReservationUnavailable(Exception): + pass + + def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: if not budget_reservation: return set() @@ -99,10 +103,13 @@ async def reserve_budget_for_request( counter=counter, reserved_cost=reservation_cost, ) - reserved_value = await _reserve_counter( - counter=counter, - reservation_cost=reservation_cost, - ) + try: + reserved_value = await _reserve_counter( + counter=counter, + reservation_cost=reservation_cost, + ) + except _CounterReservationUnavailable: + continue applied_entries.append(entry) if reserved_value is not None: @@ -141,6 +148,9 @@ async def reserve_budget_for_request( ) raise + if not applied_entries: + return None + return { "reserved_cost": reservation_cost, "entries": applied_entries, @@ -572,12 +582,18 @@ async def _reserve_counter( source_cache_key=counter.source_cache_key, ) elif counter.spend_log_entity_id is not None and counter.window_start is not None: - await _ensure_window_spend_counter_initialized( + initialized = await _ensure_window_spend_counter_initialized( counter_key=counter.counter_key, entity_type=counter.entity_type, entity_id=counter.spend_log_entity_id, window_start=counter.window_start, ) + if initialized is False: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because window spend could not be loaded", + counter.counter_key, + ) + raise _CounterReservationUnavailable reserved_value = await _increment_spend_counter_cache( counter_key=counter.counter_key, @@ -707,6 +723,9 @@ async def _resize_applied_reservation( actual_cost=new_reserved_cost, default_reserved_cost=current_reserved_cost, ) + for entry in entries: + entry["reserved_cost"] = new_reserved_cost + entry["applied_adjustment"] = 0.0 def _counter_to_reservation_entry( diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index d1684b7ef1f..c4f86655932 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -757,6 +757,14 @@ async def test_should_shrink_uncapped_reservation_multiple_times( assert reservation is not None assert reservation["reserved_cost"] == pytest.approx(0.6) + assert [entry["reserved_cost"] for entry in reservation["entries"]] == [ + pytest.approx(0.6), + pytest.approx(0.6), + ] + assert [entry["applied_adjustment"] for entry in reservation["entries"]] == [ + pytest.approx(0.0), + pytest.approx(0.0), + ] assert counter_cache.in_memory_cache.get_cache( key="spend:key:key-budget-double-resize" ) == pytest.approx(0.9) @@ -842,6 +850,47 @@ async def test_should_skip_budget_window_with_unparseable_duration( ) == pytest.approx(0.9) +@pytest.mark.asyncio +async def test_should_skip_window_reservation_when_db_baseline_unavailable( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-window-db-unavailable", + budget_limits=[ + { + "budget_duration": "1h", + "max_budget": 1.0, + } + ], + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-window-db-unavailable:window:1h" + ) + is None + ) + + @pytest.mark.asyncio async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( spend_counter_state, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 53690907712..2acb5493775 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5261,8 +5261,15 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): redis_store[key] = (redis_store.get(key) or 0.0) + value return redis_store[key] + async def redis_set_cache(key, value, **_): + if key in redis_store: + return False + redis_store[key] = value + return True + fake_redis = AsyncMock() fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis @@ -5302,6 +5309,69 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): ps.prisma_client = orig_prisma +@pytest.mark.asyncio +async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + counter_key = "spend:key:key-window-concurrent-seed:window:1h" + window_start = datetime.now(timezone.utc) - timedelta(hours=1) + redis_store = {counter_key: 2.75} + redis_reads = 0 + + async def redis_get_cache(key): + nonlocal redis_reads + redis_reads += 1 + if redis_reads <= 2: + return None + return redis_store.get(key) + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache) + fake_redis.async_set_cache = AsyncMock(return_value=False) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + counter_cache.redis_cache = fake_redis + + fake_prisma = MagicMock() + fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( + return_value=[ + {"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}} + ] + ) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_window_spend_counter( + counter_key=counter_key, + entity_type="Key", + entity_id="key-window-concurrent-seed", + window_start=window_start, + increment=0.5, + ) + + fake_redis.async_set_cache.assert_awaited_once_with( + key=counter_key, + value=2.25, + nx=True, + ) + assert redis_store[counter_key] == pytest.approx(3.25) + assert counter_cache.in_memory_cache.get_cache( + key=counter_key + ) == pytest.approx(3.25) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + @pytest.mark.asyncio async def test_window_spend_counter_skips_invalid_window_start(): from litellm.caching.dual_cache import DualCache @@ -5346,13 +5416,14 @@ async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): ps.spend_counter_cache = counter_cache ps.prisma_client = None try: - await _ensure_window_spend_counter_initialized( + initialized = await _ensure_window_spend_counter_initialized( counter_key=counter_key, entity_type="Key", entity_id="key-window-db-unavailable", window_start=datetime.now(timezone.utc) - timedelta(hours=1), ) + assert initialized is False assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None finally: ps.spend_counter_cache = orig_counter From 66c0fe23da36666373918ebdb6f8dd342d0bdba5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:55:26 -0700 Subject: [PATCH 29/31] handle bad reservation counters after spend write --- litellm/proxy/proxy_server.py | 25 ++++++++-- tests/test_litellm/proxy/test_proxy_server.py | 48 +++++++++++++++++++ 2 files changed, 68 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e77561a5e59..1a9a53c9840 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1994,17 +1994,32 @@ async def _reconcile_budget_reservation_for_counter_update( from litellm.proxy.spend_tracking.budget_reservation import ( get_reserved_counter_keys, + invalidate_budget_reservation_counters, reconcile_budget_reservation, ) reserved_counter_keys = get_reserved_counter_keys( budget_reservation=budget_reservation ) - await reconcile_budget_reservation( - budget_reservation=budget_reservation, - actual_cost=response_cost or 0.0, - finalize=False, - ) + try: + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=response_cost or 0.0, + finalize=False, + ) + except Exception: + verbose_proxy_logger.warning( + "Failed to reconcile budget reservation after persisted spend; invalidating reserved counters and continuing", + exc_info=True, + ) + try: + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate reserved counters after reservation reconciliation failed" + ) return reserved_counter_keys diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 2acb5493775..669bc85856d 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5532,6 +5532,54 @@ async def test_increment_spend_counters_finalizes_none_cost_reservation(): ps.spend_counter_cache = orig_counter +@pytest.mark.asyncio +async def test_increment_spend_counters_invalidates_bad_reserved_counter_without_failing(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import increment_spend_counters + + counter_cache = DualCache() + budget_reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:key-bad-reserved-counter", + "entity_type": "Key", + "entity_id": "key-bad-reserved-counter", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + with patch( + "litellm.proxy.proxy_server.verbose_proxy_logger.warning" + ) as mock_warning: + await increment_spend_counters( + token="key-bad-reserved-counter", + team_id=None, + user_id=None, + response_cost=0.25, + budget_reservation=budget_reservation, + ) + + mock_warning.assert_called_once() + assert budget_reservation["finalized"] is True + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-bad-reserved-counter" + ) + is None + ) + finally: + ps.spend_counter_cache = orig_counter + + @pytest.mark.asyncio async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): from litellm.caching.dual_cache import DualCache From 403bbc3b88365111ef30bee5142f3737c4ca91f1 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 23:53:36 -0700 Subject: [PATCH 30/31] degrade budget reservation cache failures --- .../spend_tracking/budget_reservation.py | 65 ++++++++----- .../proxy/test_budget_reservation.py | 96 +++++++++++++++++++ 2 files changed, 139 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 3631f94a60a..47bd1a35a3e 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -573,33 +573,54 @@ async def _reserve_counter( from litellm.proxy.proxy_server import ( _ensure_spend_counter_initialized, _ensure_window_spend_counter_initialized, + _invalidate_spend_counter, _increment_spend_counter_cache, ) - if counter.source_cache_key is not None: - await _ensure_spend_counter_initialized( - counter_key=counter.counter_key, - source_cache_key=counter.source_cache_key, - ) - elif counter.spend_log_entity_id is not None and counter.window_start is not None: - initialized = await _ensure_window_spend_counter_initialized( - counter_key=counter.counter_key, - entity_type=counter.entity_type, - entity_id=counter.spend_log_entity_id, - window_start=counter.window_start, - ) - if initialized is False: - verbose_proxy_logger.warning( - "Skipping budget reservation for %s because window spend could not be loaded", - counter.counter_key, + try: + if counter.source_cache_key is not None: + await _ensure_spend_counter_initialized( + counter_key=counter.counter_key, + source_cache_key=counter.source_cache_key, ) - raise _CounterReservationUnavailable + elif ( + counter.spend_log_entity_id is not None and counter.window_start is not None + ): + initialized = await _ensure_window_spend_counter_initialized( + counter_key=counter.counter_key, + entity_type=counter.entity_type, + entity_id=counter.spend_log_entity_id, + window_start=counter.window_start, + ) + if initialized is False: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because window spend could not be loaded", + counter.counter_key, + ) + raise _CounterReservationUnavailable - reserved_value = await _increment_spend_counter_cache( - counter_key=counter.counter_key, - increment=reservation_cost, - ) - return float(reserved_value) if reserved_value is not None else None + reserved_value = await _increment_spend_counter_cache( + counter_key=counter.counter_key, + increment=reservation_cost, + ) + return float(reserved_value) if reserved_value is not None else None + except _CounterReservationUnavailable: + raise + except Exception: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because spend counter reservation failed", + counter.counter_key, + exc_info=True, + ) + try: + await _invalidate_spend_counter(counter_key=counter.counter_key) + except Exception: + verbose_proxy_logger.warning( + "Failed to invalidate spend counter after budget reservation failure for %s", + counter.counter_key, + exc_info=True, + ) + raise _CounterReservationUnavailable async def _get_current_counter_value(counter: _BudgetCounter) -> float: diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index c4f86655932..350eb952e2a 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -891,6 +891,102 @@ async def test_should_skip_window_reservation_when_db_baseline_unavailable( ) +@pytest.mark.asyncio +async def test_should_skip_reservation_when_counter_increment_fails( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-unavailable", + spend=0.0, + max_budget=1.0, + ) + + async def fail_increment_cache(*args, **kwargs): + raise RuntimeError("counter unavailable") + + monkeypatch.setattr(counter_cache, "async_increment_cache", fail_increment_cache) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.verbose_proxy_logger.warning" + ) as mock_warning, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert mock_warning.call_count >= 1 + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-unavailable" + ) + is None + ) + + +@pytest.mark.asyncio +async def test_should_skip_reservation_when_counter_initialization_fails( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-init-unavailable", + spend=0.0, + max_budget=1.0, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ), + patch( + "litellm.proxy.proxy_server._ensure_spend_counter_initialized", + side_effect=RuntimeError("redis unavailable"), + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.verbose_proxy_logger.warning" + ) as mock_warning, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert mock_warning.call_count >= 1 + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-init-unavailable" + ) + is None + ) + + @pytest.mark.asyncio async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( spend_counter_state, From 83ed317c50c80a9aa163484e0e3aca5e7a9610c8 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 1 May 2026 00:09:51 -0700 Subject: [PATCH 31/31] track reservation entry before counter write --- .../spend_tracking/budget_reservation.py | 28 +++++++-- .../proxy/test_budget_reservation.py | 60 +++++++++++++++++++ 2 files changed, 84 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 47bd1a35a3e..1d296611bfc 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -34,7 +34,14 @@ class _BudgetCounter: class _CounterReservationUnavailable(Exception): - pass + def __init__( + self, + touched_counter: bool = False, + counter_invalidated: bool = False, + ) -> None: + self.touched_counter = touched_counter + self.counter_invalidated = counter_invalidated + super().__init__("Counter reservation unavailable") def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: @@ -103,14 +110,20 @@ async def reserve_budget_for_request( counter=counter, reserved_cost=reservation_cost, ) + applied_entries.append(entry) try: reserved_value = await _reserve_counter( counter=counter, reservation_cost=reservation_cost, ) - except _CounterReservationUnavailable: + except _CounterReservationUnavailable as exc: + if exc.touched_counter and not exc.counter_invalidated: + await _release_applied_entries_best_effort( + entries=[entry], + default_reserved_cost=reservation_cost, + ) + applied_entries.remove(entry) continue - applied_entries.append(entry) if reserved_value is not None: current_spend = reserved_value @@ -577,6 +590,7 @@ async def _reserve_counter( _increment_spend_counter_cache, ) + attempted_increment = False try: if counter.source_cache_key is not None: await _ensure_spend_counter_initialized( @@ -599,6 +613,7 @@ async def _reserve_counter( ) raise _CounterReservationUnavailable + attempted_increment = True reserved_value = await _increment_spend_counter_cache( counter_key=counter.counter_key, increment=reservation_cost, @@ -612,15 +627,20 @@ async def _reserve_counter( counter.counter_key, exc_info=True, ) + counter_invalidated = False try: await _invalidate_spend_counter(counter_key=counter.counter_key) + counter_invalidated = True except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", counter.counter_key, exc_info=True, ) - raise _CounterReservationUnavailable + raise _CounterReservationUnavailable( + touched_counter=attempted_increment, + counter_invalidated=counter_invalidated, + ) async def _get_current_counter_value(counter: _BudgetCounter) -> float: diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 350eb952e2a..070b232066a 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -987,6 +987,66 @@ async def test_should_skip_reservation_when_counter_initialization_fails( ) +@pytest.mark.asyncio +async def test_should_release_tracked_entry_when_reservation_fails_after_increment( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-after-increment-failure", + spend=0.0, + max_budget=1.0, + ) + + import litellm.proxy.proxy_server as ps + + original_increment_counter = ps._increment_spend_counter_cache + first_increment = True + + async def fail_after_increment(counter_key: str, increment: float): + nonlocal first_increment + if first_increment: + first_increment = False + await counter_cache.async_increment_cache(key=counter_key, value=increment) + raise RuntimeError("lost increment response") + return await original_increment_counter( + counter_key=counter_key, + increment=increment, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ), + patch( + "litellm.proxy.proxy_server._increment_spend_counter_cache", + side_effect=fail_after_increment, + ), + patch( + "litellm.proxy.proxy_server._invalidate_spend_counter", + side_effect=RuntimeError("invalidate unavailable"), + ), + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-after-increment-failure" + ) == pytest.approx(0.0) + + @pytest.mark.asyncio async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( spend_counter_state,