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/56] 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/56] 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/56] 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/56] 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/56] 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/56] 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/56] 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/56] 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 15d4d514531e86f6c6e07d04823bf5cc0d2060cf Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:27:19 -0700 Subject: [PATCH 09/56] chore(callbacks): guard dynamic integration hosts --- litellm/integrations/langfuse/langfuse.py | 36 +++- .../integrations/langfuse/langfuse_handler.py | 1 + .../langfuse/langfuse_prompt_management.py | 19 ++- litellm/integrations/langsmith.py | 35 ++-- litellm/litellm_core_utils/litellm_logging.py | 7 +- .../vertex_ai_endpoints/langfuse_endpoints.py | 154 +++++++++++++++--- .../test_langfuse_dynamic_credentials.py | 46 ++++++ .../test_langsmith_dynamic_credentials.py | 50 ++++++ .../test_langfuse_passthrough_security.py | 102 ++++++++++++ 9 files changed, 399 insertions(+), 51 deletions(-) create mode 100644 tests/logging_callback_tests/test_langfuse_dynamic_credentials.py create mode 100644 tests/logging_callback_tests/test_langsmith_dynamic_credentials.py create mode 100644 tests/test_litellm/proxy/test_langfuse_passthrough_security.py diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index e691c490c85..aaff046a93b 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -90,6 +90,29 @@ def _extract_cache_read_input_tokens(usage_obj) -> int: return cache_read_input_tokens +def resolve_langfuse_credentials( + langfuse_public_key=None, + langfuse_secret=None, + langfuse_secret_key=None, + langfuse_host=None, + allow_env_credentials: bool = True, +): + if allow_env_credentials is False and langfuse_host is not None: + secret_key = langfuse_secret or langfuse_secret_key + public_key = langfuse_public_key + else: + secret_key = ( + langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY") + ) + public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") + + resolved_host = langfuse_host or os.getenv( + "LANGFUSE_HOST", "https://cloud.langfuse.com" + ) + + return public_key, secret_key, resolved_host + + class LangFuseLogger: # Class variables or attributes def __init__( @@ -98,6 +121,7 @@ class LangFuseLogger: langfuse_secret=None, langfuse_host=None, flush_interval=1, + allow_env_credentials: bool = True, ): try: import langfuse @@ -106,11 +130,13 @@ class LangFuseLogger: raise Exception( f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m" ) - # Instance variables - self.secret_key = langfuse_secret or os.getenv("LANGFUSE_SECRET_KEY") - self.public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") - self.langfuse_host = langfuse_host or os.getenv( - "LANGFUSE_HOST", "https://cloud.langfuse.com" + self.public_key, self.secret_key, self.langfuse_host = ( + resolve_langfuse_credentials( + langfuse_public_key=langfuse_public_key, + langfuse_secret=langfuse_secret, + langfuse_host=langfuse_host, + allow_env_credentials=allow_env_credentials, + ) ) if not ( self.langfuse_host.startswith("http://") diff --git a/litellm/integrations/langfuse/langfuse_handler.py b/litellm/integrations/langfuse/langfuse_handler.py index fbadf1a2fc7..3552054bcd4 100644 --- a/litellm/integrations/langfuse/langfuse_handler.py +++ b/litellm/integrations/langfuse/langfuse_handler.py @@ -117,6 +117,7 @@ class LangFuseHandler: langfuse_public_key=credentials.get("langfuse_public_key"), langfuse_secret=credentials.get("langfuse_secret"), langfuse_host=credentials.get("langfuse_host"), + allow_env_credentials=credentials.get("langfuse_host") is None, ) in_memory_dynamic_logger_cache.set_cache( credentials=credentials, diff --git a/litellm/integrations/langfuse/langfuse_prompt_management.py b/litellm/integrations/langfuse/langfuse_prompt_management.py index 5f4ced3a5cb..b7a565512c6 100644 --- a/litellm/integrations/langfuse/langfuse_prompt_management.py +++ b/litellm/integrations/langfuse/langfuse_prompt_management.py @@ -20,7 +20,7 @@ from ...litellm_core_utils.specialty_caches.dynamic_logging_cache import ( DynamicLoggingCache, ) from ..prompt_management_base import PromptManagementBase -from .langfuse import LangFuseLogger +from .langfuse import LangFuseLogger, resolve_langfuse_credentials from .langfuse_handler import LangFuseHandler if TYPE_CHECKING: @@ -46,6 +46,7 @@ def langfuse_client_init( langfuse_secret_key=None, langfuse_host=None, flush_interval=1, + allow_env_credentials: bool = True, ) -> LangfuseClass: """ Initialize Langfuse client with caching to prevent multiple initializations. @@ -70,14 +71,12 @@ def langfuse_client_init( f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n\033[0m" ) - # Instance variables - - secret_key = ( - langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY") - ) - public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") - langfuse_host = langfuse_host or os.getenv( - "LANGFUSE_HOST", "https://cloud.langfuse.com" + public_key, secret_key, langfuse_host = resolve_langfuse_credentials( + langfuse_public_key=langfuse_public_key, + langfuse_secret=langfuse_secret, + langfuse_secret_key=langfuse_secret_key, + langfuse_host=langfuse_host, + allow_env_credentials=allow_env_credentials, ) if not ( @@ -222,6 +221,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge langfuse_secret=dynamic_callback_params.get("langfuse_secret"), langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"), langfuse_host=dynamic_callback_params.get("langfuse_host"), + allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None, ) langfuse_prompt_client = self._get_prompt_from_id( langfuse_prompt_id=prompt_id, @@ -246,6 +246,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge langfuse_secret=dynamic_callback_params.get("langfuse_secret"), langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"), langfuse_host=dynamic_callback_params.get("langfuse_host"), + allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None, ) langfuse_prompt_client = self._get_prompt_from_id( langfuse_prompt_id=prompt_id, diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py index 3d4fd39ebe1..3a206122373 100644 --- a/litellm/integrations/langsmith.py +++ b/litellm/integrations/langsmith.py @@ -112,17 +112,28 @@ class LangsmithLogger(CustomBatchLogger): langsmith_project: Optional[str] = None, langsmith_base_url: Optional[str] = None, langsmith_tenant_id: Optional[str] = None, + allow_env_credentials: bool = True, ) -> LangsmithCredentialsObject: - _credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY") - _credentials_project = ( - langsmith_project or os.getenv("LANGSMITH_PROJECT") or "litellm-completion" - ) - _credentials_base_url = ( - langsmith_base_url - or os.getenv("LANGSMITH_BASE_URL") - or "https://api.smith.langchain.com" - ) - _credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID") + if allow_env_credentials is False and langsmith_base_url is not None: + _credentials_api_key = langsmith_api_key + _credentials_project = langsmith_project or "litellm-completion" + _credentials_base_url = langsmith_base_url + _credentials_tenant_id = langsmith_tenant_id + else: + _credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY") + _credentials_project = ( + langsmith_project + or os.getenv("LANGSMITH_PROJECT") + or "litellm-completion" + ) + _credentials_base_url = ( + langsmith_base_url + or os.getenv("LANGSMITH_BASE_URL") + or "https://api.smith.langchain.com" + ) + _credentials_tenant_id = langsmith_tenant_id or os.getenv( + "LANGSMITH_TENANT_ID" + ) return LangsmithCredentialsObject( LANGSMITH_API_KEY=_credentials_api_key, @@ -540,6 +551,10 @@ class LangsmithLogger(CustomBatchLogger): langsmith_tenant_id=standard_callback_dynamic_params.get( "langsmith_tenant_id", None ), + allow_env_credentials=standard_callback_dynamic_params.get( + "langsmith_base_url", None + ) + is None, ) else: credentials = self.default_credentials diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 829c1c9ca07..e1240b436c4 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3242,10 +3242,15 @@ class Logging(LiteLLMLoggingBaseClass): ), langfuse_secret=self.standard_callback_dynamic_params.get( "langfuse_secret" - ), + ) + or self.standard_callback_dynamic_params.get("langfuse_secret_key"), langfuse_host=self.standard_callback_dynamic_params.get( "langfuse_host" ), + allow_env_credentials=self.standard_callback_dynamic_params.get( + "langfuse_host" + ) + is None, ) return langFuseLogger diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index b6454bf077b..8ce1bedcf90 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -12,11 +12,13 @@ import base64 import os from base64 import b64encode from typing import Optional +from urllib.parse import unquote import httpx -from fastapi import APIRouter, Request, Response +from fastapi import APIRouter, HTTPException, Request, Response, status import litellm +from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers @@ -27,6 +29,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( router = APIRouter() default_vertex_config = None +_DEFAULT_LANGFUSE_HOST = "https://cloud.langfuse.com" def create_request_copy(request: Request): @@ -39,6 +42,116 @@ def create_request_copy(request: Request): } +def _decode_to_convergence(value: str) -> str: + previous = value + while True: + decoded = unquote(previous) + if decoded == previous: + return decoded + previous = decoded + + +def _normalize_langfuse_base_url(base_target_url: str) -> str: + if not ( + base_target_url.startswith("http://") or base_target_url.startswith("https://") + ): + # Existing behavior allows host-only Langfuse settings. + base_target_url = "http://" + base_target_url + + try: + base_url = httpx.URL(base_target_url) + except Exception as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"Invalid Langfuse host: {str(e)}"}, + ) + + if base_url.scheme not in ("http", "https") or not base_url.host: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse host"}, + ) + + if base_url.userinfo: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Langfuse host must not include credentials"}, + ) + + return str(base_url) + + +def _validate_langfuse_proxy_path(endpoint: str) -> str: + decoded_endpoint = _decode_to_convergence(endpoint) + if any(ord(char) < 32 for char in decoded_endpoint): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse endpoint path"}, + ) + if "\\" in decoded_endpoint or decoded_endpoint.startswith("//"): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse endpoint path"}, + ) + + endpoint_path = "/" + decoded_endpoint.lstrip("/") + if any(segment in (".", "..") for segment in endpoint_path.split("/")): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse endpoint path"}, + ) + return endpoint_path + + +def _get_langfuse_proxy_credentials( + *, + dynamic_host_supplied: bool, + dynamic_langfuse_public_key: Optional[str], + dynamic_langfuse_secret_key: Optional[str], +): + if dynamic_host_supplied: + if not dynamic_langfuse_public_key or not dynamic_langfuse_secret_key: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "Dynamic Langfuse hosts must include dynamic Langfuse credentials" + }, + ) + return dynamic_langfuse_public_key, dynamic_langfuse_secret_key + + return ( + dynamic_langfuse_public_key + or litellm.utils.get_secret(secret_name="LANGFUSE_PUBLIC_KEY"), + dynamic_langfuse_secret_key + or litellm.utils.get_secret(secret_name="LANGFUSE_SECRET_KEY"), + ) + + +def _build_langfuse_proxy_target( + *, + endpoint: str, + base_target_url: str, + dynamic_host_supplied: bool, +): + endpoint_path = _validate_langfuse_proxy_path(endpoint) + base_url = httpx.URL(_normalize_langfuse_base_url(base_target_url)) + updated_url = base_url.copy_with(path=endpoint_path) + custom_headers = {} + + if dynamic_host_supplied and getattr(litellm, "user_url_validation", True): + try: + target_url, host_header = validate_url(str(updated_url)) + except SSRFError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"Invalid Langfuse host: {str(e)}"}, + ) + custom_headers["Host"] = host_header + return target_url, custom_headers + + return str(updated_url), custom_headers + + @router.api_route( "/langfuse/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -91,44 +204,33 @@ async def langfuse_proxy_route( elif k == "langfuse_host": dynamic_langfuse_host = v + dynamic_host_supplied = dynamic_langfuse_host is not None base_target_url: str = ( dynamic_langfuse_host - or os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com") - or "https://cloud.langfuse.com" + or os.getenv("LANGFUSE_HOST", _DEFAULT_LANGFUSE_HOST) + or _DEFAULT_LANGFUSE_HOST ) - if not ( - base_target_url.startswith("http://") or base_target_url.startswith("https://") - ): - # add http:// if unset, assume communicating over private network - e.g. render - base_target_url = "http://" + base_target_url - - encoded_endpoint = httpx.URL(endpoint).path - - # Ensure endpoint starts with '/' for proper URL construction - if not encoded_endpoint.startswith("/"): - encoded_endpoint = "/" + encoded_endpoint - - # Construct the full target URL using httpx - base_url = httpx.URL(base_target_url) - updated_url = base_url.copy_with(path=encoded_endpoint) - - # Add or update query parameters - langfuse_public_key = dynamic_langfuse_public_key or litellm.utils.get_secret( - secret_name="LANGFUSE_PUBLIC_KEY" + langfuse_public_key, langfuse_secret_key = _get_langfuse_proxy_credentials( + dynamic_host_supplied=dynamic_host_supplied, + dynamic_langfuse_public_key=dynamic_langfuse_public_key, + dynamic_langfuse_secret_key=dynamic_langfuse_secret_key, ) - langfuse_secret_key = dynamic_langfuse_secret_key or litellm.utils.get_secret( - secret_name="LANGFUSE_SECRET_KEY" + target_url, target_headers = _build_langfuse_proxy_target( + endpoint=endpoint, + base_target_url=base_target_url, + dynamic_host_supplied=dynamic_host_supplied, ) langfuse_combined_key = "Basic " + b64encode( f"{langfuse_public_key}:{langfuse_secret_key}".encode("utf-8") ).decode("ascii") + target_headers["Authorization"] = langfuse_combined_key ## CREATE PASS-THROUGH endpoint_func = create_pass_through_route( endpoint=endpoint, - target=str(updated_url), - custom_headers={"Authorization": langfuse_combined_key}, + target=target_url, + custom_headers=target_headers, query_params=dict(request.query_params), # type: ignore ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( diff --git a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py new file mode 100644 index 00000000000..ac4486639da --- /dev/null +++ b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py @@ -0,0 +1,46 @@ +from litellm.integrations.langfuse.langfuse import resolve_langfuse_credentials + + +def test_resolve_langfuse_credentials_does_not_use_env_for_dynamic_host(monkeypatch): + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key, host = resolve_langfuse_credentials( + langfuse_host="https://attacker.example", + allow_env_credentials=False, + ) + + assert public_key is None + assert secret_key is None + assert host == "https://attacker.example" + + +def test_resolve_langfuse_credentials_accepts_secret_key_alias_for_dynamic_host( + monkeypatch, +): + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key, host = resolve_langfuse_credentials( + langfuse_public_key="dynamic-public", + langfuse_secret_key="dynamic-secret", + langfuse_host="https://team-langfuse.example", + allow_env_credentials=False, + ) + + assert public_key == "dynamic-public" + assert secret_key == "dynamic-secret" + assert host == "https://team-langfuse.example" + + +def test_resolve_langfuse_credentials_keeps_env_for_global_config(monkeypatch): + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key, host = resolve_langfuse_credentials( + langfuse_host="https://admin-configured.example", + allow_env_credentials=True, + ) + + assert public_key == "global-public" + assert secret_key == "global-secret" + assert host == "https://admin-configured.example" diff --git a/tests/logging_callback_tests/test_langsmith_dynamic_credentials.py b/tests/logging_callback_tests/test_langsmith_dynamic_credentials.py new file mode 100644 index 00000000000..f1912c58464 --- /dev/null +++ b/tests/logging_callback_tests/test_langsmith_dynamic_credentials.py @@ -0,0 +1,50 @@ +import pytest + +from litellm.integrations.langsmith import LangsmithLogger + + +@pytest.mark.asyncio +async def test_get_credentials_from_env_does_not_use_env_for_dynamic_base_url( + monkeypatch, +): + monkeypatch.setenv("LANGSMITH_API_KEY", "global-key") + monkeypatch.setenv("LANGSMITH_PROJECT", "global-project") + monkeypatch.setenv("LANGSMITH_TENANT_ID", "global-tenant") + logger = LangsmithLogger( + langsmith_api_key="default-key", + langsmith_project="default-project", + langsmith_base_url="https://default.example", + ) + + credentials = logger.get_credentials_from_env( + langsmith_base_url="https://attacker.example", + allow_env_credentials=False, + ) + + assert credentials["LANGSMITH_API_KEY"] is None + assert credentials["LANGSMITH_PROJECT"] == "litellm-completion" + assert credentials["LANGSMITH_BASE_URL"] == "https://attacker.example" + assert credentials["LANGSMITH_TENANT_ID"] is None + + +@pytest.mark.asyncio +async def test_dynamic_langsmith_base_url_does_not_inherit_default_api_key( + monkeypatch, +): + monkeypatch.setenv("LANGSMITH_API_KEY", "global-key") + logger = LangsmithLogger( + langsmith_api_key="default-key", + langsmith_project="default-project", + langsmith_base_url="https://default.example", + ) + + credentials = logger._get_credentials_to_use_for_request( + kwargs={ + "standard_callback_dynamic_params": { + "langsmith_base_url": "https://attacker.example" + } + } + ) + + assert credentials["LANGSMITH_API_KEY"] is None + assert credentials["LANGSMITH_BASE_URL"] == "https://attacker.example" diff --git a/tests/test_litellm/proxy/test_langfuse_passthrough_security.py b/tests/test_litellm/proxy/test_langfuse_passthrough_security.py new file mode 100644 index 00000000000..5ef3c38c09d --- /dev/null +++ b/tests/test_litellm/proxy/test_langfuse_passthrough_security.py @@ -0,0 +1,102 @@ +import socket + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.proxy.vertex_ai_endpoints.langfuse_endpoints import ( + _build_langfuse_proxy_target, + _get_langfuse_proxy_credentials, +) + + +def test_dynamic_langfuse_host_requires_dynamic_credentials(monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True, raising=False) + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + with pytest.raises(HTTPException) as exc: + _get_langfuse_proxy_credentials( + dynamic_host_supplied=True, + dynamic_langfuse_public_key=None, + dynamic_langfuse_secret_key=None, + ) + + assert exc.value.status_code == 400 + + +def test_global_langfuse_host_can_use_env_credentials(monkeypatch): + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key = _get_langfuse_proxy_credentials( + dynamic_host_supplied=False, + dynamic_langfuse_public_key=None, + dynamic_langfuse_secret_key=None, + ) + + assert public_key == "global-public" + assert secret_key == "global-secret" + + +@pytest.mark.parametrize( + "endpoint", + [ + "../api/public/projects", + "%2e%2e/api/public/projects", + "%252e%252e%252fapi/public/projects", + "api\\public\\projects", + "%2f%2fattacker.example/api", + ], +) +def test_langfuse_proxy_target_rejects_traversal_paths(endpoint): + with pytest.raises(HTTPException) as exc: + _build_langfuse_proxy_target( + endpoint=endpoint, + base_target_url="https://cloud.langfuse.com", + dynamic_host_supplied=False, + ) + + assert exc.value.status_code == 400 + + +def test_dynamic_langfuse_proxy_target_rejects_internal_host(monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True, raising=False) + + with pytest.raises(HTTPException) as exc: + _build_langfuse_proxy_target( + endpoint="api/public/projects", + base_target_url="http://127.0.0.1:3000", + dynamic_host_supplied=True, + ) + + assert exc.value.status_code == 400 + + +def test_dynamic_langfuse_proxy_target_preserves_host_header_for_http(monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True, raising=False) + + def fake_getaddrinfo(host, port, proto): + assert host == "langfuse.example" + assert port == 80 + assert proto == socket.IPPROTO_TCP + return [ + ( + socket.AF_INET, + socket.SOCK_STREAM, + socket.IPPROTO_TCP, + "", + ("8.8.8.8", 80), + ) + ] + + monkeypatch.setattr(socket, "getaddrinfo", fake_getaddrinfo) + + target_url, headers = _build_langfuse_proxy_target( + endpoint="api/public/projects", + base_target_url="http://langfuse.example", + dynamic_host_supplied=True, + ) + + assert target_url == "http://8.8.8.8/api/public/projects" + assert headers["Host"] == "langfuse.example" From d19f342af783b4e104013fa92249753669dedb47 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:32:41 -0700 Subject: [PATCH 10/56] chore(callbacks): satisfy langfuse type check --- litellm/integrations/langfuse/langfuse.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index aaff046a93b..0efc7d66876 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -186,9 +186,10 @@ class LangFuseLogger: project_id = None if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None: + upstream_langfuse_debug_env = os.getenv("UPSTREAM_LANGFUSE_DEBUG") upstream_langfuse_debug = ( - str_to_bool(self.upstream_langfuse_debug) - if self.upstream_langfuse_debug is not None + str_to_bool(upstream_langfuse_debug_env) + if upstream_langfuse_debug_env is not None else None ) self.upstream_langfuse_secret_key = os.getenv( @@ -199,7 +200,7 @@ class LangFuseLogger: ) self.upstream_langfuse_host = os.getenv("UPSTREAM_LANGFUSE_HOST") self.upstream_langfuse_release = os.getenv("UPSTREAM_LANGFUSE_RELEASE") - self.upstream_langfuse_debug = os.getenv("UPSTREAM_LANGFUSE_DEBUG") + self.upstream_langfuse_debug = upstream_langfuse_debug_env self.upstream_langfuse = Langfuse( public_key=self.upstream_langfuse_public_key, secret_key=self.upstream_langfuse_secret_key, From 258edac7276aeae50903c9a4ce3e687e591f6111 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:34:39 -0700 Subject: [PATCH 11/56] test(callbacks): cover upstream langfuse debug env --- .../test_langfuse_dynamic_credentials.py | 37 +++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py index ac4486639da..14478649ad3 100644 --- a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py +++ b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py @@ -1,3 +1,7 @@ +import sys +from types import ModuleType, SimpleNamespace + +import litellm from litellm.integrations.langfuse.langfuse import resolve_langfuse_credentials @@ -44,3 +48,36 @@ def test_resolve_langfuse_credentials_keeps_env_for_global_config(monkeypatch): assert public_key == "global-public" assert secret_key == "global-secret" assert host == "https://admin-configured.example" + + +def test_upstream_langfuse_debug_env_is_passed(monkeypatch): + from litellm.integrations.langfuse.langfuse import LangFuseLogger + + class FakeLangfuse: + instances = [] + + def __init__(self, **kwargs): + self.kwargs = kwargs + FakeLangfuse.instances.append(self) + + fake_langfuse_module = ModuleType("langfuse") + fake_langfuse_module.Langfuse = FakeLangfuse + fake_langfuse_module.version = SimpleNamespace(__version__="2.6.0") + + monkeypatch.setitem(sys.modules, "langfuse", fake_langfuse_module) + monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "upstream-secret") + monkeypatch.setenv("UPSTREAM_LANGFUSE_PUBLIC_KEY", "upstream-public") + monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example") + monkeypatch.setenv("UPSTREAM_LANGFUSE_RELEASE", "release") + monkeypatch.setenv("UPSTREAM_LANGFUSE_DEBUG", "true") + + logger = LangFuseLogger( + langfuse_public_key="public", + langfuse_secret="secret", + langfuse_host="https://langfuse.example", + ) + + assert logger.upstream_langfuse_debug == "true" + assert FakeLangfuse.instances[-1].kwargs["debug"] is True From bb6d7c9715dc44518b606d51411c50ef3538b526 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:36:51 -0700 Subject: [PATCH 12/56] fix(callbacks): preserve langfuse secret alias --- .../integrations/langfuse/langfuse_handler.py | 3 +- .../test_langfuse_dynamic_credentials.py | 46 +++++++++++++++++++ 2 files changed, 48 insertions(+), 1 deletion(-) diff --git a/litellm/integrations/langfuse/langfuse_handler.py b/litellm/integrations/langfuse/langfuse_handler.py index 3552054bcd4..4a809726424 100644 --- a/litellm/integrations/langfuse/langfuse_handler.py +++ b/litellm/integrations/langfuse/langfuse_handler.py @@ -115,7 +115,8 @@ class LangFuseHandler: langfuse_logger = LangFuseLogger( langfuse_public_key=credentials.get("langfuse_public_key"), - langfuse_secret=credentials.get("langfuse_secret"), + langfuse_secret=credentials.get("langfuse_secret") + or credentials.get("langfuse_secret_key"), langfuse_host=credentials.get("langfuse_host"), allow_env_credentials=credentials.get("langfuse_host") is None, ) diff --git a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py index 14478649ad3..1b198623381 100644 --- a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py +++ b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py @@ -3,6 +3,7 @@ from types import ModuleType, SimpleNamespace import litellm from litellm.integrations.langfuse.langfuse import resolve_langfuse_credentials +from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler def test_resolve_langfuse_credentials_does_not_use_env_for_dynamic_host(monkeypatch): @@ -81,3 +82,48 @@ def test_upstream_langfuse_debug_env_is_passed(monkeypatch): assert logger.upstream_langfuse_debug == "true" assert FakeLangfuse.instances[-1].kwargs["debug"] is True + + +def test_langfuse_handler_accepts_secret_key_alias(monkeypatch): + captured = {} + + class FakeLangFuseLogger: + def __init__( + self, + *, + langfuse_public_key=None, + langfuse_secret=None, + langfuse_host=None, + allow_env_credentials=True, + ): + captured["langfuse_public_key"] = langfuse_public_key + captured["langfuse_secret"] = langfuse_secret + captured["langfuse_host"] = langfuse_host + captured["allow_env_credentials"] = allow_env_credentials + + class FakeDynamicLoggingCache: + def set_cache(self, *, credentials, service_name, logging_obj): + captured["cached_credentials"] = credentials + captured["cached_service_name"] = service_name + captured["cached_logging_obj"] = logging_obj + + monkeypatch.setattr( + "litellm.integrations.langfuse.langfuse_handler.LangFuseLogger", + FakeLangFuseLogger, + ) + + logger = LangFuseHandler._create_langfuse_logger_from_credentials( + credentials={ + "langfuse_public_key": "dynamic-public", + "langfuse_secret_key": "dynamic-secret", + "langfuse_host": "https://langfuse.example", + }, + in_memory_dynamic_logger_cache=FakeDynamicLoggingCache(), + ) + + assert captured["langfuse_public_key"] == "dynamic-public" + assert captured["langfuse_secret"] == "dynamic-secret" + assert captured["langfuse_host"] == "https://langfuse.example" + assert captured["allow_env_credentials"] is False + assert captured["cached_service_name"] == "langfuse" + assert captured["cached_logging_obj"] is logger 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 13/56] 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 ad3a251eb8fbc768179741d18db63c75395fe796 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:39:05 -0700 Subject: [PATCH 14/56] chore(proxy): refresh lazy openapi snapshot --- litellm/proxy/_lazy_openapi_snapshot.json | 34 +++++++++++------------ 1 file changed, 17 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8331f748c6e..3a746b9a27e 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_head", "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_head", "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_head", "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_head", "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_head", "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_head", "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_head", "parameters": [ { "in": "path", From 3800596d08d6b455addee969708b550557613330 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 14:42:50 -0700 Subject: [PATCH 15/56] fix(proxy): stabilize lazy openapi snapshot ids --- litellm/proxy/_lazy_openapi_snapshot.json | 14 ++++++------ litellm/proxy/_lazy_openapi_snapshot.py | 22 ++++++++++++++++++- .../proxy/test_lazy_openapi_snapshot.py | 16 ++++++++++++++ 3 files changed, 44 insertions(+), 8 deletions(-) create mode 100644 tests/test_litellm/proxy/test_lazy_openapi_snapshot.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3a746b9a27e..b8e9eb6c261 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -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_head", + "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_head", + "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_head", + "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_head", + "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_head", + "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_head", + "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_head", + "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..fbd1a49aacf 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -8,9 +8,10 @@ any drift as a neutral check. """ import json +import re import sys from pathlib import Path -from typing import Dict, Optional +from typing import Dict, Iterable, Optional SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" @@ -25,6 +26,24 @@ def load_snapshot() -> Optional[Dict[str, Dict]]: return None +def _stable_generate_unique_id(route) -> str: + operation_id = f"{route.name}{route.path_format}" + operation_id = re.sub(r"\W", "_", operation_id) + methods = sorted(route.methods or []) + if not methods: + return operation_id + return f"{operation_id}_{methods[0].lower()}" + + +def _set_stable_operation_ids(routes: Iterable) -> None: + for route in routes: + if getattr(route, "operation_id", None) is not None: + continue + if getattr(route, "methods", None) is None: + continue + route.operation_id = _stable_generate_unique_id(route) + + def generate_snapshot() -> Dict[str, Dict]: import importlib @@ -51,6 +70,7 @@ def generate_snapshot() -> Dict[str, Dict]: ] if not feat_routes: continue + _set_stable_operation_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(): diff --git a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py new file mode 100644 index 00000000000..33f8ded84c9 --- /dev/null +++ b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py @@ -0,0 +1,16 @@ +from types import SimpleNamespace + +from litellm.proxy._lazy_openapi_snapshot import _stable_generate_unique_id + + +def test_stable_generate_unique_id_sorts_route_methods(): + route = SimpleNamespace( + name="langfuse_proxy_route", + path_format="/langfuse/{endpoint}", + methods={"POST", "GET", "DELETE", "PATCH", "PUT"}, + ) + + assert ( + _stable_generate_unique_id(route) + == "langfuse_proxy_route_langfuse__endpoint__delete" + ) 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 16/56] 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 17/56] 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 d3ab59e0595ed1b10c3de6da24a80a52907d31de Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 15:04:25 -0700 Subject: [PATCH 18/56] chore(vector stores): tighten managed store access --- litellm/proxy/common_request_processing.py | 54 +++ .../llm_passthrough_endpoints.py | 18 + litellm/proxy/rag_endpoints/endpoints.py | 49 +++ .../proxy/vector_store_endpoints/endpoints.py | 63 ++- litellm/proxy/vector_store_endpoints/utils.py | 96 +++++ .../vector_store_files_endpoints/endpoints.py | 29 +- litellm/vector_stores/main.py | 8 +- .../test_vector_store_tenant_guard.py | 362 ++++++++++++++++++ 8 files changed, 638 insertions(+), 41 deletions(-) create mode 100644 tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 76c52f83ee4..e3ad714d45d 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -97,6 +97,55 @@ def _serialize_http_exception_detail( return str(detail), None +def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[str]: + vector_store_ids: set[str] = set() + tools = data.get("tools") + if not isinstance(tools, list): + return vector_store_ids + + for tool in tools: + if not isinstance(tool, dict) or tool.get("type") != "file_search": + continue + ids = tool.get("vector_store_ids") or [] + if not isinstance(ids, list): + raise HTTPException( + status_code=400, + detail={ + "error": "file_search.vector_store_ids must be a list of strings" + }, + ) + for vector_store_id in ids: + if not isinstance(vector_store_id, str) or not vector_store_id: + raise HTTPException( + status_code=400, + detail={ + "error": "file_search.vector_store_ids must be a list of strings" + }, + ) + vector_store_ids.add(vector_store_id) + + return vector_store_ids + + +async def _authorize_response_file_search_vector_stores( + data: Dict[str, Any], + user_api_key_dict: UserAPIKeyAuth, +) -> None: + vector_store_ids = _collect_response_file_search_vector_store_ids(data) + if not vector_store_ids: + return + + from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store_id, + ) + + for vector_store_id in sorted(vector_store_ids): + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) + + async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]: """Parses an event line and returns an error code if present, else None.""" event_line = ( @@ -786,6 +835,11 @@ class ProxyBaseLLMRequestProcessing: version=version, proxy_config=proxy_config, ) + if route_type in {"aresponses", "_aresponses_websocket"}: + await _authorize_response_file_search_vector_stores( + data=self.data, + user_api_key_dict=user_api_key_dict, + ) # Calculate request queue time after add_litellm_data_to_request # which sets arrival_time in proxy_server_request diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 6521abffb85..ddb6717cb2e 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -47,6 +47,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( ) from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store, is_allowed_to_call_vector_store_endpoint, ) from litellm.secret_managers.main import get_secret_str @@ -533,6 +534,10 @@ async def milvus_proxy_route( ) if vector_store is None: raise Exception(f"Vector store not found for {vector_store_name}") + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + ) litellm_params = vector_store.get("litellm_params") or {} auth_credentials = provider_config.get_auth_credentials( litellm_params=litellm_params @@ -1438,6 +1443,10 @@ async def azure_proxy_route( ) if vector_store is None: raise Exception(f"Vector store not found for {vector_store_name}") + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + ) litellm_params = vector_store.get("litellm_params") or {} auth_credentials = provider_config.get_auth_credentials( litellm_params=litellm_params @@ -1777,6 +1786,11 @@ async def _base_vertex_proxy_route( request=request, api_key=api_key_to_use, ) + if router_credentials is not None: + await assert_user_can_access_vector_store( + vector_store=router_credentials, + user_api_key_dict=user_api_key_dict, + ) vertex_project: Optional[str] = get_vertex_project_id_from_url(endpoint) vertex_location: Optional[str] = get_vertex_location_from_url(endpoint) @@ -1929,6 +1943,10 @@ async def vertex_discovery_proxy_route( "Vector store ID %s found in endpoint but no credentials found in registry", vector_store_id, ) + raise HTTPException( + status_code=403, + detail="Access denied: You do not have permission to access this vector store", + ) discovery_handler = get_vertex_pass_through_handler(call_type="discovery") return await _base_vertex_proxy_route( diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 95ca51612fc..9e6093a47a1 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -22,10 +22,45 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, get_form_data, ) +from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store_id, +) router = APIRouter() +def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: + vector_store_ids: set[str] = set() + + if isinstance(payload, dict): + for key, value in payload.items(): + if key == "vector_store_id": + if not isinstance(value, str) or not value: + raise HTTPException( + status_code=400, + detail={"error": "vector_store_id must be a non-empty string"}, + ) + vector_store_ids.add(value) + continue + vector_store_ids.update(_collect_vector_store_ids_from_payload(value)) + elif isinstance(payload, list): + for item in payload: + vector_store_ids.update(_collect_vector_store_ids_from_payload(item)) + + return vector_store_ids + + +async def _authorize_nested_vector_store_ids( + payload: Any, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + for vector_store_id in sorted(_collect_vector_store_ids_from_payload(payload)): + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) + + def _build_file_metadata_entry( response: Any, file_data: Optional[Tuple[str, bytes, str]] = None, @@ -385,6 +420,11 @@ async def rag_ingest( }, ) + await _authorize_nested_vector_store_ids( + payload=ingest_options, + user_api_key_dict=user_api_key_dict, + ) + # Add litellm data request_data: Dict[str, Any] = {} request_data = await add_litellm_data_to_request( @@ -537,11 +577,20 @@ async def rag_query( status_code=400, detail={"error": "retrieval_config is required"}, ) + if not isinstance(retrieval_config, dict): + raise HTTPException( + status_code=400, + detail={"error": "retrieval_config must be an object"}, + ) if "vector_store_id" not in retrieval_config: raise HTTPException( status_code=400, detail={"error": "retrieval_config must contain 'vector_store_id'"}, ) + await _authorize_nested_vector_store_ids( + payload=retrieval_config, + user_api_key_dict=user_api_key_dict, + ) # Add litellm data request_data: Dict[str, Any] = {} diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 1fdfad8c96c..05423d9843d 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -1,8 +1,6 @@ from typing import Any, Dict, Optional from fastapi import APIRouter, Depends, HTTPException, Request, Response - -import litellm from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, ) @@ -10,7 +8,10 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.utils import jsonify_object -from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store +from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store, + get_litellm_managed_vector_store, +) from litellm.types.vector_stores import IndexCreateRequest router = APIRouter() @@ -32,9 +33,14 @@ async def _check_vector_store_access( - key-level and team-level ``object_permission.vector_stores`` allowlists - team_id match between key and store """ - return await can_user_access_vector_store( - vector_store=vector_store, user_api_key_dict=user_api_key_dict - ) + try: + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + ) + return True + except HTTPException: + return False async def _update_request_data_with_litellm_managed_vector_store_registry( @@ -53,35 +59,27 @@ async def _update_request_data_with_litellm_managed_vector_store_registry( Raises: HTTPException: If user doesn't have access to the vector store """ - if litellm.vector_store_registry is not None: - vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = ( - litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( - vector_store_id=vector_store_id + vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = ( + await get_litellm_managed_vector_store(vector_store_id=vector_store_id) + ) + if vector_store_to_run is not None: + if user_api_key_dict is not None: + await assert_user_can_access_vector_store( + vector_store=vector_store_to_run, + user_api_key_dict=user_api_key_dict, ) - ) - if vector_store_to_run is not None: - if user_api_key_dict is not None: - if not await _check_vector_store_access( - vector_store_to_run, user_api_key_dict - ): - raise HTTPException( - status_code=403, - detail="Access denied: You do not have permission to access this vector store", - ) - if "custom_llm_provider" in vector_store_to_run: - data["custom_llm_provider"] = vector_store_to_run.get( - "custom_llm_provider" - ) + if "custom_llm_provider" in vector_store_to_run: + data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider") - if "litellm_credential_name" in vector_store_to_run: - data["litellm_credential_name"] = vector_store_to_run.get( - "litellm_credential_name" - ) + if "litellm_credential_name" in vector_store_to_run: + data["litellm_credential_name"] = vector_store_to_run.get( + "litellm_credential_name" + ) - if "litellm_params" in vector_store_to_run: - litellm_params = vector_store_to_run.get("litellm_params", {}) or {} - data.update(litellm_params) + if "litellm_params" in vector_store_to_run: + litellm_params = vector_store_to_run.get("litellm_params", {}) or {} + data.update(litellm_params) return data @@ -121,8 +119,7 @@ async def vector_store_search( ) data = await _read_request_body(request=request) - if "vector_store_id" not in data: - data["vector_store_id"] = vector_store_id + data["vector_store_id"] = vector_store_id # Check for legacy vector store registry (non-managed vector stores) data = await _update_request_data_with_litellm_managed_vector_store_registry( diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 061a8aaa240..827bbba630e 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -1,7 +1,9 @@ +import json from typing import Any, Dict, Literal, Optional from fastapi import HTTPException, Request +import litellm from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -13,6 +15,21 @@ from litellm.types.vector_stores import LiteLLM_ManagedVectorStore from litellm.utils import ProviderConfigManager +def _normalize_litellm_params( + vector_store: LiteLLM_ManagedVectorStore, +) -> LiteLLM_ManagedVectorStore: + litellm_params = vector_store.get("litellm_params") + if isinstance(litellm_params, str): + normalized = LiteLLM_ManagedVectorStore(**dict(vector_store)) + try: + parsed = json.loads(litellm_params) + normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {} + except (TypeError, ValueError): + normalized["litellm_params"] = {} + return normalized + return vector_store + + def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: return ( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN @@ -120,6 +137,85 @@ async def can_user_access_vector_store( return False +async def get_litellm_managed_vector_store( + vector_store_id: str, +) -> Optional[LiteLLM_ManagedVectorStore]: + """ + Resolve a LiteLLM-managed vector store from the registry or database. + + Provider-native vector store IDs will not be present in either location and + return None, preserving direct provider behavior while still protecting + LiteLLM-managed multi-tenant stores. + """ + if not vector_store_id: + return None + + if litellm.vector_store_registry is not None: + try: + vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( + vector_store_id=vector_store_id + ) + if vector_store is not None: + return _normalize_litellm_params(vector_store) + except Exception as e: + verbose_proxy_logger.debug( + "Failed to resolve vector store id=%s from registry: %s", + vector_store_id, + e, + ) + + try: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return None + row = await prisma_client.db.litellm_managedvectorstorestable.find_unique( + where={"vector_store_id": vector_store_id} + ) + if row is None: + return None + return _normalize_litellm_params(LiteLLM_ManagedVectorStore(**row.model_dump())) + except Exception as e: + verbose_proxy_logger.debug( + "Failed to resolve vector store id=%s from database: %s", + vector_store_id, + e, + ) + return None + + +async def assert_user_can_access_vector_store( + vector_store: LiteLLM_ManagedVectorStore, + user_api_key_dict: UserAPIKeyAuth, + detail: str = "Access denied: You do not have permission to access this vector store", +) -> None: + """Raise 403 unless the caller can access the resolved vector store.""" + if not await can_user_access_vector_store(vector_store, user_api_key_dict): + raise HTTPException(status_code=403, detail=detail) + + +async def assert_user_can_access_vector_store_id( + vector_store_id: str, + user_api_key_dict: UserAPIKeyAuth, + detail: str = "Access denied: You do not have permission to access this vector store", +) -> Optional[LiteLLM_ManagedVectorStore]: + """ + Resolve a managed vector store id and enforce ownership if it exists. + + Unknown ids are treated as provider-native ids and are not rejected here. + """ + vector_store = await get_litellm_managed_vector_store( + vector_store_id=vector_store_id + ) + if vector_store is not None: + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + detail=detail, + ) + return vector_store + + def _does_endpoint_match(endpoint_path: str, request_path: str) -> bool: if endpoint_path in request_path: return True diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index 7cdf865692b..ae8dc602e82 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -17,6 +17,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( prepare_data_with_credentials, ) from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store_id, is_allowed_to_call_vector_store_files_endpoint, ) from litellm.types.utils import LlmProviders @@ -363,8 +364,11 @@ async def vector_store_file_create( ) data = await _read_request_body(request=request) - if "vector_store_id" not in data: - data["vector_store_id"] = vector_store_id + data["vector_store_id"] = vector_store_id + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs if present in request body original_managed_file_id = None @@ -459,6 +463,11 @@ async def vector_store_file_list( query_params = dict(request.query_params) data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id} data.update(query_params) + data["vector_store_id"] = vector_store_id + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) data = _update_request_data_with_litellm_managed_vector_store_registry( data=data, vector_store_id=vector_store_id, llm_router=llm_router @@ -541,6 +550,10 @@ async def vector_store_file_retrieve( "vector_store_id": vector_store_id, "file_id": file_id, } + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -635,6 +648,10 @@ async def vector_store_file_content( "vector_store_id": vector_store_id, "file_id": file_id, } + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -729,6 +746,10 @@ async def vector_store_file_update( data = await _read_request_body(request=request) data["vector_store_id"] = vector_store_id data["file_id"] = file_id + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -823,6 +844,10 @@ async def vector_store_file_delete( "vector_store_id": vector_store_id, "file_id": file_id, } + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index 6d28d670979..13f2f27d3fa 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -377,15 +377,11 @@ def search( _is_async = kwargs.pop("asearch", False) is True # pull credentials from registry if available - vector_store_id_for_credentials = kwargs.get("vector_store_id", vector_store_id) - if ( - litellm.vector_store_registry is not None - and vector_store_id_for_credentials is not None - ): + if litellm.vector_store_registry is not None and vector_store_id is not None: try: registry_credentials = ( litellm.vector_store_registry.get_credentials_for_vector_store( - vector_store_id_for_credentials + vector_store_id ) ) kwargs.update(registry_credentials) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py new file mode 100644 index 00000000000..c160c5aceb3 --- /dev/null +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -0,0 +1,362 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException, Request, Response + +import litellm +from litellm.proxy._types import UserAPIKeyAuth + + +def _mock_request() -> MagicMock: + request = MagicMock(spec=Request) + request.headers = {} + request.method = "POST" + request.query_params = {} + request.url.path = "/v1/vector_stores/vs_path/search" + return request + + +@pytest.mark.asyncio +async def test_vector_store_search_forces_path_id_over_body_id(): + from litellm.proxy.vector_store_endpoints.endpoints import vector_store_search + + captured_data = {} + + async def fake_base_process(self, **kwargs): + captured_data.update(self.data) + return {"ok": True} + + request = _mock_request() + with ( + patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock( + return_value={ + "vector_store_id": "vs_body_victim", + "query": "test", + } + ), + ), + patch.object(litellm, "vector_store_registry", None), + patch( + "litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=fake_base_process, + ), + ): + response = await vector_store_search( + request=request, + vector_store_id="vs_path_allowed", + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert response == {"ok": True} + assert captured_data["vector_store_id"] == "vs_path_allowed" + + +@pytest.mark.asyncio +async def test_vector_store_file_create_forces_path_id_over_body_id(): + from litellm.proxy.vector_store_files_endpoints.endpoints import ( + vector_store_file_create, + ) + + captured_data = {} + + async def fake_base_process(self, **kwargs): + captured_data.update(self.data) + return {"ok": True} + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_path_allowed", + "custom_llm_provider": "openai", + "team_id": "team-a", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock( + return_value={ + "vector_store_id": "vs_body_victim", + "file_id": "file_123", + } + ), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=fake_base_process, + ), + ): + response = await vector_store_file_create( + vector_store_id="vs_path_allowed", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert response == {"ok": True} + assert captured_data["vector_store_id"] == "vs_path_allowed" + mock_registry.get_litellm_managed_vector_store_from_registry.assert_any_call( + vector_store_id="vs_path_allowed" + ) + + +@pytest.mark.asyncio +async def test_vector_store_file_create_denies_other_team_path_store(): + from litellm.proxy.vector_store_files_endpoints.endpoints import ( + vector_store_file_create, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock(return_value={"file_id": "file_123"}), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=AsyncMock(), + ) as mock_base_process, + ): + with pytest.raises(HTTPException) as exc_info: + await vector_store_file_create( + vector_store_id="vs_other_team", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + mock_base_process.assert_not_called() + + +@pytest.mark.asyncio +async def test_rag_query_denies_nested_other_team_vector_store(): + from litellm.proxy.rag_endpoints.endpoints import rag_query + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.rag_endpoints.endpoints._read_request_body", + new=AsyncMock( + return_value={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "retrieval_config": {"vector_store_id": "vs_other_team"}, + } + ), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", + new=AsyncMock(), + ) as mock_aquery, + ): + with pytest.raises(HTTPException) as exc_info: + await rag_query( + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + mock_aquery.assert_not_called() + + +@pytest.mark.asyncio +async def test_rag_ingest_denies_nested_other_team_vector_store(): + from litellm.proxy.rag_endpoints.endpoints import rag_ingest + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.rag_endpoints.endpoints.parse_rag_ingest_request", + new=AsyncMock( + return_value=( + { + "vector_store": { + "custom_llm_provider": "openai", + "vector_store_id": "vs_other_team", + } + }, + None, + "https://example.com/file.txt", + None, + ) + ), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new=AsyncMock(), + ) as mock_aingest, + ): + with pytest.raises(HTTPException) as exc_info: + await rag_ingest( + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + mock_aingest.assert_not_called() + + +@pytest.mark.asyncio +async def test_responses_file_search_denies_other_team_vector_store(): + from litellm.proxy.common_request_processing import ( + _authorize_response_file_search_vector_stores, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + with patch.object(litellm, "vector_store_registry", mock_registry): + with pytest.raises(HTTPException) as exc_info: + await _authorize_response_file_search_vector_stores( + data={ + "tools": [ + { + "type": "file_search", + "vector_store_ids": ["vs_other_team"], + } + ] + }, + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_vertex_discovery_denies_other_team_vector_store_credentials(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _base_vertex_proxy_route, + ) + + request = _mock_request() + request.method = "GET" + vector_store_credentials = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "vertex_ai", + "team_id": "team-b", + } + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth", + new=AsyncMock(return_value=UserAPIKeyAuth(team_id="team-a")), + ): + with pytest.raises(HTTPException) as exc_info: + await _base_vertex_proxy_route( + endpoint="projects/p/locations/us-central1/dataStores/vs_other_team", + request=request, + fastapi_response=Response(), + get_vertex_pass_through_handler=MagicMock(), + router_credentials=vector_store_credentials, + ) + + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_vertex_discovery_denies_unregistered_vector_store_id(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + vertex_discovery_proxy_route, + ) + + request = _mock_request() + request.method = "GET" + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_vector_store_credentials", + return_value=None, + ): + with pytest.raises(HTTPException) as exc_info: + await vertex_discovery_proxy_route( + endpoint="projects/p/locations/us-central1/dataStores/vs_unknown", + request=request, + fastapi_response=Response(), + ) + + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_milvus_passthrough_denies_other_team_vector_store_index(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + milvus_proxy_route, + ) + + request = _mock_request() + request.url.path = "/milvus/v2/vectordb/entities/search" + + index_object = MagicMock() + index_object.litellm_params.vector_store_name = "tenant-b-store" + index_object.litellm_params.vector_store_index = "tenant_b_collection" + + mock_index_registry = MagicMock() + mock_index_registry.is_vector_store_index.return_value = True + mock_index_registry.get_vector_store_index_by_name.return_value = index_object + + mock_vector_registry = MagicMock() + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "milvus", + "team_id": "team-b", + "litellm_params": {"api_base": "https://milvus.example.com"}, + } + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body", + new=AsyncMock(return_value={"collectionName": "managed_index"}), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint", + return_value=True, + ), + patch.object(litellm, "vector_store_index_registry", mock_index_registry), + patch.object(litellm, "vector_store_registry", mock_vector_registry), + ): + with pytest.raises(HTTPException) as exc_info: + await milvus_proxy_route( + endpoint="v2/vectordb/entities/search", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 From 363c0de6f70367b27925e41a2d40e91f54c2c945 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 15:13:24 -0700 Subject: [PATCH 19/56] chore(vector stores): address tenant guard followups --- litellm/proxy/_lazy_openapi_snapshot.json | 34 ++++---- .../llm_passthrough_endpoints.py | 17 ++-- litellm/proxy/rag_endpoints/endpoints.py | 34 ++++---- litellm/proxy/vector_store_endpoints/utils.py | 26 +++++-- .../test_vector_store_tenant_guard.py | 77 ++++++++++++++++--- 5 files changed, 128 insertions(+), 60 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8331f748c6e..dfebd1baf30 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__post", "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__post", "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__post", "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__post", "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__post", "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__post", "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__post", "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__post", "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__post", "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__post", "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_options", "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_options", "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_options", "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_options", "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_options", "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_options", "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_options", "parameters": [ { "in": "path", diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index ddb6717cb2e..ce103f806e1 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -48,6 +48,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( assert_user_can_access_vector_store, + get_litellm_managed_vector_store, is_allowed_to_call_vector_store_endpoint, ) from litellm.secret_managers.main import get_secret_str @@ -1927,11 +1928,11 @@ async def vertex_discovery_proxy_route( "Extracted vector store ID from endpoint: %s", vector_store_id ) - # Retrieve vector store credentials from the registry - vector_store_credentials = ( - passthrough_endpoint_router.get_vector_store_credentials( - vector_store_id=vector_store_id - ) + # Retrieve LiteLLM-managed vector store credentials if the datastore id + # is registered with LiteLLM. Unknown datastore ids keep the existing + # direct Vertex pass-through behavior. + vector_store_credentials = await get_litellm_managed_vector_store( + vector_store_id=vector_store_id ) if vector_store_credentials: @@ -1939,14 +1940,10 @@ async def vertex_discovery_proxy_route( "Found vector store credentials for ID: %s", vector_store_id ) else: - verbose_proxy_logger.warning( + verbose_proxy_logger.debug( "Vector store ID %s found in endpoint but no credentials found in registry", vector_store_id, ) - raise HTTPException( - status_code=403, - detail="Access denied: You do not have permission to access this vector store", - ) discovery_handler = get_vertex_pass_through_handler(call_type="discovery") return await _base_vertex_proxy_route( diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 9e6093a47a1..3a50b703bc4 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -31,21 +31,27 @@ router = APIRouter() def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids: set[str] = set() + payload_stack = [payload] - if isinstance(payload, dict): - for key, value in payload.items(): - if key == "vector_store_id": - if not isinstance(value, str) or not value: - raise HTTPException( - status_code=400, - detail={"error": "vector_store_id must be a non-empty string"}, - ) - vector_store_ids.add(value) - continue - vector_store_ids.update(_collect_vector_store_ids_from_payload(value)) - elif isinstance(payload, list): - for item in payload: - vector_store_ids.update(_collect_vector_store_ids_from_payload(item)) + while payload_stack: + current_payload = payload_stack.pop() + + if isinstance(current_payload, dict): + for key, value in current_payload.items(): + if key == "vector_store_id": + if not isinstance(value, str) or not value: + raise HTTPException( + status_code=400, + detail={ + "error": "vector_store_id must be a non-empty string" + }, + ) + vector_store_ids.add(value) + continue + if isinstance(value, (dict, list)): + payload_stack.append(value) + elif isinstance(current_payload, list): + payload_stack.extend(current_payload) return vector_store_ids diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 827bbba630e..c09810d06dc 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -141,7 +141,7 @@ async def get_litellm_managed_vector_store( vector_store_id: str, ) -> Optional[LiteLLM_ManagedVectorStore]: """ - Resolve a LiteLLM-managed vector store from the registry or database. + Resolve a LiteLLM-managed vector store from the registry or shared cache. Provider-native vector store IDs will not be present in either location and return None, preserving direct provider behavior while still protecting @@ -165,19 +165,31 @@ async def get_litellm_managed_vector_store( ) try: - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.auth.auth_checks import ( + get_managed_vector_store_rows_by_uuids, + ) + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if prisma_client is None: return None - row = await prisma_client.db.litellm_managedvectorstorestable.find_unique( - where={"vector_store_id": vector_store_id} + rows = await get_managed_vector_store_rows_by_uuids( + uuids=[vector_store_id], + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, ) - if row is None: + if not rows: return None - return _normalize_litellm_params(LiteLLM_ManagedVectorStore(**row.model_dump())) + return _normalize_litellm_params( + LiteLLM_ManagedVectorStore(**rows[0].model_dump()) + ) except Exception as e: verbose_proxy_logger.debug( - "Failed to resolve vector store id=%s from database: %s", + "Failed to resolve vector store id=%s from shared cache: %s", vector_store_id, e, ) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index c160c5aceb3..2d6295e4029 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -4,7 +4,7 @@ import pytest from fastapi import HTTPException, Request, Response import litellm -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth def _mock_request() -> MagicMock: @@ -288,7 +288,53 @@ async def test_vertex_discovery_denies_other_team_vector_store_credentials(): @pytest.mark.asyncio -async def test_vertex_discovery_denies_unregistered_vector_store_id(): +async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback(): + from litellm.proxy.vector_store_endpoints.utils import ( + get_litellm_managed_vector_store, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = None + cache_helper = AsyncMock( + return_value=[ + LiteLLM_ManagedVectorStoresTable( + vector_store_id="vs_cached", + custom_llm_provider="openai", + vector_store_name=None, + vector_store_description=None, + vector_store_metadata=None, + created_at=None, + updated_at=None, + litellm_credential_name=None, + litellm_params={"api_base": "https://example.com"}, + team_id="team-a", + user_id=None, + ) + ] + ) + + with ( + patch.object(litellm, "vector_store_registry", mock_registry), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + new=cache_helper, + ), + ): + vector_store = await get_litellm_managed_vector_store( + vector_store_id="vs_cached" + ) + + assert vector_store is not None + assert vector_store["vector_store_id"] == "vs_cached" + assert vector_store["team_id"] == "team-a" + cache_helper.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id(): from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( vertex_discovery_proxy_route, ) @@ -296,18 +342,25 @@ async def test_vertex_discovery_denies_unregistered_vector_store_id(): request = _mock_request() request.method = "GET" - with patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_vector_store_credentials", - return_value=None, + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_managed_vector_store", + new=AsyncMock(return_value=None), + ) as mock_lookup, + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._base_vertex_proxy_route", + new=AsyncMock(return_value={"ok": True}), + ) as mock_base_route, ): - with pytest.raises(HTTPException) as exc_info: - await vertex_discovery_proxy_route( - endpoint="projects/p/locations/us-central1/dataStores/vs_unknown", - request=request, - fastapi_response=Response(), - ) + response = await vertex_discovery_proxy_route( + endpoint="projects/p/locations/us-central1/dataStores/vs_unknown", + request=request, + fastapi_response=Response(), + ) - assert exc_info.value.status_code == 403 + assert response == {"ok": True} + mock_lookup.assert_awaited_once_with(vector_store_id="vs_unknown") + assert mock_base_route.call_args.kwargs["router_credentials"] is None @pytest.mark.asyncio From aef71ae2d5262d0e7621210989ccd4a231f8658e Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 15:17:20 -0700 Subject: [PATCH 20/56] chore(proxy): stabilize lazy openapi snapshot --- litellm/proxy/_lazy_openapi_snapshot.json | 28 +++++++++++------------ litellm/proxy/_lazy_openapi_snapshot.py | 20 +++++++++++++++- 2 files changed, 33 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index dfebd1baf30..46a514c0870 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__post", + "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__post", + "operationId": "anthropic_proxy_route_anthropic__endpoint__get", "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__post", + "operationId": "anthropic_proxy_route_anthropic__endpoint__patch", "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__post", + "operationId": "anthropic_proxy_route_anthropic__endpoint__put", "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__post", + "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__post", + "operationId": "langfuse_proxy_route_langfuse__endpoint__get", "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__post", + "operationId": "langfuse_proxy_route_langfuse__endpoint__patch", "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__post", + "operationId": "langfuse_proxy_route_langfuse__endpoint__put", "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_options", + "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_options", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_get", "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_options", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_head", "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_options", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_patch", "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_options", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_post", "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_options", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", "parameters": [ { "in": "path", diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 315f6a9742a..9c317a5365b 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -13,6 +13,16 @@ from pathlib import Path from typing import Dict, Optional SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" +HTTP_METHOD_SUFFIXES = { + "delete", + "get", + "head", + "options", + "patch", + "post", + "put", + "trace", +} def load_snapshot() -> Optional[Dict[str, Dict]]: @@ -54,8 +64,16 @@ def generate_snapshot() -> Dict[str, Dict]: 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(): - for op in path_ops.values(): + for method, op in path_ops.items(): if isinstance(op, dict): + operation_id = op.get("operationId") + if isinstance(operation_id, str): + for suffix in HTTP_METHOD_SUFFIXES: + if operation_id.endswith(f"_{suffix}"): + op["operationId"] = ( + operation_id[: -len(suffix)] + method + ) + break op["tags"] = [feat.name] fragments[feat.name] = { "paths": full.get("paths", {}), From ce0c55701298830f32296fc7badf7bf7662f0402 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:09:26 -0700 Subject: [PATCH 21/56] chore(vector stores): address access review followups --- litellm/proxy/rag_endpoints/endpoints.py | 16 +++- .../proxy/vector_store_endpoints/endpoints.py | 23 ----- litellm/proxy/vector_store_endpoints/utils.py | 13 ++- .../vector_store_files_endpoints/endpoints.py | 83 +++++++++++++------ .../test_vector_store_tenant_guard.py | 40 ++++++++- 5 files changed, 118 insertions(+), 57 deletions(-) diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 3a50b703bc4..ecb13638648 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -15,6 +15,7 @@ from fastapi.responses import ORJSONResponse import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( @@ -31,10 +32,17 @@ router = APIRouter() def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids: set[str] = set() - payload_stack = [payload] + payload_stack = [(payload, 0)] while payload_stack: - current_payload = payload_stack.pop() + current_payload, depth = payload_stack.pop() + if depth > DEFAULT_MAX_RECURSE_DEPTH: + raise HTTPException( + status_code=400, + detail={ + "error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values" + }, + ) if isinstance(current_payload, dict): for key, value in current_payload.items(): @@ -49,9 +57,9 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids.add(value) continue if isinstance(value, (dict, list)): - payload_stack.append(value) + payload_stack.append((value, depth + 1)) elif isinstance(current_payload, list): - payload_stack.extend(current_payload) + payload_stack.extend((item, depth + 1) for item in current_payload) return vector_store_ids diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 05423d9843d..86e316e7f40 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -20,29 +20,6 @@ router = APIRouter() ######################################################## -async def _check_vector_store_access( - vector_store: LiteLLM_ManagedVectorStore, - user_api_key_dict: UserAPIKeyAuth, -) -> bool: - """ - Check if the user has access to the vector store. - - Delegates to :func:`can_user_access_vector_store`, which honors: - - PROXY_ADMIN bypass - - legacy vector stores with no team_id - - key-level and team-level ``object_permission.vector_stores`` allowlists - - team_id match between key and store - """ - try: - await assert_user_can_access_vector_store( - vector_store=vector_store, - user_api_key_dict=user_api_key_dict, - ) - return True - except HTTPException: - return False - - async def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index c09810d06dc..657b520b271 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -158,11 +158,15 @@ async def get_litellm_managed_vector_store( if vector_store is not None: return _normalize_litellm_params(vector_store) except Exception as e: - verbose_proxy_logger.debug( + verbose_proxy_logger.warning( "Failed to resolve vector store id=%s from registry: %s", vector_store_id, e, ) + raise HTTPException( + status_code=500, + detail="Unable to validate vector store access", + ) from e try: from litellm.proxy.auth.auth_checks import ( @@ -188,12 +192,15 @@ async def get_litellm_managed_vector_store( LiteLLM_ManagedVectorStore(**rows[0].model_dump()) ) except Exception as e: - verbose_proxy_logger.debug( + verbose_proxy_logger.warning( "Failed to resolve vector store id=%s from shared cache: %s", vector_store_id, e, ) - return None + raise HTTPException( + status_code=500, + detail="Unable to validate vector store access", + ) from e async def assert_user_can_access_vector_store( diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index ae8dc602e82..346a847c5dd 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -21,6 +21,7 @@ from litellm.proxy.vector_store_endpoints.utils import ( is_allowed_to_call_vector_store_files_endpoint, ) from litellm.types.utils import LlmProviders +from litellm.types.vector_stores import LiteLLM_ManagedVectorStore if TYPE_CHECKING: from litellm.router import Router @@ -194,6 +195,8 @@ def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, llm_router: Optional["Router"] = None, + managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None, + should_lookup_registry: bool = True, ) -> Dict: """ Update request data with model routing information from managed vector store. @@ -263,23 +266,27 @@ def _update_request_data_with_litellm_managed_vector_store_registry( return data - # Legacy path: Check vector store registry for non-managed vector stores - if litellm.vector_store_registry is not None: + # Legacy path: Check vector store registry for non-managed vector stores. + vector_store_to_run = managed_vector_store + if ( + vector_store_to_run is None + and should_lookup_registry + and litellm.vector_store_registry is not None + ): vector_store_to_run = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( vector_store_id=vector_store_id ) - if vector_store_to_run is not None: - if "custom_llm_provider" in vector_store_to_run: - data["custom_llm_provider"] = vector_store_to_run.get( - "custom_llm_provider" - ) - if "litellm_credential_name" in vector_store_to_run: - data["litellm_credential_name"] = vector_store_to_run.get( - "litellm_credential_name" - ) - if "litellm_params" in vector_store_to_run: - litellm_params = vector_store_to_run.get("litellm_params", {}) or {} - data.update(litellm_params) + + if vector_store_to_run is not None: + if "custom_llm_provider" in vector_store_to_run: + data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider") + if "litellm_credential_name" in vector_store_to_run: + data["litellm_credential_name"] = vector_store_to_run.get( + "litellm_credential_name" + ) + if "litellm_params" in vector_store_to_run: + litellm_params = vector_store_to_run.get("litellm_params", {}) or {} + data.update(litellm_params) return data @@ -365,7 +372,7 @@ async def vector_store_file_create( data = await _read_request_body(request=request) data["vector_store_id"] = vector_store_id - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -379,7 +386,11 @@ async def vector_store_file_create( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -464,13 +475,17 @@ async def vector_store_file_list( data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id} data.update(query_params) data["vector_store_id"] = vector_store_id - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -550,7 +565,7 @@ async def vector_store_file_retrieve( "vector_store_id": vector_store_id, "file_id": file_id, } - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -562,7 +577,11 @@ async def vector_store_file_retrieve( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -648,7 +667,7 @@ async def vector_store_file_content( "vector_store_id": vector_store_id, "file_id": file_id, } - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -660,7 +679,11 @@ async def vector_store_file_content( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -746,7 +769,7 @@ async def vector_store_file_update( data = await _read_request_body(request=request) data["vector_store_id"] = vector_store_id data["file_id"] = file_id - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -758,7 +781,11 @@ async def vector_store_file_update( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -844,7 +871,7 @@ async def vector_store_file_delete( "vector_store_id": vector_store_id, "file_id": file_id, } - await assert_user_can_access_vector_store_id( + managed_vector_store = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict, ) @@ -856,7 +883,11 @@ async def vector_store_file_delete( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index 2d6295e4029..380b3963c94 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -99,7 +99,8 @@ async def test_vector_store_file_create_forces_path_id_over_body_id(): assert response == {"ok": True} assert captured_data["vector_store_id"] == "vs_path_allowed" - mock_registry.get_litellm_managed_vector_store_from_registry.assert_any_call( + assert captured_data["custom_llm_provider"] == "openai" + mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with( vector_store_id="vs_path_allowed" ) @@ -227,6 +228,25 @@ async def test_rag_ingest_denies_nested_other_team_vector_store(): mock_aingest.assert_not_called() +def test_rag_payload_scan_rejects_excessive_nesting(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 1): + current["nested"] = {} + current = current["nested"] + current["vector_store_id"] = "vs_too_deep" + + with pytest.raises(HTTPException) as exc_info: + _collect_vector_store_ids_from_payload(payload) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio async def test_responses_file_search_denies_other_team_vector_store(): from litellm.proxy.common_request_processing import ( @@ -333,6 +353,24 @@ async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback cache_helper.assert_awaited_once() +@pytest.mark.asyncio +async def test_get_managed_vector_store_fails_closed_on_lookup_error(): + from litellm.proxy.vector_store_endpoints.utils import ( + get_litellm_managed_vector_store, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = ( + RuntimeError("registry unavailable") + ) + + with patch.object(litellm, "vector_store_registry", mock_registry): + with pytest.raises(HTTPException) as exc_info: + await get_litellm_managed_vector_store(vector_store_id="vs_registry_only") + + assert exc_info.value.status_code == 500 + + @pytest.mark.asyncio async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id(): from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( 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 22/56] 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 1201a0ba5cfe3fd607f6eed288cb34f55d519910 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:16:37 -0700 Subject: [PATCH 23/56] test(vector stores): pin no-db registry fallback case --- .../vector_store_endpoints/test_vector_store_endpoints.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 6dd0e0e68a6..e67a04c749a 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -156,8 +156,11 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry(): vector_store_id="test_store_id" ) - # Test with no vector store registry - with patch.object(litellm, "vector_store_registry", None): + # Test with no vector store registry or DB fallback + with ( + patch.object(litellm, "vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + ): original_data = {"existing_key": "existing_value"} result = await _update_request_data_with_litellm_managed_vector_store_registry( data=original_data, vector_store_id=vector_store_id 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 24/56] 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 49ccb3369c0e3b203063dd655eae7a9f6df308ab Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:33:42 -0700 Subject: [PATCH 25/56] test(vector stores): pin rag scan depth boundary --- litellm/proxy/rag_endpoints/endpoints.py | 26 +++++++++++++------ .../test_vector_store_tenant_guard.py | 16 ++++++++++++ 2 files changed, 34 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index ecb13638648..df774c1d321 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -30,6 +30,15 @@ from litellm.proxy.vector_store_endpoints.utils import ( router = APIRouter() +def _raise_vector_store_scan_depth_exceeded() -> None: + raise HTTPException( + status_code=400, + detail={ + "error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values" + }, + ) + + def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids: set[str] = set() payload_stack = [(payload, 0)] @@ -37,12 +46,7 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: while payload_stack: current_payload, depth = payload_stack.pop() if depth > DEFAULT_MAX_RECURSE_DEPTH: - raise HTTPException( - status_code=400, - detail={ - "error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values" - }, - ) + _raise_vector_store_scan_depth_exceeded() if isinstance(current_payload, dict): for key, value in current_payload.items(): @@ -57,9 +61,15 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids.add(value) continue if isinstance(value, (dict, list)): - payload_stack.append((value, depth + 1)) + next_depth = depth + 1 + if next_depth > DEFAULT_MAX_RECURSE_DEPTH: + _raise_vector_store_scan_depth_exceeded() + payload_stack.append((value, next_depth)) elif isinstance(current_payload, list): - payload_stack.extend((item, depth + 1) for item in current_payload) + next_depth = depth + 1 + if current_payload and next_depth > DEFAULT_MAX_RECURSE_DEPTH: + _raise_vector_store_scan_depth_exceeded() + payload_stack.extend((item, next_depth) for item in current_payload) return vector_store_ids diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index 380b3963c94..ecde853b0af 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -247,6 +247,22 @@ def test_rag_payload_scan_rejects_excessive_nesting(): assert exc_info.value.status_code == 400 +def test_rag_payload_scan_accepts_vector_store_id_at_depth_limit(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH): + current["nested"] = {} + current = current["nested"] + current["vector_store_id"] = "vs_at_limit" + + assert _collect_vector_store_ids_from_payload(payload) == {"vs_at_limit"} + + @pytest.mark.asyncio async def test_responses_file_search_denies_other_team_vector_store(): from litellm.proxy.common_request_processing import ( From 32272908d3b0d610801e60a90643d77ca3841ed3 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 16:41:13 -0700 Subject: [PATCH 26/56] test(vector stores): isolate provider-native guard case --- .../vector_store_endpoints/test_vector_store_tenant_guard.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index ecde853b0af..6cbf260457b 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -38,6 +38,7 @@ async def test_vector_store_search_forces_path_id_over_body_id(): ), ), patch.object(litellm, "vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", None), patch( "litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", new=fake_base_process, 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 27/56] 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 2922da9b644e675b0114609fbef0bcf011bcc6c6 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 17:00:43 -0700 Subject: [PATCH 28/56] test(vector stores): cover azure passthrough guard --- .../test_vector_store_tenant_guard.py | 54 +++++++++++++++++++ 1 file changed, 54 insertions(+) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index 6cbf260457b..0ec94b3337d 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -468,3 +468,57 @@ async def test_milvus_passthrough_denies_other_team_vector_store_index(): ) assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_azure_passthrough_denies_other_team_vector_store_index(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + azure_proxy_route, + ) + + request = _mock_request() + request.url.path = "/azure/indexes/managed_index/docs/search" + + index_object = MagicMock() + index_object.litellm_params.vector_store_name = "tenant-b-store" + + mock_index_registry = MagicMock() + mock_index_registry.is_vector_store_index.side_effect = ( + lambda vector_store_index_name: vector_store_index_name == "managed_index" + ) + mock_index_registry.get_vector_store_index_by_name.return_value = index_object + + mock_vector_registry = MagicMock() + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "azure_ai", + "team_id": "team-b", + "litellm_params": {"api_base": "https://azure.example.com"}, + } + + with ( + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", + return_value=False, + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint", + return_value=True, + ), + patch.object(litellm, "vector_store_index_registry", mock_index_registry), + patch.object(litellm, "vector_store_registry", mock_vector_registry), + ): + with pytest.raises(HTTPException) as exc_info: + await azure_proxy_route( + endpoint="indexes/managed_index/docs/search", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 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 29/56] 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 30/56] 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 06502d19a7d468689d860398b5d5833c6cc6ab50 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 17:28:02 -0700 Subject: [PATCH 31/56] test(vector stores): allow primitive rag depth boundary --- litellm/proxy/rag_endpoints/endpoints.py | 36 ++++++++++++++----- .../test_vector_store_tenant_guard.py | 16 +++++++++ 2 files changed, 44 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index df774c1d321..498d77f7535 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -39,6 +39,23 @@ def _raise_vector_store_scan_depth_exceeded() -> None: ) +def _append_payload_to_scan_stack( + payload_stack: list[tuple[Any, int]], + value: Any, + next_depth: int, +) -> None: + if isinstance(value, dict): + if next_depth > DEFAULT_MAX_RECURSE_DEPTH: + _raise_vector_store_scan_depth_exceeded() + payload_stack.append((value, next_depth)) + elif isinstance(value, list): + if next_depth > DEFAULT_MAX_RECURSE_DEPTH: + if any(isinstance(item, (dict, list)) for item in value): + _raise_vector_store_scan_depth_exceeded() + return + payload_stack.append((value, next_depth)) + + def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids: set[str] = set() payload_stack = [(payload, 0)] @@ -61,15 +78,18 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: vector_store_ids.add(value) continue if isinstance(value, (dict, list)): - next_depth = depth + 1 - if next_depth > DEFAULT_MAX_RECURSE_DEPTH: - _raise_vector_store_scan_depth_exceeded() - payload_stack.append((value, next_depth)) + _append_payload_to_scan_stack( + payload_stack=payload_stack, + value=value, + next_depth=depth + 1, + ) elif isinstance(current_payload, list): - next_depth = depth + 1 - if current_payload and next_depth > DEFAULT_MAX_RECURSE_DEPTH: - _raise_vector_store_scan_depth_exceeded() - payload_stack.extend((item, next_depth) for item in current_payload) + for item in current_payload: + _append_payload_to_scan_stack( + payload_stack=payload_stack, + value=item, + next_depth=depth + 1, + ) return vector_store_ids diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index 0ec94b3337d..48262afd363 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -264,6 +264,22 @@ def test_rag_payload_scan_accepts_vector_store_id_at_depth_limit(): assert _collect_vector_store_ids_from_payload(payload) == {"vs_at_limit"} +def test_rag_payload_scan_ignores_primitive_list_beyond_depth_limit(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH): + current["nested"] = {} + current = current["nested"] + current["labels"] = ["alpha", "beta"] + + assert _collect_vector_store_ids_from_payload(payload) == set() + + @pytest.mark.asyncio async def test_responses_file_search_denies_other_team_vector_store(): from litellm.proxy.common_request_processing import ( 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 32/56] 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 33/56] 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 34/56] 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 35/56] 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 36/56] 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 37/56] 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 38/56] 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 39/56] 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 40/56] 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 41/56] 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 42/56] 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 43/56] 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 336fe8276f2ae69a352eb73b83d7297d5746815c Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 21:59:56 -0700 Subject: [PATCH 44/56] chore(proxy): align resource model auth checks --- litellm/proxy/auth/auth_checks.py | 9 +- litellm/proxy/auth/auth_utils.py | 253 +++++++++++++++++- litellm/proxy/auth/user_api_key_auth.py | 117 ++++++-- .../proxy/auth/test_auth_utils.py | 112 ++++++++ .../proxy/auth/test_user_api_key_auth.py | 60 ++++- 5 files changed, 501 insertions(+), 50 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 65638ed6c1e..c0ce82b8916 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -61,6 +61,10 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.http_parsing_utils import ( + _safe_get_request_headers, + _safe_get_request_query_params, +) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( TOOL_CAPABLE_CALL_TYPES, @@ -485,7 +489,10 @@ async def common_checks( # noqa: PLR0915 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache _model: Optional[Union[str, List[str]]] = get_model_from_request( - request_body, route + request_data=request_body, + route=route, + request_headers=_safe_get_request_headers(request=request), + request_query_params=_safe_get_request_query_params(request=request), ) # 1. If team is blocked diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 91c8f2dd7c9..ba858a89fbc 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -2,7 +2,7 @@ import os import re import sys from functools import lru_cache -from typing import Any, List, Optional, Tuple +from typing import Any, Dict, List, Mapping, Optional, Tuple, Union from fastapi import HTTPException, Request, status @@ -942,20 +942,249 @@ def get_end_user_id_from_request_body( return None -def get_model_from_request( - request_data: dict, route: str +MODEL_ROUTING_HEADER_NAME = "x-litellm-model" +_MODEL_ROUTING_ROUTE_MARKERS = ( + "/files", + "/batches", + "/vector_stores", + "/skills", + "/evals", + "/fine_tuning", + "/videos", +) +_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS = ( + "/files", + "/batches", + "/skills", + "/evals", +) +_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS = ( + "/files", + "/batches", + "/fine_tuning", +) +_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = ( + "/files", + "/batches", + "/vector_stores", +) +_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",) +_MODEL_ROUTING_ID_FIELDS = ( + "file_id", + "input_file_id", + "output_file_id", + "error_file_id", + "batch_id", + "fine_tuning_job_id", + "training_file", + "validation_file", + "vector_store_id", + "video_id", + "character_id", +) + + +def _append_model_candidates(candidates: List[str], value: Any) -> None: + if value is None: + return + + if isinstance(value, str): + model_names = [model.strip() for model in value.split(",")] + elif isinstance(value, (list, tuple, set)): + for item in value: + _append_model_candidates(candidates=candidates, value=item) + return + else: + model_names = [str(value).strip()] + + candidates.extend(model for model in model_names if model) + + +def _dedupe_model_candidates(candidates: List[str]) -> List[str]: + deduped: List[str] = [] + for model in candidates: + if model not in deduped: + deduped.append(model) + return deduped + + +def _get_case_insensitive_mapping_value( + mapping: Optional[Mapping[str, Any]], key: str +) -> Any: + if not mapping: + return None + if key in mapping: + return mapping[key] + key_lower = key.lower() + for mapping_key, value in mapping.items(): + if str(mapping_key).lower() == key_lower: + return value + return None + + +def _route_matches_any_marker(route: str, markers: Tuple[str, ...]) -> bool: + normalized_route = route.lower() + return any(marker in normalized_route for marker in markers) + + +def _route_uses_model_routing_sources(route: str) -> bool: + return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS) + + +def _extract_models_from_managed_resource_id(resource_id: Any) -> List[str]: + if not isinstance(resource_id, str) or not resource_id: + return [] + + candidates: List[str] = [] + + try: + from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, + decode_model_from_file_id, + get_model_id_from_unified_batch_id, + get_models_from_unified_file_id, + ) + + _append_model_candidates( + candidates=candidates, value=decode_model_from_file_id(resource_id) + ) + unified_file_id = _is_base64_encoded_unified_file_id(resource_id) + if unified_file_id: + _append_model_candidates( + candidates=candidates, + value=get_models_from_unified_file_id(unified_file_id), + ) + _append_model_candidates( + candidates=candidates, + value=get_model_id_from_unified_batch_id(unified_file_id), + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from managed file/batch ID: %s", str(e) + ) + + try: + from litellm.llms.base_llm.managed_resources.utils import parse_unified_id + + parsed_id = parse_unified_id(resource_id) + if parsed_id: + _append_model_candidates( + candidates=candidates, value=parsed_id.get("model_id") + ) + _append_model_candidates( + candidates=candidates, value=parsed_id.get("target_model_names") + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from unified managed resource ID: %s", str(e) + ) + + try: + from litellm.types.videos.utils import ( + decode_character_id_with_provider, + decode_video_id_with_provider, + ) + + _append_model_candidates( + candidates=candidates, + value=decode_video_id_with_provider(resource_id).get("model_id"), + ) + _append_model_candidates( + candidates=candidates, + value=decode_character_id_with_provider(resource_id).get("model_id"), + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from managed video/character ID: %s", str(e) + ) + + return _dedupe_model_candidates(candidates) + + +def _extract_model_candidates_from_request( + request_data: dict, + route: str, + request_headers: Optional[Mapping[str, Any]] = None, + request_query_params: Optional[Mapping[str, Any]] = None, +) -> List[str]: + candidates: List[str] = [] + uses_model_routing_sources = _route_uses_model_routing_sources(route=route) + uses_header_or_query_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS + ) + uses_query_target_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS + ) + uses_body_target_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS + ) + uses_completion_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS + ) + + body_model = request_data.get("model") + _append_model_candidates(candidates, body_model) + if uses_body_target_model_sources or not body_model: + _append_model_candidates(candidates, request_data.get("target_model_names")) + if uses_completion_model_sources and isinstance( + request_data.get("completion"), dict + ): + _append_model_candidates(candidates, request_data["completion"].get("model")) + + if uses_model_routing_sources: + if uses_header_or_query_model_sources: + _append_model_candidates( + candidates, + _get_case_insensitive_mapping_value(request_query_params, "model"), + ) + _append_model_candidates( + candidates, + _get_case_insensitive_mapping_value( + request_headers, MODEL_ROUTING_HEADER_NAME + ), + ) + if uses_query_target_model_sources: + _append_model_candidates( + candidates, + _get_case_insensitive_mapping_value( + request_query_params, "target_model_names" + ), + ) + + for field in _MODEL_ROUTING_ID_FIELDS: + _append_model_candidates( + candidates, + _extract_models_from_managed_resource_id(request_data.get(field)), + ) + + return _dedupe_model_candidates(candidates) + + +def _format_model_candidates( + candidates: List[str], ) -> Optional[Union[str, List[str]]]: - # First try to get model from request_data - model = request_data.get("model") or request_data.get("target_model_names") + if not candidates: + return None + if len(candidates) == 1: + return candidates[0] + return candidates - if model is not None: - model_names = model.split(",") - if len(model_names) == 1: - model = model_names[0].strip() - else: - model = [m.strip() for m in model_names] - # If model not in request_data, try to extract from route +def get_model_from_request( + request_data: dict, + route: str, + request_headers: Optional[Mapping[str, Any]] = None, + request_query_params: Optional[Mapping[str, Any]] = None, +) -> Optional[Union[str, List[str]]]: + candidates = _extract_model_candidates_from_request( + request_data=request_data, + route=route, + request_headers=request_headers, + request_query_params=request_query_params, + ) + model = _format_model_candidates(candidates) + + # If no explicit model was found, try to extract from route if model is None: # Parse model from route that follows the pattern /openai/deployments/{model}/* match = re.match(r"/openai/deployments/([^/]+)", route) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index b7700feb5bb..9005327bfb2 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -11,7 +11,7 @@ import asyncio import re import secrets from datetime import datetime, timezone -from typing import Any, List, Optional, Tuple, cast +from typing import Any, List, Optional, Tuple, Union, cast import fastapi from fastapi import HTTPException, Request, WebSocket, status @@ -63,6 +63,7 @@ from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordin from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, + _safe_get_request_query_params, populate_request_with_path_params, ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body @@ -118,6 +119,29 @@ azure_apim_header = APIKeyHeader( ) +def _get_model_from_request_context( + request_data: dict, + route: str, + request: Optional[Request], +) -> Optional[Union[str, List[str]]]: + return get_model_from_request( + request_data=request_data, + route=route, + request_headers=_safe_get_request_headers(request=request), + request_query_params=_safe_get_request_query_params(request=request), + ) + + +def _get_model_names_for_budget_checks( + model: Optional[Union[str, List[str]]], +) -> List[str]: + if model is None: + return [] + if isinstance(model, str): + return [model] + return model + + def _get_bearer_token_or_received_api_key(api_key: str) -> str: if api_key.startswith("Bearer "): # ensure Bearer token passed in api_key = api_key.replace("Bearer ", "") # extract the token @@ -884,7 +908,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) # Check if model has zero cost - if so, skip all budget checks - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) skip_budget_checks = False if model is not None and llm_router is not None: from litellm.proxy.auth.auth_checks import _is_model_cost_zero @@ -1252,6 +1280,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 valid_token=valid_token, request_data=request_data, route=route, + request=request, llm_model_list=llm_model_list, llm_router=llm_router, ) @@ -1277,7 +1306,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 user_obj = None # Check 2a. Check if model has zero cost - if so, skip all budget checks - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) skip_budget_checks = False if model is not None and llm_router is not None: from litellm.proxy.auth.auth_checks import _is_model_cost_zero @@ -1395,21 +1428,29 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Check 5. Token Model Spend is under Model budget max_budget_per_model = valid_token.model_max_budget - current_model = request_data.get("model", None) + current_model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) + current_models = _get_model_names_for_budget_checks( + model=current_model + ) if ( max_budget_per_model is not None and isinstance(max_budget_per_model, dict) and len(max_budget_per_model) > 0 and prisma_client is not None - and current_model is not None + and current_models and valid_token.token is not None ): ## GET THE SPEND FOR THIS MODEL - await model_max_budget_limiter.is_key_within_model_budget( - user_api_key_dict=valid_token, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_key_within_model_budget( + user_api_key_dict=valid_token, + model=model_name, + ) # Check 5b. End-user model max budget end_user_mmb = valid_token.end_user_model_max_budget @@ -1417,14 +1458,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 end_user_mmb is not None and isinstance(end_user_mmb, dict) and len(end_user_mmb) > 0 - and current_model is not None + and current_models and valid_token.end_user_id is not None ): - await model_max_budget_limiter.is_end_user_within_model_budget( - end_user_id=valid_token.end_user_id, - end_user_model_max_budget=end_user_mmb, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_end_user_within_model_budget( + end_user_id=valid_token.end_user_id, + end_user_model_max_budget=end_user_mmb, + model=model_name, + ) # Check 6: Additional Common Checks across jwt + key auth if valid_token.team_id is not None: @@ -1851,7 +1893,11 @@ async def _run_centralized_common_checks( user_api_key_auth_obj.project_alias = project_object.project_alias skip_budget_checks = False - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) if model is not None and llm_router is not None: skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) @@ -2122,6 +2168,7 @@ async def _enforce_key_and_fallback_model_access( valid_token: UserAPIKeyAuth, request_data: dict, route: str, + request: Optional[Request], llm_model_list: Optional[list], llm_router: Optional[Any], ) -> None: @@ -2140,7 +2187,11 @@ async def _enforce_key_and_fallback_model_access( ): pass else: - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) fallback_models = cast( Optional[List[ALL_FALLBACK_MODEL_VALUES]], request_data.get("fallbacks", None), @@ -2227,11 +2278,17 @@ async def _run_post_custom_auth_checks( valid_token=valid_token, request_data=request_data, route=route, + request=request, llm_model_list=llm_model_list, llm_router=llm_router, ) - current_model = request_data.get("model", None) + current_model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) + current_models = _get_model_names_for_budget_checks(model=current_model) # 3. Check key-level model_max_budget max_budget_per_model = valid_token.model_max_budget @@ -2239,13 +2296,14 @@ async def _run_post_custom_auth_checks( max_budget_per_model is not None and isinstance(max_budget_per_model, dict) and len(max_budget_per_model) > 0 - and current_model is not None + and current_models and valid_token.token is not None ): - await model_max_budget_limiter.is_key_within_model_budget( - user_api_key_dict=valid_token, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_key_within_model_budget( + user_api_key_dict=valid_token, + model=model_name, + ) # 4. Check end-user model_max_budget end_user_mmb = valid_token.end_user_model_max_budget @@ -2253,14 +2311,15 @@ async def _run_post_custom_auth_checks( end_user_mmb is not None and isinstance(end_user_mmb, dict) and len(end_user_mmb) > 0 - and current_model is not None + and current_models and valid_token.end_user_id is not None ): - await model_max_budget_limiter.is_end_user_within_model_budget( - end_user_id=valid_token.end_user_id, - end_user_model_max_budget=end_user_mmb, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_end_user_within_model_budget( + end_user_id=valid_token.end_user_id, + end_user_model_max_budget=end_user_mmb, + model=model_name, + ) # team / user / end_user / project context objects are fetched by # the centralized common_checks gate in user_api_key_auth after diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 91f300b88ce..b3b6fdde670 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -2,6 +2,7 @@ Unit tests for auth_utils functions related to rate limiting and customer ID extraction. """ +import base64 from typing import Optional from unittest.mock import MagicMock, patch @@ -258,6 +259,117 @@ def test_get_model_from_request_vertex_passthrough_still_works(): assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro" +def test_get_model_from_request_includes_file_endpoint_header_model(): + assert ( + get_model_from_request( + request_data={}, + route="/v1/files", + request_headers={"X-LiteLLM-Model": "restricted-model"}, + ) + == "restricted-model" + ) + + +def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes(): + assert ( + get_model_from_request( + request_data={"model": "allowed-model"}, + route="/v1/chat/completions", + request_headers={"x-litellm-model": "restricted-model"}, + ) + == "allowed-model" + ) + + +def test_get_model_from_request_authorizes_all_file_routing_model_sources(): + models = get_model_from_request( + request_data={"model": "body-model"}, + route="/v1/files", + request_headers={"x-litellm-model": "header-model"}, + request_query_params={"target_model_names": "query-model-a,query-model-b"}, + ) + assert isinstance(models, list) + assert set(models) == { + "body-model", + "query-model-a", + "query-model-b", + "header-model", + } + + +def test_get_model_from_request_extracts_simple_encoded_file_id_model(): + from litellm.proxy.openai_files_endpoints.common_utils import ( + encode_file_id_with_model, + ) + + file_id = encode_file_id_with_model( + file_id="file-provider-id", + model="restricted-model", + ) + + assert ( + get_model_from_request( + request_data={"file_id": file_id}, + route="/v1/files/{file_id}", + ) + == "restricted-model" + ) + + +def test_get_model_from_request_extracts_unified_file_id_models(): + raw_unified_file_id = ( + "litellm_proxy:application/octet-stream;unified_id,test-id;" + "target_model_names,model-a,model-b;llm_output_file_id,file-provider-id" + ) + encoded_unified_file_id = ( + base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=") + ) + + assert get_model_from_request( + request_data={"file_id": encoded_unified_file_id}, + route="/v1/files/{file_id}", + ) == ["model-a", "model-b"] + + +def test_get_model_from_request_extracts_eval_completion_model(): + assert ( + get_model_from_request( + request_data={"completion": {"model": "judge-model"}}, + route="/v1/evals/{eval_id}/runs", + ) + == "judge-model" + ) + + +def test_get_model_from_request_includes_fine_tuning_target_model_query(): + assert ( + get_model_from_request( + request_data={}, + route="/v1/fine_tuning/jobs", + request_query_params={"target_model_names": "fine-tune-model"}, + ) + == "fine-tune-model" + ) + + +def test_get_model_from_request_extracts_video_id_model(): + from litellm.types.videos.utils import encode_video_id_with_provider + + video_id = encode_video_id_with_provider( + video_id="video-provider-id", + provider="openai", + model_id="video-model", + ) + + assert ( + get_model_from_request( + request_data={"video_id": video_id}, + route="/v1/videos/{video_id}", + ) + == "video-model" + ) + + def test_get_customer_user_header_returns_none_when_no_customer_role(): from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping 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 08f4bd0ebff..679b8fa6d2c 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,7 @@ -import asyncio import json import os import sys -from typing import Tuple +from types import SimpleNamespace from unittest.mock import ANY, AsyncMock, MagicMock, patch sys.path.insert( @@ -32,6 +31,13 @@ from litellm.proxy.auth.user_api_key_auth import ( ) +class _RoutingRequest: + def __init__(self, headers=None, query_params=None): + self.headers = headers or {} + self.query_params = query_params or {} + self.state = SimpleNamespace() + + def test_get_api_key(): bearer_token = "Bearer sk-12345678" api_key = "sk-12345678" @@ -107,6 +113,39 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit ) +@pytest.mark.asyncio +async def test_custom_auth_enforces_key_model_access_from_file_route_header_with_opt_in(): + valid_token = UserAPIKeyAuth(token="test_token", models=["allowed-model"]) + request = _RoutingRequest(headers={"x-litellm-model": "restricted-model"}) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + new_callable=AsyncMock, + ) as mock_can_key, + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"custom_auth_run_common_checks": True}, + ), + ): + await _run_post_custom_auth_checks( + valid_token=valid_token, + request=request, + request_data={}, + route="/v1/files", + parent_otel_span=None, + ) + mock_can_key.assert_awaited_once_with( + model="restricted-model", + llm_model_list=ANY, + valid_token=valid_token, + llm_router=ANY, + ) + + @pytest.mark.asyncio async def test_custom_auth_honors_key_level_model_access_restriction_denied_with_opt_in(): valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"]) @@ -1752,7 +1791,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 +1876,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 # --------------------------------------------------------------------------- From 3b54012b7be39e9981008a15ec27288a7d258c99 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:04:14 -0700 Subject: [PATCH 45/56] chore(proxy): satisfy auth model checks CI --- litellm/proxy/_lazy_openapi_snapshot.json | 34 +++++++++++------------ litellm/proxy/auth/auth_utils.py | 19 ++++++------- 2 files changed, 26 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8331f748c6e..43c18922e1c 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__patch", "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__patch", "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__patch", "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__patch", "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__patch", "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__patch", "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__patch", "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__patch", "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__patch", "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__patch", "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_patch", "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_patch", "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_patch", "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_patch", "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_patch", "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_patch", "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_patch", "parameters": [ { "in": "path", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index ba858a89fbc..97870cfcf04 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -988,16 +988,15 @@ def _append_model_candidates(candidates: List[str], value: Any) -> None: if value is None: return - if isinstance(value, str): - model_names = [model.strip() for model in value.split(",")] - elif isinstance(value, (list, tuple, set)): - for item in value: - _append_model_candidates(candidates=candidates, value=item) - return - else: - model_names = [str(value).strip()] - - candidates.extend(model for model in model_names if model) + values = value if isinstance(value, (list, tuple, set)) else [value] + for item in values: + if item is None: + continue + if isinstance(item, str): + model_names = [model.strip() for model in item.split(",")] + else: + model_names = [str(item).strip()] + candidates.extend(model for model in model_names if model) def _dedupe_model_candidates(candidates: List[str]) -> List[str]: From a5135b1b556cbe1c71160759542fab56643309ab Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:14:47 -0700 Subject: [PATCH 46/56] chore(proxy): stabilize lazy openapi snapshot --- litellm/proxy/_lazy_openapi_snapshot.json | 90 +++++++++++------------ litellm/proxy/_lazy_openapi_snapshot.py | 6 +- litellm/proxy/proxy_server.py | 83 +++++++++++++++++++++ 3 files changed, 132 insertions(+), 47 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 43c18922e1c..e8b2d701691 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__patch", + "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__patch", + "operationId": "anthropic_proxy_route_anthropic__endpoint__get", "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__patch", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", "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__patch", + "operationId": "anthropic_proxy_route_anthropic__endpoint__put", "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__patch", + "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__patch", + "operationId": "langfuse_proxy_route_langfuse__endpoint__get", "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__patch", + "operationId": "langfuse_proxy_route_langfuse__endpoint__post", "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__patch", + "operationId": "langfuse_proxy_route_langfuse__endpoint__put", "parameters": [ { "in": "path", @@ -14008,7 +14008,7 @@ "/mcp-rest/test/connection": { "post": { "description": "Test if we can connect to the provided MCP server before adding it", - "operationId": "test_connection_mcp_rest_test_connection_post", + "operationId": "test_connection_mcp_rest_test_connection_post_2", "requestBody": { "content": { "application/json": { @@ -14053,7 +14053,7 @@ "/mcp-rest/test/tools/list": { "post": { "description": "Preview tools available from MCP server before adding it", - "operationId": "test_tools_list_mcp_rest_test_tools_list_post", + "operationId": "test_tools_list_mcp_rest_test_tools_list_post_2", "requestBody": { "content": { "application/json": { @@ -14098,7 +14098,7 @@ "/mcp-rest/tools/call": { "post": { "description": "REST API to call a specific MCP tool with the provided arguments", - "operationId": "call_tool_rest_api_mcp_rest_tools_call_post", + "operationId": "call_tool_rest_api_mcp_rest_tools_call_post_2", "responses": { "200": { "content": { @@ -14123,7 +14123,7 @@ "/mcp-rest/tools/list": { "get": { "description": "List all available tools with information about the server they belong to.\n\nExample response:\n{\n \"tools\": [\n {\n \"name\": \"create_zap\",\n \"description\": \"Create a new zap\",\n \"inputSchema\": \"tool_input_schema\",\n \"mcp_info\": {\n \"server_name\": \"zapier\",\n \"logo_url\": \"https://www.zapier.com/logo.png\",\n }\n }\n ],\n \"error\": null,\n \"message\": \"Successfully retrieved tools\"\n}", - "operationId": "list_tool_rest_api_mcp_rest_tools_list_get", + "operationId": "list_tool_rest_api_mcp_rest_tools_list_get_2", "parameters": [ { "description": "The server id to list tools for", @@ -21896,7 +21896,7 @@ "/policies/usage/overview": { "get": { "description": "Return policy performance overview for the dashboard.", - "operationId": "policies_usage_overview_policies_usage_overview_get", + "operationId": "policies_usage_overview_policies_usage_overview_get_2", "parameters": [ { "description": "YYYY-MM-DD", @@ -22521,7 +22521,7 @@ "/policies/attachments/estimate-impact": { "post": { "description": "Estimate how many keys and teams would be affected by a policy attachment.\n\nUse this before creating an attachment to preview the blast radius.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/attachments/estimate-impact\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"policy_name\": \"hipaa-compliance\",\n \"tags\": [\"healthcare\", \"health-*\"]\n }'\n```", - "operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post", + "operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post_2", "requestBody": { "content": { "application/json": { @@ -22568,7 +22568,7 @@ "/policies/resolve": { "post": { "description": "Resolve which policies and guardrails apply for a given context.\n\nUse this endpoint to debug \"what guardrails would apply to a request\nwith this team/key/model/tags combination?\"\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/resolve\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"tags\": [\"healthcare\"],\n \"model\": \"gpt-4\"\n }'\n```", - "operationId": "resolve_policies_for_context_policies_resolve_post", + "operationId": "resolve_policies_for_context_policies_resolve_post_2", "parameters": [ { "description": "Force a DB sync before resolving. Default uses in-memory cache.", @@ -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_patch", + "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_patch", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_get", "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_patch", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_head", "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_patch", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", "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_patch", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_post", "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_patch", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", "parameters": [ { "in": "path", @@ -28329,7 +28329,7 @@ "/v1/vector_stores": { "get": { "description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list", - "operationId": "vector_store_list_v1_vector_stores_get", + "operationId": "vector_store_list_v1_vector_stores_get_2", "parameters": [ { "in": "query", @@ -28430,7 +28430,7 @@ }, "post": { "description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```", - "operationId": "vector_store_create_v1_vector_stores_post", + "operationId": "vector_store_create_v1_vector_stores_post_2", "responses": { "200": { "content": { @@ -28455,7 +28455,7 @@ "/v1/vector_stores/{vector_store_id}": { "delete": { "description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete", - "operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete", + "operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete_2", "parameters": [ { "in": "path", @@ -28499,7 +28499,7 @@ }, "get": { "description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve", - "operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get", + "operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get_2", "parameters": [ { "in": "path", @@ -28543,7 +28543,7 @@ }, "post": { "description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify", - "operationId": "vector_store_update_v1_vector_stores__vector_store_id__post", + "operationId": "vector_store_update_v1_vector_stores__vector_store_id__post_2", "parameters": [ { "in": "path", @@ -28588,7 +28588,7 @@ }, "/v1/vector_stores/{vector_store_id}/files": { "get": { - "operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get", + "operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get_2", "parameters": [ { "in": "path", @@ -28631,7 +28631,7 @@ ] }, "post": { - "operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post", + "operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post_2", "parameters": [ { "in": "path", @@ -28676,7 +28676,7 @@ }, "/v1/vector_stores/{vector_store_id}/files/{file_id}": { "delete": { - "operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete", + "operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete_2", "parameters": [ { "in": "path", @@ -28728,7 +28728,7 @@ ] }, "get": { - "operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get", + "operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get_2", "parameters": [ { "in": "path", @@ -28780,7 +28780,7 @@ ] }, "post": { - "operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post", + "operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post_2", "parameters": [ { "in": "path", @@ -28834,7 +28834,7 @@ }, "/v1/vector_stores/{vector_store_id}/files/{file_id}/content": { "get": { - "operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get", + "operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get_2", "parameters": [ { "in": "path", @@ -28889,7 +28889,7 @@ "/v1/vector_stores/{vector_store_id}/search": { "post": { "description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search", - "operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post", + "operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post_2", "parameters": [ { "in": "path", @@ -28935,7 +28935,7 @@ "/vector_stores": { "get": { "description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list", - "operationId": "vector_store_list_vector_stores_get", + "operationId": "vector_store_list_vector_stores_get_2", "parameters": [ { "in": "query", @@ -29036,7 +29036,7 @@ }, "post": { "description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```", - "operationId": "vector_store_create_vector_stores_post", + "operationId": "vector_store_create_vector_stores_post_2", "responses": { "200": { "content": { @@ -29061,7 +29061,7 @@ "/vector_stores/{vector_store_id}": { "delete": { "description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete", - "operationId": "vector_store_delete_vector_stores__vector_store_id__delete", + "operationId": "vector_store_delete_vector_stores__vector_store_id__delete_2", "parameters": [ { "in": "path", @@ -29105,7 +29105,7 @@ }, "get": { "description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve", - "operationId": "vector_store_retrieve_vector_stores__vector_store_id__get", + "operationId": "vector_store_retrieve_vector_stores__vector_store_id__get_2", "parameters": [ { "in": "path", @@ -29149,7 +29149,7 @@ }, "post": { "description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify", - "operationId": "vector_store_update_vector_stores__vector_store_id__post", + "operationId": "vector_store_update_vector_stores__vector_store_id__post_2", "parameters": [ { "in": "path", @@ -29194,7 +29194,7 @@ }, "/vector_stores/{vector_store_id}/files": { "get": { - "operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get", + "operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get_2", "parameters": [ { "in": "path", @@ -29237,7 +29237,7 @@ ] }, "post": { - "operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post", + "operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post_2", "parameters": [ { "in": "path", @@ -29282,7 +29282,7 @@ }, "/vector_stores/{vector_store_id}/files/{file_id}": { "delete": { - "operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete", + "operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete_2", "parameters": [ { "in": "path", @@ -29334,7 +29334,7 @@ ] }, "get": { - "operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get", + "operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get_2", "parameters": [ { "in": "path", @@ -29386,7 +29386,7 @@ ] }, "post": { - "operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post", + "operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post_2", "parameters": [ { "in": "path", @@ -29440,7 +29440,7 @@ }, "/vector_stores/{vector_store_id}/files/{file_id}/content": { "get": { - "operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get", + "operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get_2", "parameters": [ { "in": "path", @@ -29495,7 +29495,7 @@ "/vector_stores/{vector_store_id}/search": { "post": { "description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search", - "operationId": "vector_store_search_vector_stores__vector_store_id__search_post", + "operationId": "vector_store_search_vector_stores__vector_store_id__search_post_2", "parameters": [ { "in": "path", diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 315f6a9742a..51cbd6eb989 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -10,7 +10,7 @@ any drift as a neutral check. import json import sys from pathlib import Path -from typing import Dict, Optional +from typing import Dict, Optional, Set SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" @@ -31,7 +31,7 @@ def generate_snapshot() -> Dict[str, Dict]: from fastapi.openapi.utils import get_openapi from litellm.proxy._lazy_features import LAZY_FEATURES - from litellm.proxy.proxy_server import app + from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids for feat in LAZY_FEATURES: if feat.module_path in sys.modules: @@ -43,6 +43,7 @@ def generate_snapshot() -> Dict[str, Dict]: sys.stderr.write(f"warning: skip {feat.name}: {exc}\n") fragments: Dict[str, Dict] = {} + used_operation_ids: Set[str] = set() for feat in LAZY_FEATURES: feat_routes = [ r @@ -57,6 +58,7 @@ def generate_snapshot() -> Dict[str, Dict]: for op in path_ops.values(): if isinstance(op, dict): op["tags"] = [feat.name] + full = ensure_unique_openapi_operation_ids(full, used_operation_ids) fragments[feat.name] = { "paths": full.get("paths", {}), "components": {"schemas": full.get("components", {}).get("schemas", {})}, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6cba6a3e96b..ed8365f95ce 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6,6 +6,7 @@ import inspect import io import os import random +import re import secrets import shutil import subprocess @@ -950,6 +951,85 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues] +def _generate_stable_operation_id(route: Any) -> str: + operation_id = re.sub(r"\W", "_", f"{route.name}{route.path_format}") + route_methods = sorted(route.methods or []) + if len(route_methods) == 1: + operation_id = f"{operation_id}_{route_methods[0].lower()}" + return operation_id + + +_OPENAPI_HTTP_METHODS = { + "delete", + "get", + "head", + "options", + "patch", + "post", + "put", + "trace", +} + + +def _strip_operation_id_method_suffix(operation_id: str) -> str: + base, separator, suffix = operation_id.rpartition("_") + if separator and suffix in _OPENAPI_HTTP_METHODS: + return base + return operation_id + + +def ensure_unique_openapi_operation_ids( + openapi_schema: Dict[str, Any], + reserved_operation_ids: Optional[Set[str]] = None, +) -> Dict[str, Any]: + operation_entries = [] + operation_id_counts: Dict[str, int] = {} + for path_item in openapi_schema.get("paths", {}).values(): + if not isinstance(path_item, dict): + continue + for method, operation in path_item.items(): + if method not in _OPENAPI_HTTP_METHODS or not isinstance(operation, dict): + continue + operation_id = operation.get("operationId") + if not isinstance(operation_id, str): + continue + operation_entries.append((method, operation, operation_id)) + operation_id_counts[operation_id] = ( + operation_id_counts.get(operation_id, 0) + 1 + ) + + used_operation_ids = set(reserved_operation_ids or set()) + seen_operation_ids: Set[str] = set() + for method, operation, operation_id in operation_entries: + should_rewrite = ( + operation_id_counts[operation_id] > 1 + or operation_id in used_operation_ids + or operation_id in seen_operation_ids + ) + if not should_rewrite: + seen_operation_ids.add(operation_id) + used_operation_ids.add(operation_id) + continue + + base_operation_id = _strip_operation_id_method_suffix(operation_id) + new_operation_id = f"{base_operation_id}_{method}" + suffix = 2 + while ( + new_operation_id in used_operation_ids + or new_operation_id in seen_operation_ids + ): + new_operation_id = f"{base_operation_id}_{method}_{suffix}" + suffix += 1 + operation["operationId"] = new_operation_id + seen_operation_ids.add(new_operation_id) + used_operation_ids.add(new_operation_id) + + if reserved_operation_ids is not None: + reserved_operation_ids.update(used_operation_ids) + + return openapi_schema + + app = FastAPI( docs_url=_get_docs_url(), redoc_url=_get_redoc_url(), @@ -959,6 +1039,7 @@ app = FastAPI( version=version, root_path=server_root_path, lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues] + generate_unique_id_function=_generate_stable_operation_id, ) vertex_live_passthrough_vertex_base = VertexBase() @@ -1038,6 +1119,7 @@ def get_openapi_schema(): from litellm.proxy._lazy_features import inject_lazy_stubs openapi_schema = inject_lazy_stubs(openapi_schema) + openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema) # Fix Swagger UI execute path error when server_root_path is set if server_root_path: @@ -1069,6 +1151,7 @@ def custom_openapi(): from litellm.proxy._lazy_features import inject_lazy_stubs openapi_schema = inject_lazy_stubs(openapi_schema) + openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema) # Fix Swagger UI execute path error when server_root_path is set if server_root_path: From 0704f672c55cf72f8ae377172b58a0b8cffea61f Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:21:57 -0700 Subject: [PATCH 47/56] test(proxy): cover resource model extraction fallbacks --- .../proxy/auth/test_auth_utils.py | 41 ++++++++++++++++++- 1 file changed, 40 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index b3b6fdde670..9fb33099fd8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -11,11 +11,12 @@ import pytest from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( _get_customer_id_from_standard_headers, + abbreviate_api_key, check_complete_credentials, get_end_user_id_from_request_body, - get_model_from_request, get_key_model_rpm_limit, get_key_model_tpm_limit, + get_model_from_request, get_project_model_rpm_limit, get_project_model_tpm_limit, is_request_body_safe, @@ -259,6 +260,16 @@ def test_get_model_from_request_vertex_passthrough_still_works(): assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro" +def test_get_model_from_request_openai_deployment_route_still_works(): + assert ( + get_model_from_request( + request_data={}, + route="/openai/deployments/my-azure-deployment/chat/completions", + ) + == "my-azure-deployment" + ) + + def test_get_model_from_request_includes_file_endpoint_header_model(): assert ( get_model_from_request( @@ -370,6 +381,34 @@ def test_get_model_from_request_extracts_video_id_model(): ) +def test_get_model_from_request_handles_managed_id_decoder_failures(): + with ( + patch( + "litellm.proxy.openai_files_endpoints.common_utils.decode_model_from_file_id", + side_effect=Exception("decode failed"), + ), + patch( + "litellm.llms.base_llm.managed_resources.utils.parse_unified_id", + side_effect=Exception("parse failed"), + ), + patch( + "litellm.types.videos.utils.decode_video_id_with_provider", + side_effect=Exception("video decode failed"), + ), + ): + assert ( + get_model_from_request( + request_data={"file_id": "not-a-managed-resource-id"}, + route="/v1/files/{file_id}", + ) + is None + ) + + +def test_abbreviate_api_key(): + assert abbreviate_api_key("sk-test-1234") == "sk-...1234" + + def test_get_customer_user_header_returns_none_when_no_customer_role(): from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping From 6ef26945fa2434c8656a476eff8e85f146fcaa80 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:55:00 -0700 Subject: [PATCH 48/56] test(proxy): narrow media resource decoding --- litellm/proxy/auth/auth_utils.py | 47 ++++++++++------- .../proxy/auth/test_auth_utils.py | 51 +++++++++++++++++++ 2 files changed, 79 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 97870cfcf04..4b72d813ee9 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1030,7 +1030,9 @@ def _route_uses_model_routing_sources(route: str) -> bool: return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS) -def _extract_models_from_managed_resource_id(resource_id: Any) -> List[str]: +def _extract_models_from_managed_resource_id( + resource_id: Any, resource_id_field: Optional[str] = None +) -> List[str]: if not isinstance(resource_id, str) or not resource_id: return [] @@ -1078,24 +1080,29 @@ def _extract_models_from_managed_resource_id(resource_id: Any) -> List[str]: "Unable to extract model from unified managed resource ID: %s", str(e) ) - try: - from litellm.types.videos.utils import ( - decode_character_id_with_provider, - decode_video_id_with_provider, - ) + if resource_id_field in ("video_id", "character_id"): + try: + from litellm.types.videos.utils import ( + decode_character_id_with_provider, + decode_video_id_with_provider, + ) - _append_model_candidates( - candidates=candidates, - value=decode_video_id_with_provider(resource_id).get("model_id"), - ) - _append_model_candidates( - candidates=candidates, - value=decode_character_id_with_provider(resource_id).get("model_id"), - ) - except Exception as e: - verbose_proxy_logger.debug( - "Unable to extract model from managed video/character ID: %s", str(e) - ) + if resource_id_field == "video_id": + _append_model_candidates( + candidates=candidates, + value=decode_video_id_with_provider(resource_id).get("model_id"), + ) + else: + _append_model_candidates( + candidates=candidates, + value=decode_character_id_with_provider(resource_id).get( + "model_id" + ), + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from managed video/character ID: %s", str(e) + ) return _dedupe_model_candidates(candidates) @@ -1153,7 +1160,9 @@ def _extract_model_candidates_from_request( for field in _MODEL_ROUTING_ID_FIELDS: _append_model_candidates( candidates, - _extract_models_from_managed_resource_id(request_data.get(field)), + _extract_models_from_managed_resource_id( + request_data.get(field), resource_id_field=field + ), ) return _dedupe_model_candidates(candidates) diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 9fb33099fd8..cf02d6f95df 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -381,6 +381,50 @@ def test_get_model_from_request_extracts_video_id_model(): ) +def test_get_model_from_request_only_runs_media_decoders_for_matching_fields(): + with ( + patch( + "litellm.types.videos.utils.decode_video_id_with_provider", + return_value={"model_id": "video-model"}, + ) as video_decoder, + patch( + "litellm.types.videos.utils.decode_character_id_with_provider", + return_value={"model_id": "character-model"}, + ) as character_decoder, + ): + assert ( + get_model_from_request( + request_data={"file_id": "file-provider-id"}, + route="/v1/files/{file_id}", + ) + is None + ) + video_decoder.assert_not_called() + character_decoder.assert_not_called() + + assert ( + get_model_from_request( + request_data={"video_id": "video-provider-id"}, + route="/v1/videos/{video_id}", + ) + == "video-model" + ) + video_decoder.assert_called_once_with("video-provider-id") + character_decoder.assert_not_called() + + video_decoder.reset_mock() + character_decoder.reset_mock() + assert ( + get_model_from_request( + request_data={"character_id": "character-provider-id"}, + route="/v1/videos/{character_id}", + ) + == "character-model" + ) + video_decoder.assert_not_called() + character_decoder.assert_called_once_with("character-provider-id") + + def test_get_model_from_request_handles_managed_id_decoder_failures(): with ( patch( @@ -403,6 +447,13 @@ def test_get_model_from_request_handles_managed_id_decoder_failures(): ) is None ) + assert ( + get_model_from_request( + request_data={"video_id": "not-a-managed-resource-id"}, + route="/v1/videos/{video_id}", + ) + is None + ) def test_abbreviate_api_key(): 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 49/56] 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 f51dd68ff0607e1a95c1ceef876a7176add31536 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 23:00:25 -0700 Subject: [PATCH 50/56] test(proxy): cover lazy openapi operation ids --- .../proxy/test_lazy_openapi_snapshot.py | 82 +++++++++++++++++++ 1 file changed, 82 insertions(+) create mode 100644 tests/test_litellm/proxy/test_lazy_openapi_snapshot.py diff --git a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py new file mode 100644 index 00000000000..3062cf94c76 --- /dev/null +++ b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py @@ -0,0 +1,82 @@ +import sys +from types import ModuleType, SimpleNamespace + + +def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch): + from litellm.proxy import _lazy_openapi_snapshot + + route_a = SimpleNamespace(path="/feature-a/items") + route_b = SimpleNamespace(path="/feature-b/items") + fake_app = SimpleNamespace( + title="LiteLLM test", + version="0.0.0", + routes=[route_a, route_b], + ) + + fake_feature_a_module = ModuleType("fake_feature_a") + fake_feature_b_module = ModuleType("fake_feature_b") + monkeypatch.setitem(sys.modules, "fake_feature_a", fake_feature_a_module) + monkeypatch.setitem(sys.modules, "fake_feature_b", fake_feature_b_module) + + fake_lazy_features_module = ModuleType("litellm.proxy._lazy_features") + fake_lazy_features_module.LAZY_FEATURES = [ + SimpleNamespace( + name="feature-a", + module_path="fake_feature_a", + path_prefixes=("/feature-a",), + register_fn=lambda app, module: None, + ), + SimpleNamespace( + name="feature-b", + module_path="fake_feature_b", + path_prefixes=("/feature-b",), + register_fn=lambda app, module: None, + ), + ] + monkeypatch.setitem( + sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module + ) + + def fake_get_openapi(title, version, routes): + path = routes[0].path + return { + "paths": {path: {"get": {"operationId": "shared_operation_id_get"}}}, + "components": {"schemas": {"Example": {"type": "object"}}}, + } + + def fake_ensure_unique_openapi_operation_ids(schema, reserved_operation_ids): + for path_item in schema["paths"].values(): + operation = path_item["get"] + operation_id = operation["operationId"] + if operation_id in reserved_operation_ids: + operation_id = f"{operation_id}_2" + operation["operationId"] = operation_id + reserved_operation_ids.add(operation_id) + return schema + + fake_proxy_server_module = ModuleType("litellm.proxy.proxy_server") + fake_proxy_server_module.app = fake_app + fake_proxy_server_module.ensure_unique_openapi_operation_ids = ( + fake_ensure_unique_openapi_operation_ids + ) + monkeypatch.setitem( + sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module + ) + monkeypatch.setattr("fastapi.openapi.utils.get_openapi", fake_get_openapi) + + fragments = _lazy_openapi_snapshot.generate_snapshot() + + assert ( + fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["operationId"] + == "shared_operation_id_get" + ) + assert ( + fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["operationId"] + == "shared_operation_id_get_2" + ) + assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["tags"] == [ + "feature-a" + ] + assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == [ + "feature-b" + ] 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 51/56] 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 52/56] 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, From 4825d94a9dd09e36bdda16f2bc9490fba8e07fbb Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 1 May 2026 14:24:53 -0700 Subject: [PATCH 53/56] [Fix] Tests: Move Misplaced Import in Lazy OpenAPI Snapshot Test The GitHub merge conflict resolver concatenated both test sets but left `from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids` stranded between functions instead of at the top of the file. --- tests/test_litellm/proxy/test_lazy_openapi_snapshot.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py index 4fa6cb85123..64cb931888b 100644 --- a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py +++ b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py @@ -1,6 +1,8 @@ import sys from types import ModuleType, SimpleNamespace +from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids + def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch): from litellm.proxy import _lazy_openapi_snapshot @@ -80,7 +82,6 @@ def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch): assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == [ "feature-b" ] -from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids def test_normalize_operation_ids_uses_each_http_method(): From d8f556e18cefbf15087fde80de0df03dd84cc829 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 1 May 2026 14:39:24 -0700 Subject: [PATCH 54/56] [Fix] Auth: Restore Request-Context Model Resolution In Skip-Budget Helper Merge of #26845 kept the PR's _should_skip_budget_checks helper but lost staging's upgrade to _get_model_from_request_context, so zero-cost models resolved from request headers/query params no longer skipped budget checks. Route the helper through _get_model_from_request_context so this path matches the other 8 model-resolution sites in the file. --- litellm/proxy/auth/user_api_key_auth.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index d7cfb8962a9..2e5140e0e34 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1907,6 +1907,7 @@ async def _run_centralized_common_checks( skip_budget_checks = _should_skip_budget_checks( request_data=request_data, route=route, + request=request, llm_router=llm_router, ) @@ -1988,9 +1989,14 @@ async def _reserve_budget_after_common_checks( def _should_skip_budget_checks( request_data: dict, route: str, + request: Optional[Request], llm_router: Optional[Any], ) -> bool: - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) if model is not None and llm_router is not None: return _is_model_cost_zero(model=model, llm_router=llm_router) return False From 61a3923b718ac86927006197c7babfde76ce72f6 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 1 May 2026 15:11:56 -0700 Subject: [PATCH 55/56] [Fix] Proxy: Repair Stale HTTP_METHODS Reference In Lazy OpenAPI Snapshot _normalize_operation_ids referenced HTTP_METHODS but only HTTP_METHOD_SUFFIXES is defined, raising NameError on snapshot generation and failing test_lazy_openapi_snapshot. The constant was renamed in an earlier merge without updating these two references; values are identical sets of HTTP method names. --- litellm/proxy/_lazy_openapi_snapshot.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index de118a66d98..c63ff8d0733 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -61,12 +61,12 @@ def _normalize_operation_ids(paths: Dict[str, Dict]) -> None: if not isinstance(path_ops, dict): continue - methods = {method for method in path_ops if method in HTTP_METHODS} + methods = {method for method in path_ops if method in HTTP_METHOD_SUFFIXES} if not methods: continue for method, operation in path_ops.items(): - if method not in HTTP_METHODS or not isinstance(operation, dict): + if method not in HTTP_METHOD_SUFFIXES or not isinstance(operation, dict): continue operation_id = operation.get("operationId") From a12b4249bd65b4712682cd73bf96055e3389f7a2 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 1 May 2026 15:56:07 -0700 Subject: [PATCH 56/56] [Fix] Proxy: Skip Personal Budget Hook When Reservation Covers Counter The reservation path (PR #26845) atomically pre-fills `spend:user:{user_id}` and admits at the strict-`<` boundary. The legacy `_PROXY_MaxBudgetLimiter` pre-call hook re-reads the same counter with `>=`, so a reservation that fills the counter to exactly `max_budget` (e.g. a request without a `max_tokens` cap that falls back to reserving the smallest remaining headroom) is rejected by the hook even though the reservation already admitted it. Skip the hook when the request's active `budget_reservation` covers `spend:user:{user_id}`. The reservation is the source of truth for that counter cross-pod; the legacy `>=` path remains in place for requests without a reservation (e.g. paths that bypass the reservation entirely). Reproduces as `tests/otel_tests/test_prometheus.py::test_user_budget_metrics` on a fresh user with `max_budget=10` calling `fake-openai-endpoint` without `max_tokens`. Adds focused unit coverage in `tests/test_litellm/proxy/hooks/test_max_budget_limiter.py`. --- litellm/proxy/hooks/max_budget_limiter.py | 17 +- .../proxy/hooks/test_max_budget_limiter.py | 208 ++++++++++++++++++ 2 files changed, 224 insertions(+), 1 deletion(-) create mode 100644 tests/test_litellm/proxy/hooks/test_max_budget_limiter.py diff --git a/litellm/proxy/hooks/max_budget_limiter.py b/litellm/proxy/hooks/max_budget_limiter.py index 7789fa6a349..9a7e5117945 100644 --- a/litellm/proxy/hooks/max_budget_limiter.py +++ b/litellm/proxy/hooks/max_budget_limiter.py @@ -32,10 +32,25 @@ class _PROXY_MaxBudgetLimiter(CustomLogger): if user_api_key_dict.team_id is not None: return + # The reservation path admits at the strict-`<` boundary and + # atomically pre-fills the same counter we'd read here. Re-checking + # with `>=` would reject a request the reservation already admitted + # when the reservation fills the counter to exactly max_budget. + # Imported lazily to avoid a circular import via proxy.utils. + from litellm.proxy.spend_tracking.budget_reservation import ( + get_reserved_counter_keys, + ) + + user_counter_key = f"spend:user:{user_id}" + if user_counter_key in get_reserved_counter_keys( + user_api_key_dict.budget_reservation + ): + return + from litellm.proxy.proxy_server import get_current_spend curr_spend = await get_current_spend( - counter_key=f"spend:user:{user_id}", + counter_key=user_counter_key, fallback_spend=user_api_key_dict.user_spend or 0.0, ) diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_limiter.py new file mode 100644 index 00000000000..0074d7062b8 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_max_budget_limiter.py @@ -0,0 +1,208 @@ +""" +Unit tests for the personal-budget pre-call hook. + +The reservation path (added in PR #26845) atomically pre-fills the same +`spend:user:{user_id}` counter this hook reads, admitting at a strict-`<` +boundary. Re-checking with `>=` after reservation would reject requests the +reservation already admitted when the reservation fills the counter to +exactly `max_budget` (e.g. requests with no `max_tokens` cap fall back to +reserving the smallest remaining headroom). + +These tests pin the skip-when-reserved behavior and guard against drift. +""" + +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import HTTPException + +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter + + +def _make_user_api_key_auth( + user_id: str = "user-1", + user_max_budget: float = 10.0, + user_spend: float = 0.0, + team_id=None, + budget_reservation=None, +) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + user_id=user_id, + user_max_budget=user_max_budget, + user_spend=user_spend, + team_id=team_id, + budget_reservation=budget_reservation, + ) + + +@pytest.mark.asyncio +async def test_under_budget_passes(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=3.0), + ): + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + + +@pytest.mark.asyncio +async def test_over_budget_rejects_without_reservation(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=10.0), + ): + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert exc_info.value.status_code == 429 + assert "Max budget limit reached." in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_skips_when_user_counter_is_reserved(): + """ + Reservation atomically pre-fills `spend:user:{user_id}` and admits the + request. The legacy `>=` check must not double-enforce on the same + counter — that's what produced the boundary regression where a fresh + user with no `max_tokens` cap got 429'd on their first request. + """ + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth( + user_id="user-1", + user_max_budget=10.0, + budget_reservation={ + "reserved_cost": 10.0, + "entries": [ + { + "counter_key": "spend:user:user-1", + "entity_type": "User", + "entity_id": "user-1", + "reserved_cost": 10.0, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + }, + ) + + # `get_current_spend` would return 10.0 here (counter pre-filled by the + # reservation). The hook must skip without reading it. + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=10.0), + ) as mock_get_spend: + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + mock_get_spend.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_does_not_skip_when_reservation_covers_a_different_counter(): + """ + A reservation that only covers e.g. `spend:team:{team_id}` (not the user + counter) must not exempt the user-budget check. + """ + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth( + user_id="user-1", + user_max_budget=10.0, + budget_reservation={ + "reserved_cost": 5.0, + "entries": [ + { + "counter_key": "spend:team:team-x", + "entity_type": "Team", + "entity_id": "team-x", + "reserved_cost": 5.0, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + }, + ) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=10.0), + ): + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert exc_info.value.status_code == 429 + + +@pytest.mark.asyncio +async def test_team_keys_skip_personal_budget(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth( + user_max_budget=10.0, + team_id="team-1", + ) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=999.0), + ) as mock_get_spend: + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + mock_get_spend.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_no_max_budget_passes(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", + user_id="user-1", + ) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=999.0), + ) as mock_get_spend: + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + mock_get_spend.assert_not_awaited()