diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index a77b313bd7d..85f96480016 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -201,18 +201,28 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): ) return _current_spend - def _get_request_model_budget_config( + def get_request_model_budget_key( self, model: str, internal_model_max_budget: GenericBudgetConfigType - ) -> BudgetConfig | None: + ) -> str | None: """ - Get the budget config for the request model + Which `model_max_budget` key a request model is charged against. 1. Check if `model` is in `internal_model_max_budget` 2. If not, check if `model` without custom llm provider is in `internal_model_max_budget` """ - return internal_model_max_budget.get(model, None) or internal_model_max_budget.get( - self._get_model_without_custom_llm_provider(model), None + if model in internal_model_max_budget: + return model + stripped: Final = self._get_model_without_custom_llm_provider(model) + return stripped if stripped in internal_model_max_budget else None + + def _get_request_model_budget_config( + self, model: str, internal_model_max_budget: GenericBudgetConfigType + ) -> BudgetConfig | None: + """Get the budget config for the request model.""" + matched: Final = self.get_request_model_budget_key( + model=model, internal_model_max_budget=internal_model_max_budget ) + return internal_model_max_budget.get(matched, None) if matched is not None else None def _get_model_without_custom_llm_provider(self, model: str) -> str: if "/" in model: diff --git a/litellm/proxy/management_endpoints/key_budget_resolver.py b/litellm/proxy/management_endpoints/key_budget_resolver.py index d5e6ba58c0e..55b57a01cad 100644 --- a/litellm/proxy/management_endpoints/key_budget_resolver.py +++ b/litellm/proxy/management_endpoints/key_budget_resolver.py @@ -15,7 +15,7 @@ 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 typing import Final, Protocol, TypeVar from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError from typing_extensions import assert_never @@ -90,12 +90,28 @@ _ENTITY_TYPE_BY_SCOPE: Final[Mapping[BudgetScope, Litellm_EntityType]] = Mapping } ) +_RESERVATION_COVERED_SCOPES: Final[frozenset[BudgetScope]] = frozenset( + {"key", "key_window", "team", "team_window", "team_member", "user", "organization", "tag", "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" +_MODEL_BUDGET_NOTE: Final = ( + "per-model spend is counted separately for every request model that maps onto this cap, " + "and each counter is compared against it on its own, so the highest is reported" +) +_MODEL_BUDGET_COLD_NOTE: Final = ( + "no per-model counter exists yet; these budgets are cache-only and fail open until one does" +) +_RESERVATION_NOTE: Final = ( + "the reservation layer blocks this scope as soon as spend reaches the limit, before the read-time check would trip" +) _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" +_CUSTOM_AUTH_END_USER_NOTE: Final = ( + "a custom auth callable can set a request-scoped end user cap that overrides this one and is not visible here" +) _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" @@ -121,7 +137,11 @@ class SpendReader(Protocol): class ModelSpendReader(Protocol): - async def __call__(self, *, entity_id: str, model: str, budget_config: BudgetConfig) -> float | None: ... + async def __call__(self, *, entity_id: str, models: Sequence[str], budget_config: BudgetConfig) -> float | None: ... + + +class ModelBudgetKeyMatcher(Protocol): + def __call__(self, *, model: str, configured: Mapping[str, BudgetConfig]) -> str | None: ... async def _read_counter_spend( @@ -147,35 +167,75 @@ async def _read_counter_spend( ) -async def _read_key_model_spend(*, entity_id: str, model: str, budget_config: BudgetConfig) -> float | None: +def _highest(spends: Sequence[float | None]) -> float | None: + """Each request model has its own counter compared against the same cap, so the highest is the one that blocks.""" + found: Final = tuple(spend for spend in spends if spend is not None) + return max(found) if found else None + + +async def _read_key_model_spend(*, entity_id: str, models: Sequence[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, + return _highest( + await asyncio.gather( + *( + model_max_budget_limiter.get_virtual_key_spend_for_model( + user_api_key_hash=entity_id, + model=model, + key_budget_config=budget_config, + ) + for model in models + ) + ) ) -async def _read_end_user_model_spend(*, entity_id: str, model: str, budget_config: BudgetConfig) -> float | None: +async def _read_end_user_model_spend( + *, entity_id: str, models: Sequence[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, + return _highest( + await asyncio.gather( + *( + model_max_budget_limiter.get_end_user_spend_for_model( + end_user_id=entity_id, + model=model, + key_budget_config=budget_config, + ) + for model in models + ) + ) ) +def _match_model_budget_key(*, model: str, configured: Mapping[str, BudgetConfig]) -> str | None: + from litellm.proxy.proxy_server import model_max_budget_limiter + + return model_max_budget_limiter.get_request_model_budget_key( + model=model, internal_model_max_budget=dict(configured) + ) + + +def _request_models(allowed_models: tuple[str, ...]) -> tuple[str, ...]: + """Model names a request could carry, since per-model counters are keyed by the request model.""" + from litellm.proxy.proxy_server import llm_router + + router_models: Final = tuple(llm_router.get_model_names()) if llm_router is not None else () + return tuple(dict.fromkeys((*allowed_models, *router_models))) + + @dataclass(frozen=True, slots=True) class KeyBudgetResolverDeps: prisma_client: PrismaClient user_api_key_cache: UserApiKeyCache proxy_logging_obj: ProxyLogging general_settings: Mapping[str, object] + custom_auth_enabled: bool = False 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) + match_model_budget_key: ModelBudgetKeyMatcher = field(default=_match_model_budget_key) @dataclass(frozen=True, slots=True) @@ -196,14 +256,14 @@ class _RecordedSpend: @dataclass(frozen=True, slots=True) class _KeyModelSpend: key_hash: str - model: str + models: tuple[str, ...] budget_config: BudgetConfig @dataclass(frozen=True, slots=True) class _EndUserModelSpend: end_user_id: str - model: str + models: tuple[str, ...] budget_config: BudgetConfig @@ -238,21 +298,42 @@ class _MetadataFields(BaseModel): metadata: _MetadataTags = _MetadataTags() -class _WindowFields(BaseModel): +class _KeyModelsFields(BaseModel): model_config = ConfigDict(extra="ignore") - budget_limits: tuple[BudgetLimitEntry, ...] = () + models: tuple[str, ...] = () + + +class _WindowFields(BaseModel): + """Entries stay unparsed here so that one malformed window cannot discard the rest.""" + + model_config = ConfigDict(extra="ignore") + + budget_limits: tuple[object, ...] = () class _ModelBudgetFields(BaseModel): model_config = ConfigDict(extra="ignore", protected_namespaces=()) - model_max_budget: Mapping[str, BudgetConfig] = MappingProxyType({}) + model_max_budget: Mapping[str, object] = MappingProxyType({}) +_T = TypeVar("_T") + _METADATA_FIELDS: Final = TypeAdapter(_MetadataFields) _WINDOW_FIELDS: Final = TypeAdapter(_WindowFields) _MODEL_BUDGET_FIELDS: Final = TypeAdapter(_ModelBudgetFields) +_BUDGET_LIMIT_ENTRY: Final = TypeAdapter(BudgetLimitEntry) +_BUDGET_CONFIG: Final = TypeAdapter(BudgetConfig) +_KEY_MODELS: Final = TypeAdapter(_KeyModelsFields) + + +def _validated(adapter: TypeAdapter[_T], value: object, field_name: str) -> _T | None: + try: + return adapter.validate_python(value) + except ValidationError: + verbose_proxy_logger.exception("Skipping malformed %s entry during budget resolution", field_name) + return None @dataclass(frozen=True, slots=True) @@ -262,14 +343,17 @@ class _TokenBudgetInputs: tags: tuple[str, ...] budget_limits: tuple[BudgetLimitEntry, ...] model_max_budget: Mapping[str, BudgetConfig] + models: tuple[str, ...] def _token_budget_inputs(valid_token: UserAPIKeyAuth) -> _TokenBudgetInputs: dumped: Final = valid_token.model_dump() + container: Final = _validated(_KEY_MODELS, dumped, "models") return _TokenBudgetInputs( tags=_key_tags(dumped), budget_limits=_budget_windows(dumped), model_max_budget=_model_budgets(dumped), + models=container.models if container is not None else (), ) @@ -284,21 +368,26 @@ def _key_tags(dumped: Mapping[str, object]) -> tuple[str, ...]: 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") + """Enforcement still applies the windows that parse, so a bad neighbour must not hide them.""" + container: Final = _validated(_WINDOW_FIELDS, dumped, "budget_limits") + if container is None: return () + parsed: Final = (_validated(_BUDGET_LIMIT_ENTRY, entry, "budget_limits") for entry in container.budget_limits) + return tuple(entry for entry in parsed if entry is not None) def _model_budgets(dumped: Mapping[str, object] | None) -> Mapping[str, BudgetConfig]: + """Same as ``_budget_windows``: drop only the per-model caps that fail to parse.""" 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") + container: Final = _validated(_MODEL_BUDGET_FIELDS, dumped, "model_max_budget") + if container is None: return MappingProxyType({}) + parsed: Final = ( + (model, _validated(_BUDGET_CONFIG, config, "model_max_budget")) + for model, config in container.model_max_budget.items() + ) + return MappingProxyType({model: config for model, config in parsed if config is not None}) @dataclass(frozen=True, slots=True) @@ -320,6 +409,10 @@ class _KeyBudgetContext: valid_token: UserAPIKeyAuth token_inputs: _TokenBudgetInputs end_user_id: str | None + token_end_user_max_budget: float | None + custom_auth_enabled: bool + request_models: tuple[str, ...] + match_model_budget_key: ModelBudgetKeyMatcher general_settings: Mapping[str, object] proxy: _ProxyBudget | None team: LiteLLM_TeamTable | None @@ -337,12 +430,22 @@ async def resolve_key_budgets( valid_token: UserAPIKeyAuth, end_user_id: str | None, deps: KeyBudgetResolverDeps, + token_end_user_max_budget: float | None = None, ) -> 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) + context: Final = await _load_context( + valid_token=valid_token, + end_user_id=end_user_id, + deps=deps, + token_end_user_max_budget=token_end_user_max_budget, + ) 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)) + reservation_enabled: Final = deps.general_settings.get("disable_budget_reservation") is not True + return tuple( + _to_entry(plan=plan, spend=spend, reservation_enabled=reservation_enabled) + for plan, spend in zip(plans, spends, strict=True) + ) async def _read_spend(plan: _PlannedBudget, deps: KeyBudgetResolverDeps) -> float | None: @@ -363,11 +466,11 @@ async def _read_spend(plan: _PlannedBudget, deps: KeyBudgetResolverDeps) -> floa ) case _KeyModelSpend(): return await deps.read_key_model_spend( - entity_id=source.key_hash, model=source.model, budget_config=source.budget_config + entity_id=source.key_hash, models=source.models, 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 + entity_id=source.end_user_id, models=source.models, budget_config=source.budget_config ) case _: assert_never(source) @@ -376,18 +479,39 @@ async def _read_spend(plan: _PlannedBudget, deps: KeyBudgetResolverDeps) -> floa return None -def _to_entry(plan: _PlannedBudget, spend: float | None) -> KeyBudgetEntry: +def _effective_comparison(plan: _PlannedBudget, reservation_enabled: bool) -> BudgetComparison: + """ + Reservation runs ahead of the read-time check and blocks once spend reaches the cap, so for + every scope it covers it, not the read-time comparison, decides when a request stops going through. + """ + if plan.enforcement == "soft" or not reservation_enabled: + return plan.comparison + if plan.scope not in _RESERVATION_COVERED_SCOPES: + return plan.comparison + return ">=" + + +def _entry_note(plan: _PlannedBudget, spend: float | None, comparison: BudgetComparison) -> str | None: + if spend is None: + return ( + _MODEL_BUDGET_COLD_NOTE + if isinstance(plan.spend_source, _KeyModelSpend | _EndUserModelSpend) + else _SPEND_UNREADABLE_NOTE + ) + if comparison == plan.comparison: + return plan.note + return _RESERVATION_NOTE if plan.note is None else f"{_RESERVATION_NOTE}; {plan.note}" + + +def _to_entry(plan: _PlannedBudget, spend: float | None, reservation_enabled: bool) -> KeyBudgetEntry: + comparison: Final = _effective_comparison(plan=plan, reservation_enabled=reservation_enabled) 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) + and (spend >= plan.max_budget if 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 - ) + note: Final = _entry_note(plan=plan, spend=spend, comparison=comparison) return KeyBudgetEntry( scope=plan.scope, entity_type=_ENTITY_TYPE_BY_SCOPE[plan.scope], @@ -397,7 +521,7 @@ def _to_entry(plan: _PlannedBudget, spend: float | None) -> KeyBudgetEntry: 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, + comparison=comparison, budget_duration=plan.budget_duration, budget_reset_at=plan.budget_reset_at, window_start=plan.window_start, @@ -411,6 +535,7 @@ async def _load_context( valid_token: UserAPIKeyAuth, end_user_id: str | None, deps: KeyBudgetResolverDeps, + token_end_user_max_budget: float | None, ) -> _KeyBudgetContext: token_inputs: Final = _token_budget_inputs(valid_token) proxy, team, user, project, tags, end_user = await asyncio.gather( @@ -438,6 +563,10 @@ async def _load_context( valid_token=valid_token, token_inputs=token_inputs, end_user_id=end_user_id, + token_end_user_max_budget=token_end_user_max_budget, + custom_auth_enabled=deps.custom_auth_enabled, + request_models=_request_models(token_inputs.models), + match_model_budget_key=deps.match_model_budget_key, general_settings=deps.general_settings, proxy=proxy, team=team, @@ -750,8 +879,21 @@ def _plan_key_windows(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: ) +def _model_counter_candidates( + config_key: str, configured: Mapping[str, BudgetConfig], context: _KeyBudgetContext +) -> tuple[str, ...]: + """Counters are keyed by the request model, so a config key has to be probed under each model routed to it.""" + routed: Final = ( + model + for model in context.request_models + if context.match_model_budget_key(model=model, configured=configured) == config_key + ) + return tuple(dict.fromkeys((config_key, *routed))) + + def _plan_key_models(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: token: Final = context.valid_token + configured: Final = context.token_inputs.model_max_budget return tuple( _PlannedBudget( scope="key_model", @@ -761,11 +903,15 @@ def _plan_key_models(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: 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), + spend_source=_KeyModelSpend( + key_hash=token.token or "", + models=_model_counter_candidates(config_key=model, configured=configured, context=context), + budget_config=config, + ), budget_duration=config.budget_duration, note=_MODEL_BUDGET_NOTE, ) - for model, config in context.token_inputs.model_max_budget.items() + for model, config in configured.items() ) @@ -982,17 +1128,20 @@ def _plan_end_user(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: 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 = ( + row_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" ) + # Reservation reads the request token's cap first and only falls back to the row, so mirror that order. + token_max_budget: Final = context.token_end_user_max_budget + source: Final = "token.end_user_max_budget" if token_max_budget is not None else row_source 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, + max_budget=token_max_budget if token_max_budget is not None else (budget.max_budget if budget else None), comparison=">", source=source, spend_source=_CounterSpend( @@ -1002,8 +1151,13 @@ def _plan_end_user(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: ), budget_duration=meta.budget_duration, budget_reset_at=meta.budget_reset_at, - note=_END_USER_ROUTE_NOTE, + note=( + f"{_END_USER_ROUTE_NOTE}; {_CUSTOM_AUTH_END_USER_NOTE}" + if context.custom_auth_enabled and token_max_budget is None + else _END_USER_ROUTE_NOTE + ), ) + configured: Final = _model_budgets(budget.model_dump() if budget is not None else None) per_model: Final = tuple( _PlannedBudget( scope="end_user_model", @@ -1012,11 +1166,15 @@ def _plan_end_user(context: _KeyBudgetContext) -> tuple[_PlannedBudget, ...]: 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), + source=f"{row_source}.model_max_budget[{model}]", + spend_source=_EndUserModelSpend( + end_user_id=end_user_id, + models=_model_counter_candidates(config_key=model, configured=configured, context=context), + 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() + for model, config in configured.items() ) return (primary, *per_model) 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 7221b8961bb..3318392e174 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 @@ -1,3 +1,4 @@ +import dataclasses import json import os import sys @@ -16461,6 +16462,8 @@ 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 + _MODEL_BUDGET_COLD_NOTE, + _RESERVATION_NOTE, KeyBudgetResolverDeps, resolve_key_budgets, ) @@ -16505,19 +16508,31 @@ class _RecordingSpendReader: 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) +class _RecordingModelSpendReader: + """Stands in for the per-model cache so a test can prove which request models were probed.""" + + def __init__(self, spend_by_model=None): + default = {"gpt-5": 6.0, "claude-sonnet-4-5": 2.0} + self.spend_by_model = default if spend_by_model is None else spend_by_model + self.probed = [] + + async def __call__(self, *, entity_id, models, budget_config): + self.probed.append(tuple(models)) + found = [self.spend_by_model[model] for model in models if model in self.spend_by_model] + return max(found) if found else None -def _budgets_deps(read_spend=None): +def _budgets_deps(read_spend=None, model_spend=None, match_model_budget_key=None, general_settings=None): + model_reader = model_spend or _RecordingModelSpendReader() return KeyBudgetResolverDeps( prisma_client=MagicMock(), user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), - general_settings={}, + general_settings=general_settings if general_settings is not None else {}, read_spend=read_spend or _RecordingSpendReader({}), - read_key_model_spend=_model_spend_reader, - read_end_user_model_spend=_model_spend_reader, + read_key_model_spend=model_reader, + read_end_user_model_spend=model_reader, + match_model_budget_key=match_model_budget_key or (lambda *, model, configured: None), ) @@ -16757,28 +16772,106 @@ async def test_key_budgets_emit_unlimited_rows_for_configured_but_uncapped_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( +_BUDGETS_SPEND_AT_LIMIT = { + "spend:key:" + _BUDGETS_KEY_HASH: 100.0, + "spend:key:" + _BUDGETS_KEY_HASH + ":window:1d": 20.0, + "spend:team:team-budgets": 300.0, + "spend:team:team-budgets:window:7d": 30.0, + "spend:team_member:user-budgets:team-budgets": 50.0, + "spend:user:user-budgets": 400.0, + "spend:org:org-budgets": 600.0, + "spend:tag:prod": 700.0, + "spend:end_user:end-user-budgets": 800.0, +} + +# Reservation reserves before the read-time check and refuses once spend has reached the cap, so every +# scope it covers blocks at ">=" no matter which operator the read-time check in auth_checks uses. +_BUDGETS_RESERVED_SCOPES = ( + "key", + "key_window", + "team", + "team_window", + "team_member", + "user", + "organization", + "tag", + "end_user", +) +_BUDGETS_UNRESERVED_SCOPES = ("proxy", "project", "key_model", "end_user_model") + + +async def _budgets_at_limit(general_settings=None): + world = _fully_populated_world() + world["team_member"] = TeamMemberBudget(max_budget=50.0, recorded_spend=8.0, source="budget_table:budget-member") + # The owner's personal budget only carries a limit on a team key when this is on, and the point of + # this fixture is that every reservation-covered scope has a cap sitting exactly at its spend. + settings = {"apply_user_budget_to_team_keys": True, **(general_settings or {})} + with _budgets_world(**world): + return await resolve_key_budgets( valid_token=_budgets_token(), - end_user_id=None, - deps=_budgets_deps(read_spend=reader), + end_user_id="end-user-budgets", + deps=_budgets_deps( + read_spend=_RecordingSpendReader(_BUDGETS_SPEND_AT_LIMIT), + general_settings=settings, + ), ) - 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.parametrize("scope", _BUDGETS_RESERVED_SCOPES) +@pytest.mark.asyncio +async def test_key_budgets_report_the_reservation_operator_for_every_scope_it_covers(scope): + """Spend exactly at the cap is already blocked by the reservation, so the row must not read `ok`.""" + budgets = await _budgets_at_limit() + + entry = next(e for e in budgets if e.scope == scope and e.enforcement == "hard") + assert entry.max_budget is not None, scope + assert entry.spend == entry.max_budget, scope + assert entry.comparison == ">=", scope + assert entry.status == "exceeded", scope + + +@pytest.mark.parametrize("scope", _BUDGETS_RESERVED_SCOPES) +@pytest.mark.asyncio +async def test_key_budgets_fall_back_to_the_read_time_operator_when_reservation_is_disabled(scope): + """With reservation off the read-time check is the only gate, so `>` scopes stop blocking at the cap.""" + budgets = await _budgets_at_limit(general_settings={"disable_budget_reservation": True}) + + entry = next(e for e in budgets if e.scope == scope and e.enforcement == "hard") + read_time = {"team", "tag", "end_user"} + assert entry.comparison == (">" if scope in read_time else ">="), scope + assert entry.status == ("ok" if scope in read_time else "exceeded"), scope + + +@pytest.mark.asyncio +async def test_key_budgets_leave_unreserved_scopes_on_their_read_time_operator(): + """Reservation never touches these, so reporting `>=` for them would overstate when they block.""" + budgets = await _budgets_at_limit() + + by_scope = {e.scope: e for e in budgets if e.enforcement == "hard"} + for scope in _BUDGETS_UNRESERVED_SCOPES: + assert by_scope[scope].comparison == ">", scope + + +@pytest.mark.asyncio +async def test_key_budgets_keep_soft_budgets_on_their_alerting_operator(): + """Soft budgets only ever raise an alert, so the reservation operator must not be applied to them.""" + budgets = await _budgets_at_limit() + + soft = [entry for entry in budgets if entry.enforcement == "soft"] + assert soft, "expected the fully populated world to configure soft budgets" + assert all(entry.comparison == ">=" for entry in soft) + assert all(_RESERVATION_NOTE not in (entry.note or "") for entry in soft) + + +@pytest.mark.asyncio +async def test_key_budgets_explain_a_scope_whose_operator_the_reservation_tightened(): + """A row that blocks earlier than auth_checks alone would has to say so, not silently change.""" + budgets = await _budgets_at_limit() + + team = next(e for e in budgets if e.scope == "team" and e.enforcement == "hard") + key = next(e for e in budgets if e.scope == "key" and e.enforcement == "hard") + assert _RESERVATION_NOTE in (team.note or "") + assert _RESERVATION_NOTE not in (key.note or ""), "the key check already blocks at >=, nothing was tightened" @pytest.mark.asyncio @@ -16961,3 +17054,144 @@ async def test_key_budgets_route_returns_404_for_an_unknown_key(): assert response.status_code == 404 assert "Key not found in database" in json.dumps(response.json()) + + +@pytest.mark.asyncio +async def test_key_budgets_probe_every_request_model_that_maps_onto_a_per_model_cap(): + """Counters are keyed by the request model, so reading only the config key misses a warm counter.""" + model_reader = _RecordingModelSpendReader({"openai/gpt-5": 6.0}) + token = _budgets_token( + models=["openai/gpt-5"], model_max_budget={"gpt-5": {"max_budget": 5.0, "budget_duration": "1d"}} + ) + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=token, + end_user_id=None, + deps=_budgets_deps( + model_spend=model_reader, + match_model_budget_key=lambda *, model, configured: "gpt-5" if model.endswith("gpt-5") else None, + ), + ) + + entry = next(e for e in budgets if e.scope == "key_model") + assert "openai/gpt-5" in model_reader.probed[0], "the request model's counter was never probed" + assert entry.spend == 6.0 + assert entry.status == "exceeded" + + +@pytest.mark.asyncio +async def test_key_budgets_say_a_per_model_counter_is_missing_rather_than_claiming_a_cold_cache(): + """The stale note claimed the cache was cold even when it was warm under a different model key.""" + warm = _RecordingModelSpendReader({"gpt-5": 6.0}) + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), end_user_id=None, deps=_budgets_deps(model_spend=warm) + ) + + warm_entry = next(e for e in budgets if e.scope == "key_model" and e.entity_id == "gpt-5") + assert warm_entry.spend == 6.0 + assert warm_entry.note != _MODEL_BUDGET_COLD_NOTE + + cold = _RecordingModelSpendReader({}) + with _budgets_world(**_fully_populated_world()): + cold_budgets = await resolve_key_budgets( + valid_token=_budgets_token(), end_user_id=None, deps=_budgets_deps(model_spend=cold) + ) + + cold_entry = next(e for e in cold_budgets if e.scope == "key_model" and e.entity_id == "gpt-5") + assert cold_entry.spend is None + assert cold_entry.note == _MODEL_BUDGET_COLD_NOTE + + +@pytest.mark.asyncio +async def test_key_budgets_keep_the_valid_windows_when_one_window_entry_is_malformed(): + """Enforcement still applies the good windows, so dropping the whole scope hides a live budget.""" + token = _budgets_token( + budget_limits=[ + {"max_budget": 20.0, "budget_duration": "1d", "reset_at": _BUDGETS_RESET_AT}, + {"max_budget": "not-a-number", "budget_duration": "7d"}, + ] + ) + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets(valid_token=token, end_user_id=None, deps=_budgets_deps()) + + windows = [entry for entry in budgets if entry.scope == "key_window"] + assert [entry.entity_id for entry in windows] == ["1d"] + + +@pytest.mark.asyncio +async def test_key_budgets_keep_the_valid_model_caps_when_one_model_entry_is_malformed(): + """Same rule for per-model caps: one bad entry must not blank every other model's budget.""" + token = _budgets_token( + model_max_budget={ + "gpt-5": {"max_budget": 5.0, "budget_duration": "1d"}, + "claude-sonnet-4-5": {"max_budget": "not-a-number"}, + } + ) + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets(valid_token=token, end_user_id=None, deps=_budgets_deps()) + + models = [entry.entity_id for entry in budgets if entry.scope == "key_model"] + assert models == ["gpt-5"] + + +@pytest.mark.asyncio +async def test_key_budgets_prefer_the_request_token_end_user_cap_over_the_row(): + """Reservation reads the token's cap first, so reporting the row's number would understate the block.""" + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id="end-user-budgets", + deps=_budgets_deps(), + token_end_user_max_budget=25.0, + ) + + entry = next(e for e in budgets if e.scope == "end_user") + assert entry.max_budget == 25.0 + assert entry.source == "token.end_user_max_budget" + + +@pytest.mark.asyncio +async def test_key_budgets_warn_that_custom_auth_can_hide_an_end_user_cap(): + """A custom auth callable sets that cap per request, so the report has to admit it cannot see one.""" + with _budgets_world(**_fully_populated_world()): + budgets = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id="end-user-budgets", + deps=_budgets_deps(), + ) + with_custom_auth = await resolve_key_budgets( + valid_token=_budgets_token(), + end_user_id="end-user-budgets", + deps=dataclasses.replace(_budgets_deps(), custom_auth_enabled=True), + ) + + plain = next(e for e in budgets if e.scope == "end_user") + warned = next(e for e in with_custom_auth if e.scope == "end_user") + assert "custom auth" not in (plain.note or "") + assert "custom auth" in (warned.note or "") + + +@pytest.mark.parametrize("max_budget", [0, 0.0, -5.0]) +@pytest.mark.asyncio +async def test_key_budgets_treat_a_non_positive_cap_as_unset(max_budget): + """The enforcing checks skip a cap of 0 rather than blocking everything, so it is not a limit.""" + world = _fully_populated_world() + world["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=max_budget), + ) + token = _budgets_token(model_max_budget={"gpt-5": {"max_budget": max_budget, "budget_duration": "1d"}}) + with _budgets_world(**world): + budgets = await resolve_key_budgets(valid_token=token, end_user_id=None, deps=_budgets_deps()) + + by_scope = {entry.scope: entry for entry in budgets if entry.enforcement == "hard"} + assert by_scope["organization"].max_budget is None + assert by_scope["organization"].status == "unlimited" + assert by_scope["key_model"].max_budget is None + assert by_scope["key_model"].status == "unlimited"