mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
tighten budget spend admission
This commit is contained in:
parent
d7431c9db9
commit
5a619cf879
8 changed files with 1678 additions and 105 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
673
litellm/proxy/spend_tracking/budget_reservation.py
Normal file
673
litellm/proxy/spend_tracking/budget_reservation.py
Normal 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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
509
tests/test_litellm/proxy/test_budget_reservation.py
Normal file
509
tests/test_litellm/proxy/test_budget_reservation.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue