diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8e57327b31b..3498bac0190 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -279,6 +279,8 @@ class KeyManagementRoutes(str, enum.Enum): # info and health routes KEY_INFO = "/key/info" + KEY_BUDGETS = "/key/{key_id}/budgets" + KEY_BUDGETS_SELF = "/key/budgets" KEY_HEALTH = "/key/health" # list routes @@ -568,6 +570,8 @@ class LiteLLMRoutes(enum.Enum): ) info_routes = [ "/key/info", + KeyManagementRoutes.KEY_BUDGETS.value, + KeyManagementRoutes.KEY_BUDGETS_SELF.value, "/key/health", "/team/info", "/team/list", @@ -604,6 +608,8 @@ class LiteLLMRoutes(enum.Enum): KeyManagementRoutes.KEY_UPDATE.value, KeyManagementRoutes.KEY_DELETE.value, KeyManagementRoutes.KEY_INFO.value, + KeyManagementRoutes.KEY_BUDGETS.value, + KeyManagementRoutes.KEY_BUDGETS_SELF.value, KeyManagementRoutes.KEY_REGENERATE.value, KeyManagementRoutes.KEY_GENERATE_SERVICE_ACCOUNT.value, KeyManagementRoutes.KEY_REGENERATE_WITH_PATH_PARAM.value, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8708f96339f..5736d2462a4 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -14,6 +14,7 @@ import math import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence +from dataclasses import dataclass from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast @@ -94,6 +95,17 @@ from litellm.proxy.guardrails.tool_name_extraction import ( ) from litellm.proxy.route_llm_request import route_request from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start +from litellm.proxy.spend_tracking.spend_counter_keys import ( + end_user_spend_counter, + key_spend_counter, + key_window_spend_counter, + org_spend_counter, + tag_spend_counter, + team_member_spend_counter, + team_spend_counter, + team_window_spend_counter, + user_spend_counter, +) from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.object_permission_repository import ObjectPermissionRepository @@ -684,6 +696,12 @@ BUDGET_ENFORCED_SIDE_EFFECT_ROUTES: Final = frozenset( ) +def user_budget_applies_to_key(team_object: LiteLLM_TeamTable | None, general_settings: Mapping[str, object]) -> bool: + """A team key ignores its owner's personal budget unless the operator opted in.""" + is_team_key: Final = team_object is not None and team_object.team_id is not None + return not is_team_key or general_settings.get("apply_user_budget_to_team_keys") is True + + async def common_checks( request_body: dict, team_object: LiteLLM_TeamTable | None, @@ -835,15 +853,14 @@ async def common_checks( # 4.1 personal budget if user_object is None or user_object.max_budget is None: return - is_team_key: Final = team_object is not None and team_object.team_id is not None - if is_team_key and general_settings.get("apply_user_budget_to_team_keys") is not True: + if not user_budget_applies_to_key(team_object=team_object, general_settings=general_settings): return from litellm.proxy.proxy_server import get_current_spend user_budget: Final = user_object.max_budget user_spend: Final = await get_current_spend( - counter_key=f"spend:user:{user_object.user_id}", + counter_key=user_spend_counter(user_object.user_id), fallback_spend=user_object.spend or 0.0, max_budget=user_budget, ) @@ -1288,7 +1305,7 @@ async def _check_end_user_budget( from litellm.proxy.proxy_server import get_current_spend end_user_spend: Final = await get_current_spend( - counter_key=f"spend:end_user:{end_user_obj.user_id}", + counter_key=end_user_spend_counter(end_user_obj.user_id), fallback_spend=end_user_obj.spend or 0.0, max_budget=end_user_budget, fallback_authoritative=True, @@ -4124,7 +4141,7 @@ async def _virtual_key_max_budget_check( from litellm.proxy.proxy_server import get_current_spend fallback_spend: Final = valid_token.spend or 0.0 - counter_key: Final = f"spend:key:{valid_token.token}" + counter_key: Final = key_spend_counter(valid_token.token) # Read spend from cross-pod counter (Redis-first) or cached object (fallback) spend: Final = await get_current_spend( @@ -4205,7 +4222,7 @@ async def _virtual_key_multi_budget_check( for window in valid_token.budget_limits: w: dict = window if isinstance(window, dict) else window.model_dump() - counter_key = f"spend:key:{valid_token.token}:window:{w['budget_duration']}" + counter_key = key_window_spend_counter(valid_token.token, w["budget_duration"]) window_spend = await get_current_spend( counter_key=counter_key, fallback_spend=0.0, @@ -4389,6 +4406,70 @@ async def _virtual_key_max_budget_alert_check( ) +@dataclass(frozen=True, slots=True) +class TeamMemberBudget: + """The per-member cap enforced inside a team, plus the recorded spend it is measured against. + + Resolution is shared with budget introspection, so a change to the fallback order can never + make the two disagree about which cap a request is judged by. + """ + + max_budget: float | None + recorded_spend: float + source: str + + +async def resolve_team_member_budget( + team_object: LiteLLM_TeamTable, + user_id: str, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging | None = None, +) -> TeamMemberBudget: + team_membership: Final = await get_team_membership( + user_id=user_id, + team_id=team_object.team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + recorded_spend: Final = (team_membership.spend if team_membership is not None else 0.0) or 0.0 + + # Per-member override wins; otherwise fall back to the team-level + # default configured via team.metadata["team_member_budget_id"]. + if ( + team_membership is not None + and team_membership.litellm_budget_table is not None + and team_membership.litellm_budget_table.max_budget is not None + ): + return TeamMemberBudget( + max_budget=team_membership.litellm_budget_table.max_budget, + recorded_spend=recorded_spend, + source=f"budget_table:{team_membership.budget_id}", + ) + + metadata: Final = team_object.metadata + default_budget_id: Final = metadata.get("team_member_budget_id") if metadata else None + if not isinstance(default_budget_id, str): + return TeamMemberBudget(max_budget=None, recorded_spend=recorded_spend, source="team_membership.budget_id") + + default_budget: Final = await get_team_member_default_budget( + budget_id=default_budget_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + # Treat 0 on the team default as "no cap". + # Per-member rows still respect 0 as an explicit admin disable. + if default_budget is None or default_budget.max_budget is None or default_budget.max_budget <= 0: + return TeamMemberBudget(max_budget=None, recorded_spend=recorded_spend, source="team_membership.budget_id") + + return TeamMemberBudget( + max_budget=default_budget.max_budget, + recorded_spend=recorded_spend, + source=f"team.metadata.team_member_budget_id:{default_budget_id}", + ) + + async def _check_team_member_budget( team_object: LiteLLM_TeamTable | None, user_object: LiteLLM_UserTable | None, @@ -4398,67 +4479,37 @@ async def _check_team_member_budget( proxy_logging_obj: ProxyLogging, ): """Check if team member is over their max budget within the team.""" - if ( - team_object is not None - and team_object.team_id is not None - and valid_token is not None - and valid_token.user_id is not None - ): - team_membership: Final = await get_team_membership( - user_id=valid_token.user_id, - team_id=team_object.team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, + if team_object is None or team_object.team_id is None or valid_token is None or valid_token.user_id is None: + return + + member_budget: Final = await resolve_team_member_budget( + team_object=team_object, + user_id=valid_token.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if member_budget.max_budget is None: + return + + # Read from cross-pod counter (Redis-first) if available + from litellm.proxy.proxy_server import get_current_spend + + team_member_spend: Final = await get_current_spend( + counter_key=team_member_spend_counter(valid_token.user_id, team_object.team_id), + fallback_spend=member_budget.recorded_spend, + max_budget=member_budget.max_budget, + ) + + if math.isfinite(member_budget.max_budget) and team_member_spend >= member_budget.max_budget: + raise litellm.BudgetExceededError( + current_cost=team_member_spend, + max_budget=member_budget.max_budget, + message=f"Budget has been exceeded! User={valid_token.user_id} in Team={team_object.team_id} Current cost: {team_member_spend}, Max budget: {member_budget.max_budget}", + entity_type=Litellm_EntityType.TEAM_MEMBER.value, + entity_id=f"{valid_token.user_id}:{team_object.team_id}", ) - # Per-member override wins; otherwise fall back to the team-level - # default configured via team.metadata["team_member_budget_id"]. - team_member_budget: float | None = None - if ( - team_membership is not None - and team_membership.litellm_budget_table is not None - and team_membership.litellm_budget_table.max_budget is not None - ): - team_member_budget = team_membership.litellm_budget_table.max_budget - else: - default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id") - if isinstance(default_budget_id, str): - default_budget: Final = await get_team_member_default_budget( - budget_id=default_budget_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - ) - # Treat 0 on the team default as "no cap". - # Per-member rows still respect 0 as an explicit admin disable. - if ( - default_budget is not None - and default_budget.max_budget is not None - and default_budget.max_budget > 0 - ): - team_member_budget = default_budget.max_budget - - if team_member_budget is not None: - team_member_spend = (team_membership.spend if team_membership is not None else 0.0) or 0.0 - - # Read from cross-pod counter (Redis-first) if available - from litellm.proxy.proxy_server import get_current_spend - - team_member_spend = await get_current_spend( - counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}", - fallback_spend=team_member_spend, - max_budget=team_member_budget, - ) - - if math.isfinite(team_member_budget) and team_member_spend >= team_member_budget: - raise litellm.BudgetExceededError( - current_cost=team_member_spend, - max_budget=team_member_budget, - message=f"Budget has been exceeded! User={valid_token.user_id} in Team={team_object.team_id} Current cost: {team_member_spend}, Max budget: {team_member_budget}", - entity_type=Litellm_EntityType.TEAM_MEMBER.value, - entity_id=f"{valid_token.user_id}:{team_object.team_id}", - ) - async def _check_team_member_model_access( model: str | list[str], @@ -4528,7 +4579,7 @@ async def _team_max_budget_check( # Read spend from cross-pod counter (Redis-first) or cached object (fallback) spend: Final = await get_current_spend( - counter_key=f"spend:team:{team_object.team_id}", + counter_key=team_spend_counter(team_object.team_id), fallback_spend=team_object.spend or 0.0, max_budget=team_object.max_budget, ) @@ -4578,7 +4629,7 @@ async def _team_multi_budget_check( for window in team_object.budget_limits: w: dict = window if isinstance(window, dict) else window.model_dump() - counter_key = f"spend:team:{team_object.team_id}:window:{w['budget_duration']}" + counter_key = team_window_spend_counter(team_object.team_id, w["budget_duration"]) window_spend = await get_current_spend( counter_key=counter_key, fallback_spend=0.0, @@ -4835,6 +4886,15 @@ async def delete_cached_project_object( ) +def resolve_budget_org_id(valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None) -> str | None: + """The org whose budget gates this key: the key's own org, else the org its team belongs to.""" + if valid_token is not None and valid_token.org_id is not None: + return valid_token.org_id + if team_object is not None: + return team_object.organization_id + return None + + async def _organization_max_budget_check( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, @@ -4859,14 +4919,7 @@ async def _organization_max_budget_check( if valid_token is None or prisma_client is None: return - # Determine organization_id: first try from token, then fallback to team - org_id: str | None = 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 no organization_id found, skip the check + org_id: Final = resolve_budget_org_id(valid_token=valid_token, team_object=team_object) if org_id is None: return @@ -4899,7 +4952,7 @@ async def _organization_max_budget_check( from litellm.proxy.proxy_server import get_current_spend org_spend: Final = await get_current_spend( - counter_key=f"spend:org:{org_id}", + counter_key=org_spend_counter(org_id), fallback_spend=org_table.spend or 0.0, max_budget=org_max_budget, ) @@ -4976,7 +5029,7 @@ async def _tag_max_budget_check( from litellm.proxy.proxy_server import get_current_spend tag_spend = await get_current_spend( - counter_key=f"spend:tag:{tag_name}", + counter_key=tag_spend_counter(tag_name), fallback_spend=tag_object.spend or 0.0, max_budget=tag_object.litellm_budget_table.max_budget, fallback_authoritative=True, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 99592d44f9b..7866a226ca4 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1762,7 +1762,7 @@ async def _user_api_key_auth_builder( valid_token.allowed_model_region = end_user_params.get("allowed_model_region") if valid_token is not None: - valid_token = _update_key_budget_with_temp_budget_increase(valid_token) + valid_token = update_key_budget_with_temp_budget_increase(valid_token) user_obj: LiteLLM_UserTable | None = None valid_token_dict: dict = {} @@ -2813,7 +2813,7 @@ def _get_temp_budget_increase(valid_token: UserAPIKeyAuth): return None -def _update_key_budget_with_temp_budget_increase( +def update_key_budget_with_temp_budget_increase( valid_token: UserAPIKeyAuth, ) -> UserAPIKeyAuth: if valid_token.max_budget is None: diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 215969ef899..a77b313bd7d 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -62,7 +62,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): # check if current model is within budget if _current_model_budget_info.max_budget and _current_model_budget_info.max_budget > 0: - _current_spend: Final = await self._get_virtual_key_spend_for_model( + _current_spend: Final = await self.get_virtual_key_spend_for_model( user_api_key_hash=user_api_key_dict.token, model=model, key_budget_config=_current_model_budget_info, @@ -128,7 +128,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): # check if current model is within budget if _current_model_budget_info.max_budget and _current_model_budget_info.max_budget > 0: - _current_spend: Final = await self._get_end_user_spend_for_model( + _current_spend: Final = await self.get_end_user_spend_for_model( end_user_id=end_user_id, model=model, key_budget_config=_current_model_budget_info, @@ -148,7 +148,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): return True - async def _get_end_user_spend_for_model( + async def get_end_user_spend_for_model( self, end_user_id: str, model: str, @@ -170,7 +170,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): ) return _current_spend - async def _get_virtual_key_spend_for_model( + async def get_virtual_key_spend_for_model( self, user_api_key_hash: str | None, model: str, diff --git a/litellm/proxy/management_endpoints/key_budget_resolver.py b/litellm/proxy/management_endpoints/key_budget_resolver.py new file mode 100644 index 00000000000..bbe1494ec16 --- /dev/null +++ b/litellm/proxy/management_endpoints/key_budget_resolver.py @@ -0,0 +1,1022 @@ +"""Resolve every budget that can gate requests made with one virtual key. + +Enforcement is spread over a dozen checks in ``auth_checks`` that each raise the first +time they trip, so a caller who gets a 429 cannot tell which scope produced it. This +module answers the same question up front: for one key, every applicable scope, its +live spend, and the operator the enforcing check compares with. + +The limit and counter for each scope come from the same helpers enforcement uses, so +the two cannot disagree about which budget applies or which counter it is measured +against. +""" + +import asyncio +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from datetime import datetime +from types import MappingProxyType +from typing import Final, Protocol + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing_extensions import assert_never + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.constants import LITELLM_PROXY_BUDGET_NAME +from litellm.models.budget import LiteLLM_BudgetTable +from litellm.models.end_user import LiteLLM_EndUserTable +from litellm.models.organization import LiteLLM_OrganizationTable +from litellm.models.tag import LiteLLM_TagTable +from litellm.models.team import BudgetLimitEntry, LiteLLM_TeamTable +from litellm.models.user import LiteLLM_UserTable +from litellm.proxy._types import ( + Litellm_EntityType, + LiteLLM_ProjectTableCachedObj, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import ( + TeamMemberBudget, + get_default_end_user_budget, + get_end_user_object, + get_org_object, + get_project_object, + get_tag_objects_batch, + get_team_object, + get_user_object, + resolve_budget_org_id, + resolve_team_member_budget, + user_budget_applies_to_key, +) +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start +from litellm.proxy.spend_tracking.spend_counter_keys import ( + end_user_spend_counter, + key_spend_counter, + key_window_spend_counter, + org_spend_counter, + tag_spend_counter, + team_member_spend_counter, + team_spend_counter, + team_window_spend_counter, + user_spend_counter, +) +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.user_repository import UserRepository +from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BudgetComparison, + BudgetEnforcement, + BudgetScope, + BudgetStatus, + KeyBudgetEntry, +) +from litellm.types.utils import BudgetConfig + +_ENTITY_TYPE_BY_SCOPE: Final[Mapping[BudgetScope, Litellm_EntityType]] = MappingProxyType( + { + "proxy": Litellm_EntityType.PROXY, + "key": Litellm_EntityType.KEY, + "key_window": Litellm_EntityType.KEY, + "key_model": Litellm_EntityType.KEY, + "team": Litellm_EntityType.TEAM, + "team_window": Litellm_EntityType.TEAM, + "team_member": Litellm_EntityType.TEAM_MEMBER, + "user": Litellm_EntityType.USER, + "organization": Litellm_EntityType.ORGANIZATION, + "project": Litellm_EntityType.PROJECT, + "tag": Litellm_EntityType.TAG, + "end_user": Litellm_EntityType.END_USER, + "end_user_model": Litellm_EntityType.END_USER, + } +) + +_ALERT_ONLY_NOTE: Final = "alert only, never blocks; compared against recorded spend rather than the live counter" +_ROLLING_WINDOW_NOTE: Final = "rolling window; the start moves with reset_at so consecutive windows can overlap" +_MODEL_BUDGET_NOTE: Final = "per-model budgets are cache-only and fail open when the cache is cold" +_PROJECT_SPEND_NOTE: Final = "project spend is never incremented today, so this budget cannot trip" +_TAG_NOTE: Final = "key tags are attached to every request; request-supplied tags add budgets not listed here" +_END_USER_ROUTE_NOTE: Final = "only enforced on LLM routes that name this end user" +_USER_ON_TEAM_KEY_NOTE: Final = ( + "the owner's personal budget is not applied to team keys unless " + "general_settings.apply_user_budget_to_team_keys is enabled" +) +_THROTTLE_NOTE: Final = ( + "this key opted into throttle_on_budget_exceeded, so exceeding it slows requests instead of blocking" +) +_SPEND_UNREADABLE_NOTE: Final = "live spend could not be read" + + +class SpendReader(Protocol): + async def __call__( + self, + *, + counter_key: str, + fallback_spend: float, + max_budget: float | None, + window_entity_type: str | None, + window_entity_id: str | None, + window_start: datetime | None, + fallback_authoritative: bool, + ) -> float: ... + + +class ModelSpendReader(Protocol): + async def __call__(self, *, entity_id: str, model: str, budget_config: BudgetConfig) -> float | None: ... + + +async def _read_counter_spend( + *, + counter_key: str, + fallback_spend: float, + max_budget: float | None, + window_entity_type: str | None, + window_entity_id: str | None, + window_start: datetime | None, + fallback_authoritative: bool, +) -> float: + from litellm.proxy.proxy_server import get_current_spend + + return await get_current_spend( + counter_key=counter_key, + fallback_spend=fallback_spend, + max_budget=max_budget, + window_entity_type=window_entity_type, + window_entity_id=window_entity_id, + window_start=window_start, + fallback_authoritative=fallback_authoritative, + ) + + +async def _read_key_model_spend(*, entity_id: str, model: str, budget_config: BudgetConfig) -> float | None: + from litellm.proxy.proxy_server import model_max_budget_limiter + + return await model_max_budget_limiter.get_virtual_key_spend_for_model( + user_api_key_hash=entity_id, + model=model, + key_budget_config=budget_config, + ) + + +async def _read_end_user_model_spend(*, entity_id: str, model: str, budget_config: BudgetConfig) -> float | None: + from litellm.proxy.proxy_server import model_max_budget_limiter + + return await model_max_budget_limiter.get_end_user_spend_for_model( + end_user_id=entity_id, + model=model, + key_budget_config=budget_config, + ) + + +@dataclass(frozen=True, slots=True) +class KeyBudgetResolverDeps: + prisma_client: PrismaClient + user_api_key_cache: UserApiKeyCache + proxy_logging_obj: ProxyLogging + general_settings: Mapping[str, object] + read_spend: SpendReader = field(default=_read_counter_spend) + read_key_model_spend: ModelSpendReader = field(default=_read_key_model_spend) + read_end_user_model_spend: ModelSpendReader = field(default=_read_end_user_model_spend) + + +@dataclass(frozen=True, slots=True) +class _CounterSpend: + counter_key: str + fallback_spend: float + window_entity_type: str | None = None + window_entity_id: str | None = None + window_start: datetime | None = None + fallback_authoritative: bool = False + + +@dataclass(frozen=True, slots=True) +class _RecordedSpend: + value: float + + +@dataclass(frozen=True, slots=True) +class _KeyModelSpend: + key_hash: str + model: str + budget_config: BudgetConfig + + +@dataclass(frozen=True, slots=True) +class _EndUserModelSpend: + end_user_id: str + model: str + budget_config: BudgetConfig + + +_SpendSource = _CounterSpend | _RecordedSpend | _KeyModelSpend | _EndUserModelSpend + + +@dataclass(frozen=True, slots=True) +class _PlannedBudget: + scope: BudgetScope + entity_id: str | None + entity_label: str | None + enforcement: BudgetEnforcement + max_budget: float | None + comparison: BudgetComparison + source: str + spend_source: _SpendSource + budget_duration: str | None = None + budget_reset_at: datetime | None = None + window_start: datetime | None = None + note: str | None = None + + +class _MetadataTags(BaseModel): + model_config = ConfigDict(extra="ignore") + + tags: tuple[str, ...] = () + + +class _MetadataFields(BaseModel): + model_config = ConfigDict(extra="ignore") + + metadata: _MetadataTags = _MetadataTags() + + +class _WindowFields(BaseModel): + model_config = ConfigDict(extra="ignore") + + budget_limits: tuple[BudgetLimitEntry, ...] = () + + +class _ModelBudgetFields(BaseModel): + model_config = ConfigDict(extra="ignore", protected_namespaces=()) + + model_max_budget: Mapping[str, BudgetConfig] = MappingProxyType({}) + + +_METADATA_FIELDS: Final = TypeAdapter(_MetadataFields) +_WINDOW_FIELDS: Final = TypeAdapter(_WindowFields) +_MODEL_BUDGET_FIELDS: Final = TypeAdapter(_ModelBudgetFields) + + +@dataclass(frozen=True, slots=True) +class _TokenBudgetInputs: + """Key columns the token model leaves as bare dicts, re-validated once into the shapes enforcement uses.""" + + tags: tuple[str, ...] + budget_limits: tuple[BudgetLimitEntry, ...] + model_max_budget: Mapping[str, BudgetConfig] + + +def _token_budget_inputs(valid_token: UserAPIKeyAuth) -> _TokenBudgetInputs: + dumped: Final = valid_token.model_dump() + return _TokenBudgetInputs( + tags=_key_tags(dumped), + budget_limits=_budget_windows(dumped), + model_max_budget=_model_budgets(dumped), + ) + + +def _key_tags(dumped: Mapping[str, object]) -> tuple[str, ...]: + """Key metadata tags ride every request this key makes, so their budgets always apply.""" + try: + tags: Final = _METADATA_FIELDS.validate_python(dumped).metadata.tags + except ValidationError: + verbose_proxy_logger.exception("Skipping malformed key metadata tags during budget resolution") + return () + return tuple(dict.fromkeys(tag for tag in tags if tag)) + + +def _budget_windows(dumped: Mapping[str, object]) -> tuple[BudgetLimitEntry, ...]: + try: + return _WINDOW_FIELDS.validate_python(dumped).budget_limits + except ValidationError: + verbose_proxy_logger.exception("Skipping malformed budget_limits entries during budget resolution") + return () + + +def _model_budgets(dumped: Mapping[str, object] | None) -> Mapping[str, BudgetConfig]: + if dumped is None: + return MappingProxyType({}) + try: + return _MODEL_BUDGET_FIELDS.validate_python(dumped).model_max_budget + except ValidationError: + verbose_proxy_logger.exception("Skipping malformed model_max_budget entries during budget resolution") + return MappingProxyType({}) + + +@dataclass(frozen=True, slots=True) +class _ProxyBudget: + spend: float + budget_duration: str | None + budget_reset_at: datetime | None + + +@dataclass(frozen=True, slots=True) +class _BudgetRowMeta: + budget_duration: str | None + budget_reset_at: datetime | None + max_budget: float | None = None + + +@dataclass(frozen=True, slots=True) +class _KeyBudgetContext: + valid_token: UserAPIKeyAuth + token_inputs: _TokenBudgetInputs + end_user_id: str | None + general_settings: Mapping[str, object] + proxy: _ProxyBudget | None + team: LiteLLM_TeamTable | None + user: LiteLLM_UserTable | None + project: LiteLLM_ProjectTableCachedObj | None + organization: LiteLLM_OrganizationTable | None + team_member: TeamMemberBudget | None + tags: tuple[tuple[str, LiteLLM_TagTable | None], ...] + end_user: LiteLLM_EndUserTable | None + default_end_user_budget: LiteLLM_BudgetTable | None + budget_meta: Mapping[str, _BudgetRowMeta] + + +async def resolve_key_budgets( + valid_token: UserAPIKeyAuth, + end_user_id: str | None, + deps: KeyBudgetResolverDeps, +) -> tuple[KeyBudgetEntry, ...]: + """Every budget that applies to ``valid_token``, configured or not, with live spend.""" + context: Final = await _load_context(valid_token=valid_token, end_user_id=end_user_id, deps=deps) + plans: Final = _plan_budgets(context) + spends: Final = await asyncio.gather(*(_read_spend(plan=plan, deps=deps) for plan in plans)) + return tuple(_to_entry(plan=plan, spend=spend) for plan, spend in zip(plans, spends, strict=True)) + + +async def _read_spend(plan: _PlannedBudget, deps: KeyBudgetResolverDeps) -> float | None: + source: Final = plan.spend_source + try: + match source: + case _RecordedSpend(): + return source.value + case _CounterSpend(): + return await deps.read_spend( + counter_key=source.counter_key, + fallback_spend=source.fallback_spend, + max_budget=plan.max_budget, + window_entity_type=source.window_entity_type, + window_entity_id=source.window_entity_id, + window_start=source.window_start, + fallback_authoritative=source.fallback_authoritative, + ) + case _KeyModelSpend(): + return await deps.read_key_model_spend( + entity_id=source.key_hash, model=source.model, budget_config=source.budget_config + ) + case _EndUserModelSpend(): + return await deps.read_end_user_model_spend( + entity_id=source.end_user_id, model=source.model, budget_config=source.budget_config + ) + case _: + assert_never(source) + except Exception: # noqa: BLE001 # one unreadable counter must not blank the whole report + verbose_proxy_logger.exception("Unable to read live spend for budget scope %s", plan.scope) + return None + + +def _to_entry(plan: _PlannedBudget, spend: float | None) -> KeyBudgetEntry: + exceeded: Final = ( + plan.max_budget is not None + and spend is not None + and (spend >= plan.max_budget if plan.comparison == ">=" else spend > plan.max_budget) + ) + status: Final[BudgetStatus] = "unlimited" if plan.max_budget is None else ("exceeded" if exceeded else "ok") + note: Final = ( + _SPEND_UNREADABLE_NOTE + if spend is None and not isinstance(plan.spend_source, _KeyModelSpend | _EndUserModelSpend) + else plan.note + ) + return KeyBudgetEntry( + scope=plan.scope, + entity_type=_ENTITY_TYPE_BY_SCOPE[plan.scope].value, + entity_id=plan.entity_id, + entity_label=plan.entity_label, + enforcement=plan.enforcement, + max_budget=plan.max_budget, + spend=spend, + remaining=(plan.max_budget - spend) if plan.max_budget is not None and spend is not None else None, + comparison=plan.comparison, + budget_duration=plan.budget_duration, + budget_reset_at=plan.budget_reset_at, + window_start=plan.window_start, + source=plan.source, + status=status, + note=note, + ) + + +async def _load_context( + valid_token: UserAPIKeyAuth, + end_user_id: str | None, + deps: KeyBudgetResolverDeps, +) -> _KeyBudgetContext: + token_inputs: Final = _token_budget_inputs(valid_token) + proxy, team, user, project, tags, end_user = await asyncio.gather( + _load_proxy_budget(deps), + _load_team(valid_token.team_id, deps), + _load_user(valid_token.user_id, deps), + _load_project(valid_token.project_id, deps), + _load_tags(token_inputs.tags, deps), + _load_end_user(end_user_id, deps), + ) + organization, team_member, default_end_user_budget = await asyncio.gather( + _load_organization(resolve_budget_org_id(valid_token=valid_token, team_object=team), deps), + _load_team_member(valid_token=valid_token, team=team, deps=deps), + _load_default_end_user_budget(deps), + ) + budget_ids: Final = _referenced_budget_ids( + organization=organization, + project=project, + tags=tags, + end_user=end_user, + default_end_user_budget=default_end_user_budget, + key_budget_id=valid_token.budget_id, + ) + return _KeyBudgetContext( + valid_token=valid_token, + token_inputs=token_inputs, + end_user_id=end_user_id, + general_settings=deps.general_settings, + proxy=proxy, + team=team, + user=user, + project=project, + organization=organization, + team_member=team_member, + tags=tags, + end_user=end_user, + default_end_user_budget=default_end_user_budget, + budget_meta=await _load_budget_meta(budget_ids, deps), + ) + + +def _referenced_budget_ids( + organization: LiteLLM_OrganizationTable | None, + project: LiteLLM_ProjectTableCachedObj | None, + tags: Sequence[tuple[str, LiteLLM_TagTable | None]], + end_user: LiteLLM_EndUserTable | None, + default_end_user_budget: LiteLLM_BudgetTable | None, + key_budget_id: str | None, +) -> frozenset[str]: + candidates: Final = ( + key_budget_id, + organization.budget_id if organization is not None else None, + project.budget_id if project is not None else None, + end_user.budget_id if end_user is not None else None, + default_end_user_budget.budget_id if default_end_user_budget is not None else None, + *(tag.budget_id for _, tag in tags if tag is not None), + ) + return frozenset(budget_id for budget_id in candidates if budget_id is not None) + + +async def _load_budget_meta( + budget_ids: frozenset[str], + deps: KeyBudgetResolverDeps, +) -> Mapping[str, _BudgetRowMeta]: + """Reset schedules live on the budget row but are absent from LiteLLM_BudgetTable, so read them once.""" + if not budget_ids: + return MappingProxyType({}) + try: + rows: Final = await BudgetRepository(deps.prisma_client).find_full_by_ids(sorted(budget_ids)) + except Exception: # noqa: BLE001 # missing reset metadata degrades the report, it must not fail it + verbose_proxy_logger.exception("Unable to load budget reset metadata for %s", sorted(budget_ids)) + return MappingProxyType({}) + return MappingProxyType( + { + row.budget_id: _BudgetRowMeta( + budget_duration=row.budget_duration, + budget_reset_at=row.budget_reset_at, + max_budget=row.max_budget, + ) + for row in rows + if row.budget_id is not None + } + ) + + +async def _load_proxy_budget(deps: KeyBudgetResolverDeps) -> _ProxyBudget | None: + try: + row: Final = await UserRepository(deps.prisma_client).find_by_id(LITELLM_PROXY_BUDGET_NAME) + except Exception: # noqa: BLE001 # every entity load degrades to "unknown" rather than failing the report + verbose_proxy_logger.exception("Unable to load the proxy budget row") + return None + if row is None: + return None + return _ProxyBudget( + spend=row.spend, + budget_duration=row.budget_duration, + budget_reset_at=row.budget_reset_at, + ) + + +async def _load_team(team_id: str | None, deps: KeyBudgetResolverDeps) -> LiteLLM_TeamTable | None: + if team_id is None: + return None + try: + return await get_team_object( + team_id=team_id, + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + proxy_logging_obj=deps.proxy_logging_obj, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load team %s for budget resolution", team_id) + return None + + +async def _load_user(user_id: str | None, deps: KeyBudgetResolverDeps) -> LiteLLM_UserTable | None: + if user_id is None: + return None + try: + return await get_user_object( + user_id=user_id, + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + user_id_upsert=False, + proxy_logging_obj=deps.proxy_logging_obj, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load user %s for budget resolution", user_id) + return None + + +async def _load_project(project_id: str | None, deps: KeyBudgetResolverDeps) -> LiteLLM_ProjectTableCachedObj | None: + if project_id is None: + return None + try: + return await get_project_object( + project_id=project_id, + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + proxy_logging_obj=deps.proxy_logging_obj, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load project %s for budget resolution", project_id) + return None + + +async def _load_organization(org_id: str | None, deps: KeyBudgetResolverDeps) -> LiteLLM_OrganizationTable | None: + if org_id is None: + return None + try: + return await get_org_object( + org_id=org_id, + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + proxy_logging_obj=deps.proxy_logging_obj, + include_budget_table=True, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load organization %s for budget resolution", org_id) + return None + + +async def _load_tags( + tag_names: Sequence[str], + deps: KeyBudgetResolverDeps, +) -> tuple[tuple[str, LiteLLM_TagTable | None], ...]: + if not tag_names: + return () + try: + tag_objects: Final = await get_tag_objects_batch( + tag_names=list(tag_names), + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + proxy_logging_obj=deps.proxy_logging_obj, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load tags %s for budget resolution", tag_names) + return tuple((tag_name, None) for tag_name in tag_names) + return tuple((tag_name, tag_objects.get(tag_name)) for tag_name in tag_names) + + +async def _load_end_user(end_user_id: str | None, deps: KeyBudgetResolverDeps) -> LiteLLM_EndUserTable | None: + if end_user_id is None: + return None + try: + return await get_end_user_object( + end_user_id=end_user_id, + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + proxy_logging_obj=deps.proxy_logging_obj, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load end user %s for budget resolution", end_user_id) + return None + + +async def _load_default_end_user_budget(deps: KeyBudgetResolverDeps) -> LiteLLM_BudgetTable | None: + try: + return await get_default_end_user_budget( + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to load the default end user budget") + return None + + +async def _load_team_member( + valid_token: UserAPIKeyAuth, + team: LiteLLM_TeamTable | None, + deps: KeyBudgetResolverDeps, +) -> TeamMemberBudget | None: + if team is None or valid_token.user_id is None: + return None + try: + return await resolve_team_member_budget( + team_object=team, + user_id=valid_token.user_id, + prisma_client=deps.prisma_client, + user_api_key_cache=deps.user_api_key_cache, + proxy_logging_obj=deps.proxy_logging_obj, + ) + except Exception: # noqa: BLE001 # see _load_proxy_budget + verbose_proxy_logger.exception("Unable to resolve the team-member budget for team %s", team.team_id) + return None + + +def _plan_budgets(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + return ( + *_plan_proxy(context), + *_plan_key(context), + *_plan_key_windows(context), + *_plan_key_models(context), + *_plan_team(context), + *_plan_team_windows(context), + *_plan_team_member(context), + *_plan_user(context), + *_plan_organization(context), + *_plan_project(context), + *_plan_tags(context), + *_plan_end_user(context), + ) + + +def _budget_meta(context: _KeyBudgetContext, budget_id: str | None) -> _BudgetRowMeta: + if budget_id is None: + return _BudgetRowMeta(budget_duration=None, budget_reset_at=None) + return context.budget_meta.get(budget_id, _BudgetRowMeta(budget_duration=None, budget_reset_at=None)) + + +def _plan_proxy(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + proxy: Final = context.proxy + return ( + _PlannedBudget( + scope="proxy", + entity_id=None, + entity_label=None, + enforcement="hard", + max_budget=litellm.max_budget if litellm.max_budget > 0 else None, + comparison=">", + source="litellm_settings.max_budget", + spend_source=_RecordedSpend(proxy.spend if proxy is not None else 0.0), + budget_duration=proxy.budget_duration if proxy is not None else None, + budget_reset_at=proxy.budget_reset_at if proxy is not None else None, + ), + ) + + +def _key_max_budget_source(context: _KeyBudgetContext) -> str: + """The key column wins over its linked budget row, so only name the row when the row is what applies.""" + budget_id: Final = context.valid_token.budget_id + if budget_id is None: + return "key.max_budget" + linked: Final = context.budget_meta.get(budget_id) + if linked is not None and linked.max_budget is not None and linked.max_budget == context.valid_token.max_budget: + return f"budget_table:{budget_id}" + return "key.max_budget" + + +def _plan_key(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + from litellm.proxy.auth.budget_throttle import should_throttle_budget_exceeded + + token: Final = context.valid_token + hard: Final = _PlannedBudget( + scope="key", + entity_id=token.key_alias, + entity_label=token.key_alias, + enforcement="hard", + max_budget=token.max_budget, + comparison=">=", + source=_key_max_budget_source(context), + spend_source=_CounterSpend(counter_key=key_spend_counter(token.token), fallback_spend=token.spend or 0.0), + budget_duration=token.budget_duration, + budget_reset_at=token.budget_reset_at, + note=_THROTTLE_NOTE if should_throttle_budget_exceeded(token) else None, + ) + soft: Final = _PlannedBudget( + scope="key", + entity_id=token.key_alias, + entity_label=token.key_alias, + enforcement="soft", + max_budget=token.soft_budget, + comparison=">=", + source=f"budget_table:{token.budget_id}.soft_budget" if token.budget_id else "key.budget_id.soft_budget", + spend_source=_RecordedSpend(token.spend or 0.0), + budget_duration=token.budget_duration, + budget_reset_at=token.budget_reset_at, + note=_ALERT_ONLY_NOTE, + ) + return (hard, soft) + + +def _plan_key_windows(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + token: Final = context.valid_token + return tuple( + _PlannedBudget( + scope="key_window", + entity_id=window.budget_duration, + entity_label=token.key_alias, + enforcement="hard", + max_budget=window.max_budget, + comparison=">=", + source=f"key.budget_limits[{window.budget_duration}]", + spend_source=_CounterSpend( + counter_key=key_window_spend_counter(token.token, window.budget_duration), + fallback_spend=0.0, + window_entity_type="Key", + window_entity_id=token.token, + window_start=get_budget_window_start(window.model_dump()), + ), + budget_duration=window.budget_duration, + budget_reset_at=window.reset_at, + window_start=get_budget_window_start(window.model_dump()), + note=_ROLLING_WINDOW_NOTE, + ) + for window in context.token_inputs.budget_limits + ) + + +def _plan_key_models(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + token: Final = context.valid_token + return tuple( + _PlannedBudget( + scope="key_model", + entity_id=model, + entity_label=token.key_alias, + enforcement="hard", + max_budget=_positive_or_none(config.max_budget), + comparison=">", + source=f"key.model_max_budget[{model}]", + spend_source=_KeyModelSpend(key_hash=token.token or "", model=model, budget_config=config), + budget_duration=config.budget_duration, + note=_MODEL_BUDGET_NOTE, + ) + for model, config in context.token_inputs.model_max_budget.items() + ) + + +def _positive_or_none(max_budget: float | None) -> float | None: + """Several checks treat a non-positive cap as 'unset' rather than as an immediate block.""" + return max_budget if max_budget is not None and max_budget > 0 else None + + +def _plan_team(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + team: Final = context.team + if team is None: + return () + hard: Final = _PlannedBudget( + scope="team", + entity_id=team.team_id, + entity_label=team.team_alias, + enforcement="hard", + max_budget=team.max_budget, + comparison=">", + source="team.max_budget", + spend_source=_CounterSpend(counter_key=team_spend_counter(team.team_id), fallback_spend=team.spend or 0.0), + budget_duration=team.budget_duration, + budget_reset_at=team.budget_reset_at, + ) + soft: Final = _PlannedBudget( + scope="team", + entity_id=team.team_id, + entity_label=team.team_alias, + enforcement="soft", + max_budget=team.soft_budget, + comparison=">=", + source="team.soft_budget", + spend_source=_RecordedSpend(team.spend or 0.0), + budget_duration=team.budget_duration, + budget_reset_at=team.budget_reset_at, + note=_ALERT_ONLY_NOTE, + ) + return (hard, soft) + + +def _plan_team_windows(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + team: Final = context.team + if team is None: + return () + return tuple( + _PlannedBudget( + scope="team_window", + entity_id=window.budget_duration, + entity_label=team.team_alias, + enforcement="hard", + max_budget=window.max_budget, + comparison=">=", + source=f"team.budget_limits[{window.budget_duration}]", + spend_source=_CounterSpend( + counter_key=team_window_spend_counter(team.team_id, window.budget_duration), + fallback_spend=0.0, + window_entity_type="Team", + window_entity_id=team.team_id, + window_start=get_budget_window_start(window.model_dump()), + ), + budget_duration=window.budget_duration, + budget_reset_at=window.reset_at, + window_start=get_budget_window_start(window.model_dump()), + note=_ROLLING_WINDOW_NOTE, + ) + for window in (team.budget_limits or ()) + ) + + +def _plan_team_member(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + team: Final = context.team + member: Final = context.team_member + user_id: Final = context.valid_token.user_id + if team is None or member is None or user_id is None: + return () + meta: Final = _budget_meta(context, _budget_id_from_source(member.source)) + return ( + _PlannedBudget( + scope="team_member", + entity_id=f"{user_id}:{team.team_id}", + entity_label=team.team_alias, + enforcement="hard", + max_budget=member.max_budget, + comparison=">=", + source=member.source, + spend_source=_CounterSpend( + counter_key=team_member_spend_counter(user_id, team.team_id), + fallback_spend=member.recorded_spend, + ), + budget_duration=meta.budget_duration, + budget_reset_at=meta.budget_reset_at, + ), + ) + + +def _budget_id_from_source(source: str) -> str | None: + prefix: Final = "budget_table:" + return source[len(prefix) :] if source.startswith(prefix) else None + + +def _plan_user(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + user: Final = context.user + if user is None: + return () + applies: Final = user_budget_applies_to_key(team_object=context.team, general_settings=context.general_settings) + return ( + _PlannedBudget( + scope="user", + entity_id=user.user_id, + entity_label=user.user_email or user.user_alias, + enforcement="hard", + max_budget=user.max_budget if applies else None, + comparison=">=", + source="user.max_budget", + spend_source=_CounterSpend(counter_key=user_spend_counter(user.user_id), fallback_spend=user.spend or 0.0), + budget_duration=user.budget_duration, + budget_reset_at=user.budget_reset_at, + note=None if applies else _USER_ON_TEAM_KEY_NOTE, + ), + ) + + +def _plan_organization(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + org: Final = context.organization + if org is None or org.organization_id is None: + return () + meta: Final = _budget_meta(context, org.budget_id) + linked: Final = org.litellm_budget_table + return ( + _PlannedBudget( + scope="organization", + entity_id=org.organization_id, + entity_label=org.organization_alias, + enforcement="hard", + max_budget=_positive_or_none(linked.max_budget if linked is not None else None), + comparison=">=", + source=f"budget_table:{org.budget_id}", + spend_source=_CounterSpend( + counter_key=org_spend_counter(org.organization_id), fallback_spend=org.spend or 0.0 + ), + budget_duration=meta.budget_duration, + budget_reset_at=meta.budget_reset_at, + ), + ) + + +def _plan_project(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + project: Final = context.project + if project is None: + return () + meta: Final = _budget_meta(context, project.budget_id) + linked: Final = project.litellm_budget_table + hard: Final = _PlannedBudget( + scope="project", + entity_id=project.project_id, + entity_label=project.project_alias, + enforcement="hard", + max_budget=linked.max_budget if linked is not None else None, + comparison=">", + source=f"budget_table:{project.budget_id}" if project.budget_id else "project.budget_id", + spend_source=_RecordedSpend(project.spend or 0.0), + budget_duration=meta.budget_duration, + budget_reset_at=meta.budget_reset_at, + note=_PROJECT_SPEND_NOTE, + ) + soft: Final = _PlannedBudget( + scope="project", + entity_id=project.project_id, + entity_label=project.project_alias, + enforcement="soft", + max_budget=linked.soft_budget if linked is not None else None, + comparison=">=", + source=f"budget_table:{project.budget_id}.soft_budget" if project.budget_id else "project.budget_id", + spend_source=_RecordedSpend(project.spend or 0.0), + budget_duration=meta.budget_duration, + budget_reset_at=meta.budget_reset_at, + note=_ALERT_ONLY_NOTE, + ) + return (hard, soft) + + +def _plan_tags(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + return tuple( + _PlannedBudget( + scope="tag", + entity_id=tag_name, + entity_label=None, + enforcement="hard", + max_budget=( + tag.litellm_budget_table.max_budget + if tag is not None and tag.litellm_budget_table is not None + else None + ), + comparison=">", + source=f"budget_table:{tag.budget_id}" if tag is not None and tag.budget_id else "tag.budget_id", + spend_source=_CounterSpend( + counter_key=tag_spend_counter(tag_name), + fallback_spend=(tag.spend or 0.0) if tag is not None else 0.0, + fallback_authoritative=True, + ), + budget_duration=_budget_meta(context, tag.budget_id if tag is not None else None).budget_duration, + budget_reset_at=_budget_meta(context, tag.budget_id if tag is not None else None).budget_reset_at, + note=_TAG_NOTE, + ) + for tag_name, tag in context.tags + ) + + +def _plan_end_user(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: + end_user_id: Final = context.end_user_id + if end_user_id is None: + return () + end_user: Final = context.end_user + budget: Final = end_user.litellm_budget_table if end_user is not None else context.default_end_user_budget + budget_id: Final = budget.budget_id if budget is not None else None + meta: Final = _budget_meta(context, budget_id) + source: Final = ( + f"budget_table:{budget_id}" + if end_user is not None and end_user.budget_id is not None + else "litellm_settings.max_end_user_budget_id" + ) + primary: Final = _PlannedBudget( + scope="end_user", + entity_id=end_user_id, + entity_label=end_user.alias if end_user is not None else None, + enforcement="hard", + max_budget=budget.max_budget if budget is not None else None, + comparison=">", + source=source, + spend_source=_CounterSpend( + counter_key=end_user_spend_counter(end_user_id), + fallback_spend=(end_user.spend or 0.0) if end_user is not None else 0.0, + fallback_authoritative=True, + ), + budget_duration=meta.budget_duration, + budget_reset_at=meta.budget_reset_at, + note=_END_USER_ROUTE_NOTE, + ) + per_model: Final = tuple( + _PlannedBudget( + scope="end_user_model", + entity_id=model, + entity_label=end_user_id, + enforcement="hard", + max_budget=_positive_or_none(config.max_budget), + comparison=">", + source=f"{source}.model_max_budget[{model}]", + spend_source=_EndUserModelSpend(end_user_id=end_user_id, model=model, budget_config=config), + budget_duration=config.budget_duration, + note=_MODEL_BUDGET_NOTE, + ) + for model, config in _model_budgets(budget.model_dump() if budget is not None else None).items() + ) + return (primary, *per_model) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 67a836b8c92..3913777895b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -20,7 +20,7 @@ import secrets import traceback from collections.abc import Awaitable, Callable, Mapping, Sequence from datetime import datetime, timedelta, timezone -from typing import Any, Final, Literal, Optional, Protocol, TypeVar, cast +from typing import Annotated, Any, Final, Literal, Optional, Protocol, TypeVar, cast import fastapi import yaml @@ -51,6 +51,7 @@ from litellm.proxy._types import LiteLLM_VerificationToken, hash_token from litellm.proxy.auth.auth_checks import ( _delete_cache_key_object, can_team_access_model, + get_key_object, get_org_object, get_project_object, get_team_object, @@ -59,7 +60,10 @@ from litellm.proxy.auth.auth_utils import ( abbreviate_api_key, enforce_output_token_estimates_are_admin_only, ) -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.auth.user_api_key_auth import ( + update_key_budget_with_temp_budget_increase, + user_api_key_auth, +) from litellm.proxy.common_utils.callback_utils import ( decrypt_callback_vars, encrypt_callback_vars, @@ -85,6 +89,10 @@ from litellm.proxy.management_endpoints.common_utils import ( validate_budget_duration, validate_finite_spend, ) +from litellm.proxy.management_endpoints.key_budget_resolver import ( + KeyBudgetResolverDeps, + resolve_key_budgets, +) from litellm.proxy.management_endpoints.model_management_endpoints import ( _add_model_to_db, ) @@ -139,6 +147,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import ( BulkUpdateKeyResponse, BulkUpdateTeamKeysRequest, FailedKeyUpdate, + KeyBudgetsResponse, SuccessfulKeyUpdate, ) from litellm.types.router import Deployment @@ -3744,6 +3753,141 @@ async def info_key_fn( raise handle_exception_on_proxy(e) +@router.get( + "/key/{key_id}/budgets", + tags=("key management",), + dependencies=(Depends(user_api_key_auth),), + response_model=KeyBudgetsResponse, +) +@router.get( + "/key/budgets", + tags=("key management",), + dependencies=(Depends(user_api_key_auth),), + response_model=KeyBudgetsResponse, +) +@management_endpoint_wrapper +async def key_budgets_fn( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + key_id: str | None = None, + end_user_id: Annotated[ + str | None, + fastapi.Query( + description="Resolve the budgets that apply to this end user as well. End-user budgets are " + "request-scoped, so they can only be reported for a named end user." + ), + ] = None, +) -> KeyBudgetsResponse: + """ + List every budget that can block requests made with a key, with its live spend. + + A `BudgetExceededError` names one entity, but finding out which of the key, its windows, its + per-model caps, its team, the caller's membership in that team, the owning user, org, project, + the key's tags, the end user or the proxy-wide limit produced it means reading auth source. + This returns all of them at once, including the scopes that are left unconfigured, so a scope + can be ruled out without opening every object. + + Parameters: + - key_id: str | None (path parameter) - The key to inspect. Accepts the plaintext key or its + hash. Defaults to the key in the Authorization header when omitted (`GET /key/budgets`). + - end_user_id: str | None (query parameter) - Also report the budgets that would apply to this + end user. Omitted end users produce no `end_user` rows, because nothing binds an end user to + a key outside a request. + + Returns: + - key: str - The key that was looked up, echoed back as it was passed in + - budgets: list - One entry per applicable budget + - scope: str - `proxy`, `key`, `key_window`, `key_model`, `team`, `team_window`, + `team_member`, `user`, `organization`, `project`, `tag`, `end_user` or `end_user_model` + - entity_type: str - The `Litellm_EntityType` a `BudgetExceededError` from this scope + carries, so a denial message maps back to a row here + - entity_id / entity_label: str | None - Which entity is limited, and its human-facing alias + - enforcement: str - `hard` blocks the request, `soft` only raises an alert + - max_budget: float | None - The limit in effect. `null` means this scope applies to the key + but places no limit on it + - spend: float | None - Spend as the enforcing check reads it, from the same cross-pod + counter, not the periodically-synced database column + - remaining: float | None - `max_budget - spend`, when both are known + - comparison: str - The operator the enforcing check uses, which differs per scope + - budget_duration / budget_reset_at / window_start: When spend next resets to zero + - source: str - Where the limit is configured, e.g. `key.max_budget`, `budget_table:` + - status: str - `unlimited`, `ok` or `exceeded` + - note: str | None - A caveat worth knowing before trusting the row + + Example Curl: + ``` + curl -X GET "http://0.0.0.0:4000/key/sk-test-example-key-123/budgets" \ +-H "Authorization: Bearer sk-1234" + ``` + + Example Curl - the budgets on the calling key itself + ``` + curl -X GET "http://0.0.0.0:4000/key/budgets" \ +-H "Authorization: Bearer sk-test-example-key-123" + ``` + """ + from litellm.proxy.proxy_server import ( + general_settings, + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + try: + if prisma_client is None: + raise Exception( + "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" + ) + + key: Final = key_id or user_api_key_dict.api_key + hashed_key: Final = _hash_token_if_needed(token=key) if key is not None else None + key_info: Final = ( + await VerificationTokenRepository(prisma_client).find_by_id(hashed_key) if hashed_key is not None else None + ) + if key_info is None: + raise ProxyException( + message="Key not found in database", + type=ProxyErrorTypes.not_found_error, + param="key", + code=status.HTTP_404_NOT_FOUND, + ) + + if ( + await _can_user_query_key_info( + user_api_key_dict=user_api_key_dict, + key=key, + key_info=key_info, + ) + is not True + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"You are not allowed to access this key's info. Your role={user_api_key_dict.user_role}", + ) + + # The same object auth resolves the key to, so a stale cached limit is reported as the limit + # that will actually be enforced rather than the database value that will not be. + resolved_key: Final = await get_key_object( + hashed_token=hashed_key, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_dict.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + budgets: Final = await resolve_key_budgets( + valid_token=update_key_budget_with_temp_budget_increase(resolved_key), + end_user_id=end_user_id, + deps=KeyBudgetResolverDeps( + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + general_settings=general_settings, + ), + ) + return KeyBudgetsResponse(key=key, budgets=budgets) + except Exception as e: # noqa: BLE001 # every management handler maps unexpected failures onto the proxy error contract + raise handle_exception_on_proxy(e) + + def _check_model_access_group(models: list[str] | None, llm_router: Router | None, premium_user: bool) -> Literal[True]: """ if is_model_access_group is True + is_wildcard_route is True, check if user is a premium user diff --git a/litellm/proxy/spend_tracking/spend_counter_keys.py b/litellm/proxy/spend_tracking/spend_counter_keys.py new file mode 100644 index 00000000000..c1620ace4bd --- /dev/null +++ b/litellm/proxy/spend_tracking/spend_counter_keys.py @@ -0,0 +1,42 @@ +"""Canonical cross-pod spend counter keys. + +Enforcement, budget reservation and budget introspection must read the exact same +string per scope or they silently observe different counters, so the format lives +here once instead of as an f-string per call site. +""" + + +def key_spend_counter(token: str | None) -> str: + return f"spend:key:{token}" + + +def key_window_spend_counter(token: str | None, budget_duration: str) -> str: + return f"{key_spend_counter(token)}:window:{budget_duration}" + + +def team_spend_counter(team_id: str) -> str: + return f"spend:team:{team_id}" + + +def team_window_spend_counter(team_id: str, budget_duration: str) -> str: + return f"{team_spend_counter(team_id)}:window:{budget_duration}" + + +def team_member_spend_counter(user_id: str, team_id: str) -> str: + return f"spend:team_member:{user_id}:{team_id}" + + +def user_spend_counter(user_id: str) -> str: + return f"spend:user:{user_id}" + + +def org_spend_counter(org_id: str) -> str: + return f"spend:org:{org_id}" + + +def tag_spend_counter(tag_name: str) -> str: + return f"spend:tag:{tag_name}" + + +def end_user_spend_counter(end_user_id: str) -> str: + return f"spend:end_user:{end_user_id}" diff --git a/litellm/repositories/budget_repository.py b/litellm/repositories/budget_repository.py index f6c47b2d639..5dfa461203d 100644 --- a/litellm/repositories/budget_repository.py +++ b/litellm/repositories/budget_repository.py @@ -2,10 +2,11 @@ Budget repository for database operations on LiteLLM_BudgetTable. """ +from collections.abc import Sequence from typing import Any, Final -from litellm.models.budget import LiteLLM_BudgetTable -from litellm.repositories.base_repository import BaseRepository +from litellm.models.budget import LiteLLM_BudgetTable, LiteLLM_BudgetTableFull +from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): @@ -22,6 +23,13 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): async def find_by_id(self, budget_id: str, id_field: str = "budget_id") -> LiteLLM_BudgetTable | None: return await super().find_by_id(budget_id, id_field) + async def find_full_by_ids(self, budget_ids: Sequence[str]) -> tuple[LiteLLM_BudgetTableFull, ...]: + """Reset schedules are server-managed, so they are absent from the model the generic finders return.""" + records: Final[Sequence[DbRecord]] = await self.table.find_many( + where={"budget_id": {"in": list(budget_ids)}} # mutable-ok: prisma builds its query from plain dicts + ) + return tuple(LiteLLM_BudgetTableFull.model_validate(record_to_dict(record)) for record in records) + async def create_budget( self, created_by: str, diff --git a/litellm/types/proxy/management_endpoints/key_management_endpoints.py b/litellm/types/proxy/management_endpoints/key_management_endpoints.py index 0f17f2f23ab..69c38873c1e 100644 --- a/litellm/types/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/types/proxy/management_endpoints/key_management_endpoints.py @@ -106,3 +106,53 @@ class BulkUpdateTeamKeysRequest(BaseModel): if not has_key_ids and not self.all_keys_in_team: raise ValueError("Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`.") return self + + +BudgetScope = Literal[ + "proxy", + "key", + "key_window", + "key_model", + "team", + "team_window", + "team_member", + "user", + "organization", + "project", + "tag", + "end_user", + "end_user_model", +] + +BudgetEnforcement = Literal["hard", "soft"] + +BudgetComparison = Literal[">=", ">"] + +BudgetStatus = Literal["unlimited", "ok", "exceeded"] + + +class KeyBudgetEntry(BaseModel): + """One budget that can gate requests made with a key, with its live spend.""" + + scope: BudgetScope + entity_type: str + entity_id: str | None = None + entity_label: str | None = None + enforcement: BudgetEnforcement + max_budget: float | None = None + spend: float | None = None + remaining: float | None = None + comparison: BudgetComparison + budget_duration: str | None = None + budget_reset_at: datetime | None = None + window_start: datetime | None = None + source: str + status: BudgetStatus + note: str | None = None + + +class KeyBudgetsResponse(BaseModel): + """Every budget that applies to one key, including the ones left unconfigured.""" + + key: str | None = None + budgets: tuple[KeyBudgetEntry, ...] diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index ad852c16905..c35dafd540e 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1768,7 +1768,7 @@ def test_update_key_budget_with_temp_budget_increase(): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import ( - _update_key_budget_with_temp_budget_increase, + update_key_budget_with_temp_budget_increase, ) expiry = datetime.now() + timedelta(days=1) @@ -1782,7 +1782,7 @@ def test_update_key_budget_with_temp_budget_increase(): "temp_budget_expiry": expiry_in_isoformat, }, ) - result = _update_key_budget_with_temp_budget_increase(valid_token) + result = update_key_budget_with_temp_budget_increase(valid_token) assert result.max_budget == 200 assert result is not valid_token assert valid_token.max_budget == 100 diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py index 55459721906..6642cc1a2c3 100644 --- a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py +++ b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py @@ -73,7 +73,7 @@ async def test_is_key_within_model_budget(budget_limiter): # Test when model is within budget with patch.object( - budget_limiter, "_get_virtual_key_spend_for_model", return_value=50.0 + budget_limiter, "get_virtual_key_spend_for_model", return_value=50.0 ): assert ( await budget_limiter.is_key_within_model_budget(user_api_key, "gpt-4") @@ -82,7 +82,7 @@ async def test_is_key_within_model_budget(budget_limiter): # Test when model exceeds budget with patch.object( - budget_limiter, "_get_virtual_key_spend_for_model", return_value=150.0 + budget_limiter, "get_virtual_key_spend_for_model", return_value=150.0 ): with pytest.raises(litellm.BudgetExceededError): await budget_limiter.is_key_within_model_budget(user_api_key, "gpt-4") @@ -94,20 +94,20 @@ async def test_is_key_within_model_budget(budget_limiter): ) -# Test _get_virtual_key_spend_for_model +# Test get_virtual_key_spend_for_model @pytest.mark.asyncio -async def test_get_virtual_key_spend_for_model(budget_limiter): +async def testget_virtual_key_spend_for_model(budget_limiter): budget_config = GenericBudgetInfo(budget_limit=100.0, time_period="1d") # Mock cache get with patch.object(budget_limiter.dual_cache, "async_get_cache", return_value=50.0): - spend = await budget_limiter._get_virtual_key_spend_for_model( + spend = await budget_limiter.get_virtual_key_spend_for_model( user_api_key_hash="test-key", model="gpt-4", key_budget_config=budget_config ) assert spend == 50.0 # Test with provider prefix - spend = await budget_limiter._get_virtual_key_spend_for_model( + spend = await budget_limiter.get_virtual_key_spend_for_model( user_api_key_hash="test-key", model="openai/gpt-4", key_budget_config=budget_config, @@ -165,7 +165,7 @@ async def test_async_log_success_event_uses_per_model_budget_duration(budget_lim async def test_is_end_user_within_model_budget(budget_limiter): # Test when model is within budget with patch.object( - budget_limiter, "_get_end_user_spend_for_model", return_value=50.0 + budget_limiter, "get_end_user_spend_for_model", return_value=50.0 ): assert ( await budget_limiter.is_end_user_within_model_budget( @@ -178,7 +178,7 @@ async def test_is_end_user_within_model_budget(budget_limiter): # Test when model exceeds budget with patch.object( - budget_limiter, "_get_end_user_spend_for_model", return_value=150.0 + budget_limiter, "get_end_user_spend_for_model", return_value=150.0 ): with pytest.raises(litellm.BudgetExceededError): await budget_limiter.is_end_user_within_model_budget( @@ -198,20 +198,20 @@ async def test_is_end_user_within_model_budget(budget_limiter): ) -# Test _get_end_user_spend_for_model +# Test get_end_user_spend_for_model @pytest.mark.asyncio -async def test_get_end_user_spend_for_model(budget_limiter): +async def testget_end_user_spend_for_model(budget_limiter): budget_config = GenericBudgetInfo(budget_limit=100.0, time_period="1d") # Mock cache get with patch.object(budget_limiter.dual_cache, "async_get_cache", return_value=50.0): - spend = await budget_limiter._get_end_user_spend_for_model( + spend = await budget_limiter.get_end_user_spend_for_model( end_user_id="test-user", model="gpt-4", key_budget_config=budget_config ) assert spend == 50.0 # Test with provider prefix - spend = await budget_limiter._get_end_user_spend_for_model( + spend = await budget_limiter.get_end_user_spend_for_model( end_user_id="test-user", model="openai/gpt-4", key_budget_config=budget_config, @@ -473,7 +473,7 @@ async def test_get_fallback_model_within_budget_returns_first_within_budget( budget_fallbacks={"gpt-4": ["gpt-4o-mini", "claude-haiku"]}, ) with patch.object( - budget_limiter, "_get_virtual_key_spend_for_model", return_value=1.0 + budget_limiter, "get_virtual_key_spend_for_model", return_value=1.0 ): result = await budget_limiter.get_fallback_model_within_budget( user_api_key, "gpt-4" @@ -499,7 +499,7 @@ async def test_get_fallback_model_within_budget_skips_exhausted_fallback( with patch.object( budget_limiter, - "_get_virtual_key_spend_for_model", + "get_virtual_key_spend_for_model", side_effect=_spend_for_model, ): result = await budget_limiter.get_fallback_model_within_budget( @@ -521,7 +521,7 @@ async def test_get_fallback_model_within_budget_returns_none_when_chain_exhauste budget_fallbacks={"gpt-4": ["gpt-4o-mini", "claude-haiku"]}, ) with patch.object( - budget_limiter, "_get_virtual_key_spend_for_model", return_value=150.0 + budget_limiter, "get_virtual_key_spend_for_model", return_value=150.0 ): result = await budget_limiter.get_fallback_model_within_budget( user_api_key, "gpt-4" diff --git a/tests/test_litellm/proxy/auth/test_team_member_budget.py b/tests/test_litellm/proxy/auth/test_team_member_budget.py index b38a953d189..3dcd65ec46f 100644 --- a/tests/test_litellm/proxy/auth/test_team_member_budget.py +++ b/tests/test_litellm/proxy/auth/test_team_member_budget.py @@ -436,3 +436,4 @@ async def test_team_member_budget_check_personal_key_not_team(): # Should pass and get_team_membership should not be called assert result is True mock_get_team_membership.assert_not_called() + diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ab7e3d9701c..46e02b14114 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -6227,3 +6227,4 @@ async def test_unlicensed_jwt_auth_is_forbidden_not_unauthorized(): assert error.code == "403" assert "enterprise" in error.message.lower() + diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index bdf09a95e4b..7221b8961bb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -16444,3 +16444,520 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read( assert await _authorized_models_for_key( access_groups, new_token_hash, ["ag-revoked-since", "ag-attached-since"] ) == ["attached-model"] + + +# --------------------------------------------------------------------------- +# GET /key/{key_id}/budgets +# --------------------------------------------------------------------------- + +import contextlib as _budgets_contextlib # noqa: E402 +from datetime import datetime as _budgets_datetime, timezone as _budgets_timezone # noqa: E402 + +from litellm.models.budget import LiteLLM_BudgetTableFull # noqa: E402 +from litellm.models.end_user import LiteLLM_EndUserTable # noqa: E402 +from litellm.models.organization import LiteLLM_OrganizationTable as _BudgetsOrgTable # noqa: E402 +from litellm.models.tag import LiteLLM_TagTable # noqa: E402 +from litellm.models.team import LiteLLM_TeamTable # noqa: E402 +from litellm.proxy._types import LiteLLM_ProjectTableCachedObj # noqa: E402 +from litellm.proxy.auth.auth_checks import TeamMemberBudget # noqa: E402 +from litellm.proxy.management_endpoints.key_budget_resolver import ( # noqa: E402 + KeyBudgetResolverDeps, + resolve_key_budgets, +) +from litellm.types.proxy.management_endpoints.key_management_endpoints import ( # noqa: E402 + KeyBudgetEntry, +) + +_BUDGETS_RESOLVER = "litellm.proxy.management_endpoints.key_budget_resolver" +_BUDGETS_KEY_HASH = "hash-of-the-budgets-key" +_BUDGETS_RESET_AT = _budgets_datetime(2026, 9, 1, tzinfo=_budgets_timezone.utc) + + +class _RecordingSpendReader: + """Stands in for get_current_spend so a test can prove which counter each row was read from.""" + + def __init__(self, spend_by_counter_key): + self.spend_by_counter_key = spend_by_counter_key + self.calls = [] + + async def __call__( + self, + *, + counter_key, + fallback_spend, + max_budget, + window_entity_type, + window_entity_id, + window_start, + fallback_authoritative, + ): + self.calls.append( + { + "counter_key": counter_key, + "fallback_spend": fallback_spend, + "max_budget": max_budget, + "window_entity_type": window_entity_type, + "window_entity_id": window_entity_id, + "window_start": window_start, + "fallback_authoritative": fallback_authoritative, + } + ) + return self.spend_by_counter_key.get(counter_key, 0.0) + + +async def _model_spend_reader(*, entity_id, model, budget_config): + return {"gpt-5": 6.0, "claude-sonnet-4-5": 2.0}.get(model) + + +def _budgets_deps(read_spend=None): + return KeyBudgetResolverDeps( + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + general_settings={}, + read_spend=read_spend or _RecordingSpendReader({}), + read_key_model_spend=_model_spend_reader, + read_end_user_model_spend=_model_spend_reader, + ) + + +def _budgets_token(**overrides): + defaults = dict( + token=_BUDGETS_KEY_HASH, + key_alias="reporting-key", + user_id="user-budgets", + team_id="team-budgets", + project_id="project-budgets", + max_budget=100.0, + spend=1.0, + budget_duration="30d", + budget_reset_at=_BUDGETS_RESET_AT, + metadata={"tags": ["prod"]}, + budget_limits=[ + {"max_budget": 20.0, "budget_duration": "1d", "reset_at": _BUDGETS_RESET_AT} + ], + model_max_budget={"gpt-5": {"max_budget": 5.0, "budget_duration": "1d"}}, + ) + defaults.update(overrides) + return UserAPIKeyAuth(**defaults) + + +@_budgets_contextlib.contextmanager +def _budgets_world( + *, + proxy_row=None, + team=None, + user=None, + project=None, + organization=None, + tags=None, + end_user=None, + default_end_user_budget=None, + team_member=None, + budget_rows=(), +): + user_repository = MagicMock() + user_repository.return_value.find_by_id = AsyncMock(return_value=proxy_row) + budget_repository = MagicMock() + budget_repository.return_value.find_full_by_ids = AsyncMock(return_value=tuple(budget_rows)) + with ( + patch(f"{_BUDGETS_RESOLVER}.UserRepository", user_repository), + patch(f"{_BUDGETS_RESOLVER}.BudgetRepository", budget_repository), + patch(f"{_BUDGETS_RESOLVER}.get_team_object", AsyncMock(return_value=team)), + patch(f"{_BUDGETS_RESOLVER}.get_user_object", AsyncMock(return_value=user)), + patch(f"{_BUDGETS_RESOLVER}.get_project_object", AsyncMock(return_value=project)), + patch(f"{_BUDGETS_RESOLVER}.get_org_object", AsyncMock(return_value=organization)), + patch(f"{_BUDGETS_RESOLVER}.get_tag_objects_batch", AsyncMock(return_value=tags or {})), + patch(f"{_BUDGETS_RESOLVER}.get_end_user_object", AsyncMock(return_value=end_user)), + patch( + f"{_BUDGETS_RESOLVER}.get_default_end_user_budget", + AsyncMock(return_value=default_end_user_budget), + ), + patch( + f"{_BUDGETS_RESOLVER}.resolve_team_member_budget", + AsyncMock(return_value=team_member), + ), + ): + yield + + +def _fully_populated_world(**overrides): + world = dict( + proxy_row=LiteLLM_UserTable( + user_id="litellm-proxy-budget", spend=9.0, budget_duration="1mo", budget_reset_at=_BUDGETS_RESET_AT + ), + team=LiteLLM_TeamTable( + team_id="team-budgets", + team_alias="Reporting Team", + organization_id="org-budgets", + spend=2.0, + max_budget=300.0, + soft_budget=250.0, + budget_limits=[{"max_budget": 30.0, "budget_duration": "7d", "reset_at": _BUDGETS_RESET_AT}], + ), + user=LiteLLM_UserTable(user_id="user-budgets", user_email="owner@example.com", spend=3.0, max_budget=400.0), + project=LiteLLM_ProjectTableCachedObj( + project_id="project-budgets", + project_alias="Reporting Project", + budget_id="budget-project", + spend=4.0, + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-project", max_budget=500.0, soft_budget=450.0), + ), + organization=_BudgetsOrgTable( + organization_id="org-budgets", + organization_alias="Reporting Org", + budget_id="budget-org", + spend=5.0, + created_by="admin", + updated_by="admin", + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-org", max_budget=600.0), + ), + tags={ + "prod": LiteLLM_TagTable( + tag_name="prod", + spend=6.0, + budget_id="budget-tag", + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-tag", max_budget=700.0), + ) + }, + end_user=LiteLLM_EndUserTable( + user_id="end-user-budgets", + blocked=False, + alias="End User", + spend=7.0, + budget_id="budget-end-user", + litellm_budget_table=LiteLLM_BudgetTable( + budget_id="budget-end-user", + max_budget=800.0, + model_max_budget={"claude-sonnet-4-5": {"max_budget": 8.0, "budget_duration": "1d"}}, + ), + ), + team_member=TeamMemberBudget(max_budget=50.0, recorded_spend=8.0, source="budget_table:budget-member"), + budget_rows=( + LiteLLM_BudgetTableFull( + budget_id="budget-member", + max_budget=50.0, + budget_duration="7d", + budget_reset_at=_BUDGETS_RESET_AT, + created_at=_BUDGETS_RESET_AT, + ), + ), + ) + world.update(overrides) + return world + + +@pytest.mark.asyncio +async def test_key_budgets_reports_every_scope_that_applies(): + """Every scope that can gate the key gets a row, so no scope has to be ruled out by hand.""" + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id="end-user-budgets", + deps=_budgets_deps(), + ) + + assert {entry.scope for entry in budgets} == { + "proxy", + "key", + "key_window", + "key_model", + "team", + "team_window", + "team_member", + "user", + "organization", + "project", + "tag", + "end_user", + "end_user_model", + } + by_scope = {(entry.scope, entry.enforcement): entry for entry in budgets} + assert by_scope[("team_member", "hard")].entity_id == "user-budgets:team-budgets" + assert by_scope[("team_member", "hard")].max_budget == 50.0 + assert by_scope[("team_member", "hard")].budget_reset_at == _BUDGETS_RESET_AT + assert by_scope[("organization", "hard")].entity_label == "Reporting Org" + assert by_scope[("end_user_model", "hard")].entity_id == "claude-sonnet-4-5" + assert by_scope[("project", "hard")].note is not None and "never incremented" in by_scope[("project", "hard")].note + + +@pytest.mark.asyncio +async def test_key_budgets_read_live_counter_spend_not_the_database_column(): + """The database column is only a fallback; reporting it would not match the 429 the caller just got.""" + reader = _RecordingSpendReader( + { + "spend:key:" + _BUDGETS_KEY_HASH: 91.0, + "spend:team:team-budgets": 92.0, + "spend:team_member:user-budgets:team-budgets": 93.0, + "spend:user:user-budgets": 94.0, + "spend:org:org-budgets": 95.0, + "spend:tag:prod": 96.0, + "spend:end_user:end-user-budgets": 97.0, + "spend:key:" + _BUDGETS_KEY_HASH + ":window:1d": 98.0, + "spend:team:team-budgets:window:7d": 99.0, + } + ) + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id="end-user-budgets", + deps=_budgets_deps(read_spend=reader), + ) + + hard = {entry.scope: entry for entry in budgets if entry.enforcement == "hard"} + assert hard["key"].spend == 91.0 + assert hard["team"].spend == 92.0 + assert hard["team_member"].spend == 93.0 + assert hard["user"].spend == 94.0 + assert hard["organization"].spend == 95.0 + assert hard["tag"].spend == 96.0 + assert hard["end_user"].spend == 97.0 + assert hard["key_window"].spend == 98.0 + assert hard["team_window"].spend == 99.0 + + fallbacks = {call["counter_key"]: call["fallback_spend"] for call in reader.calls} + assert fallbacks["spend:key:" + _BUDGETS_KEY_HASH] == 1.0 + assert fallbacks["spend:team:team-budgets"] == 2.0 + assert fallbacks["spend:user:user-budgets"] == 3.0 + assert fallbacks["spend:org:org-budgets"] == 5.0 + assert fallbacks["spend:tag:prod"] == 6.0 + assert fallbacks["spend:end_user:end-user-budgets"] == 7.0 + assert fallbacks["spend:team_member:user-budgets:team-budgets"] == 8.0 + + +@pytest.mark.asyncio +async def test_key_budgets_emit_unlimited_rows_for_configured_but_uncapped_scopes(): + """A scope that applies but caps nothing still gets a row; that is what lets a caller rule it out.""" + world = _fully_populated_world( + team=LiteLLM_TeamTable( + team_id="team-budgets", team_alias="Reporting Team", organization_id="org-budgets", spend=2.0 + ), + user=LiteLLM_UserTable(user_id="user-budgets", spend=3.0), + organization=_BudgetsOrgTable( + organization_id="org-budgets", + budget_id="budget-org", + spend=5.0, + created_by="admin", + updated_by="admin", + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-org"), + ), + team_member=TeamMemberBudget(max_budget=None, recorded_spend=8.0, source="team_membership.budget_id"), + ) + with _budgets_world(**world): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(max_budget=None), + end_user_id=None, + deps=_budgets_deps(), + ) + + hard = {entry.scope: entry for entry in budgets if entry.enforcement == "hard"} + for scope in ("key", "team", "team_member", "user", "organization"): + assert hard[scope].max_budget is None, scope + assert hard[scope].status == "unlimited", scope + assert hard[scope].remaining is None, scope + + +@pytest.mark.asyncio +async def test_key_budgets_status_follows_the_operator_the_enforcing_check_uses(): + """The key check blocks at `>=` and the team check at `>`, so equal spend must not read the same.""" + reader = _RecordingSpendReader( + { + "spend:key:" + _BUDGETS_KEY_HASH: 100.0, + "spend:team:team-budgets": 300.0, + } + ) + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id=None, + deps=_budgets_deps(read_spend=reader), + ) + + hard = {entry.scope: entry for entry in budgets if entry.enforcement == "hard"} + assert hard["key"].comparison == ">=" + assert hard["key"].status == "exceeded" + assert hard["team"].comparison == ">" + assert hard["team"].status == "ok" + assert hard["team"].remaining == 0.0 + + +@pytest.mark.asyncio +async def test_key_budgets_never_leak_the_token_hash_or_plaintext_key(): + """The row identifiers are aliases and entity ids; the credential itself must not ride along.""" + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id="end-user-budgets", + deps=_budgets_deps(), + ) + + rendered = json.dumps([entry.model_dump(mode="json") for entry in budgets]) + assert _BUDGETS_KEY_HASH not in rendered + + +@pytest.mark.asyncio +async def test_key_budgets_report_the_personal_budget_as_inapplicable_on_a_team_key(): + """A team key ignores its owner's personal budget, so reporting the number would be a false lead.""" + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id=None, + deps=_budgets_deps(), + ) + + user_entry = next(entry for entry in budgets if entry.scope == "user") + assert user_entry.max_budget is None + assert user_entry.note is not None and "apply_user_budget_to_team_keys" in user_entry.note + + +@pytest.mark.asyncio +async def test_key_budgets_skip_scopes_that_do_not_exist_for_the_key(): + """No team, no project, no tags and no named end user means those scopes cannot gate the key at all.""" + with _budgets_world(user=LiteLLM_UserTable(user_id="user-budgets", spend=3.0, max_budget=400.0)): + budgets = await resolve_key_budgets( + valid_token=_budgets_token( + team_id=None, + project_id=None, + metadata={}, + budget_limits=None, + model_max_budget={}, + ), + end_user_id=None, + deps=_budgets_deps(), + ) + + scopes = {entry.scope for entry in budgets} + assert scopes == {"proxy", "key", "user"} + user_entry = next(entry for entry in budgets if entry.scope == "user") + assert user_entry.max_budget == 400.0 + assert user_entry.note is None + + +@_budgets_contextlib.contextmanager +def _budgets_route_world(*, key_row, budgets=(), caller): + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth as _user_api_key_auth + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_row.model_dump() if key_row is not None else None + ) + app.dependency_overrides[_user_api_key_auth] = lambda: caller + try: + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_key_object", + new_callable=AsyncMock, + return_value=_budgets_token(), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.resolve_key_budgets", + new_callable=AsyncMock, + return_value=tuple(budgets), + ) as resolver, + ): + yield resolver + finally: + app.dependency_overrides.pop(_user_api_key_auth, None) + + +def _budgets_key_row(user_id="user-budgets", team_id="team-budgets"): + return LiteLLM_VerificationToken(token=_BUDGETS_KEY_HASH, user_id=user_id, team_id=team_id) + + +@pytest.mark.asyncio +async def test_key_budgets_route_returns_the_resolved_budgets(): + entry = KeyBudgetEntry( + scope="key", + entity_type="key", + entity_id="reporting-key", + entity_label="reporting-key", + enforcement="hard", + max_budget=100.0, + spend=91.0, + remaining=9.0, + comparison=">=", + source="key.max_budget", + status="ok", + ) + caller = UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value) + with _budgets_route_world(key_row=_budgets_key_row(), budgets=(entry,), caller=caller): + response = client.get(f"/key/{_BUDGETS_KEY_HASH}/budgets?end_user_id=end-user-budgets") + + assert response.status_code == 200 + body = response.json() + assert body["key"] == _BUDGETS_KEY_HASH + assert body["budgets"] == [ + { + "scope": "key", + "entity_type": "key", + "entity_id": "reporting-key", + "entity_label": "reporting-key", + "enforcement": "hard", + "max_budget": 100.0, + "spend": 91.0, + "remaining": 9.0, + "comparison": ">=", + "budget_duration": None, + "budget_reset_at": None, + "window_start": None, + "source": "key.max_budget", + "status": "ok", + "note": None, + } + ] + + +@pytest.mark.asyncio +async def test_key_budgets_route_passes_the_named_end_user_to_the_resolver(): + caller = UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value) + with _budgets_route_world(key_row=_budgets_key_row(), caller=caller) as resolver: + response = client.get(f"/key/{_BUDGETS_KEY_HASH}/budgets?end_user_id=end-user-budgets") + + assert response.status_code == 200 + assert resolver.await_args.kwargs["end_user_id"] == "end-user-budgets" + + +@pytest.mark.asyncio +async def test_key_budgets_route_defaults_to_the_calling_key(): + caller = UserAPIKeyAuth(api_key=_BUDGETS_KEY_HASH, user_id="user-budgets") + with _budgets_route_world(key_row=_budgets_key_row(), caller=caller): + response = client.get("/key/budgets") + + assert response.status_code == 200 + assert response.json()["key"] == _BUDGETS_KEY_HASH + + +@pytest.mark.asyncio +async def test_key_budgets_route_rejects_a_caller_who_may_not_read_the_key(): + """Budget rows expose team, org and user limits, so the same gate as /key/info has to hold.""" + caller = UserAPIKeyAuth( + api_key="sk-stranger", + user_id="someone-else", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.user_belongs_to_keys_team", + new_callable=AsyncMock, + return_value=False, + ), + _budgets_route_world(key_row=_budgets_key_row(), caller=caller), + ): + response = client.get(f"/key/{_BUDGETS_KEY_HASH}/budgets") + + assert response.status_code == 403 + assert "not allowed to access this key's info" in json.dumps(response.json()) + + +@pytest.mark.asyncio +async def test_key_budgets_route_returns_404_for_an_unknown_key(): + caller = UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value) + with _budgets_route_world(key_row=None, caller=caller): + response = client.get("/key/hash-that-does-not-exist/budgets") + + assert response.status_code == 404 + assert "Key not found in database" in json.dumps(response.json()) diff --git a/tests/test_litellm/repositories/test_repositories.py b/tests/test_litellm/repositories/test_repositories.py index 38af52f165c..3f1a509eab6 100644 --- a/tests/test_litellm/repositories/test_repositories.py +++ b/tests/test_litellm/repositories/test_repositories.py @@ -3,7 +3,7 @@ Tests for gateway repository layer. """ import json -from datetime import datetime +from datetime import datetime, timezone from typing import Any, Dict, List, Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -280,6 +280,25 @@ class TestBudgetRepository: ) assert updated.max_budget == 200.0 + @pytest.mark.asyncio + async def test_find_full_by_ids_returns_server_managed_reset_fields(self, repo): + """budget_reset_at is deliberately absent from LiteLLM_BudgetTable, so the generic finders drop it.""" + reset_at = datetime(2026, 9, 1, tzinfo=timezone.utc) + repo._prisma_client.db.litellm_budgettable._records["budget-1"] = { + "budget_id": "budget-1", + "max_budget": 100.0, + "budget_duration": "30d", + "budget_reset_at": reset_at, + "created_at": reset_at, + } + + rows = await repo.find_full_by_ids(["budget-1"]) + + assert [row.budget_id for row in rows] == ["budget-1"] + assert rows[0].budget_reset_at == reset_at + assert rows[0].budget_duration == "30d" + assert rows[0].max_budget == 100.0 + @pytest.mark.asyncio async def test_delete_budget(self, repo): repo._prisma_client.db.litellm_budgettable._records["budget-1"] = {