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)