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.
This commit is contained in:
ryan-crabbe-berri 2026-08-19 15:38:28 -07:00
parent 5b573c552d
commit 5833d99e40
15 changed files with 1968 additions and 105 deletions

View file

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

View file

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

View file

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

View file

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

File diff suppressed because it is too large Load diff

View file

@ -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:<id>`
- 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -6227,3 +6227,4 @@ async def test_unlicensed_jwt_auth_is_forbidden_not_unauthorized():
assert error.code == "403"
assert "enterprise" in error.message.lower()

View file

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

View file

@ -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"] = {