tighten budget spend admission

This commit is contained in:
user 2026-04-29 19:13:55 -07:00
parent d7431c9db9
commit 5a619cf879
8 changed files with 1678 additions and 105 deletions

View file

@ -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

View file

@ -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,

View file

@ -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,
)

View file

@ -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],

View file

@ -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)

View file

@ -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,

View file

@ -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()

View file

@ -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)