From 5833d99e409ab04c7381f03916c7191fe60c314e Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 19 Aug 2026 15:38:28 -0700 Subject: [PATCH] feat(key management): add GET /key/{key_id}/budgets A BudgetExceededError names one entity, so a caller who gets a 429 still has to read auth source to work out which of the key, its windows, its per-model caps, its team, their membership in that team, the owning user, org, project, the key's tags, the end user or the proxy-wide limit produced it. This returns all of them in one call, with the live spend and reset schedule of each, including the scopes that are left unconfigured so they can be ruled out without opening every object. GET /key/budgets reports the calling key. Both routes reuse _can_user_query_key_info, so reading another key's budgets needs the same rights as reading its info, and 404 on an unknown key matches /key/info. The report has to agree with enforcement or it is worse than nothing, so the resolver consumes the same UserAPIKeyAuth get_key_object hands the auth path, reads spend through get_current_spend, and shares the limit resolution with the checks: counter key strings now come from one spend_counter_keys module, and the team-member, personal-budget-on-team-key and budget-org-id rules were extracted out of auth_checks for both callers. Each row carries the operator its check actually uses, since they differ per scope, plus a note where a budget cannot behave the way its numbers suggest. --- litellm/proxy/_types.py | 6 + litellm/proxy/auth/auth_checks.py | 207 ++-- litellm/proxy/auth/user_api_key_auth.py | 4 +- .../proxy/hooks/model_max_budget_limiter.py | 8 +- .../key_budget_resolver.py | 1022 +++++++++++++++++ .../key_management_endpoints.py | 148 ++- .../spend_tracking/spend_counter_keys.py | 42 + litellm/repositories/budget_repository.py | 12 +- .../key_management_endpoints.py | 50 + tests/proxy_unit_tests/test_proxy_utils.py | 4 +- ...test_unit_test_max_model_budget_limiter.py | 30 +- .../proxy/auth/test_team_member_budget.py | 1 + .../proxy/auth/test_user_api_key_auth.py | 1 + .../test_key_management_endpoints.py | 517 +++++++++ .../repositories/test_repositories.py | 21 +- 15 files changed, 1968 insertions(+), 105 deletions(-) create mode 100644 litellm/proxy/management_endpoints/key_budget_resolver.py create mode 100644 litellm/proxy/spend_tracking/spend_counter_keys.py 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"] = {