From 1dd18ff5e24991819a3379f3a594b53c49d3d1c8 Mon Sep 17 00:00:00 2001 From: perseus <51974392+tcconnally@users.noreply.github.com> Date: Wed, 24 Jun 2026 10:33:00 -0500 Subject: [PATCH] feat(proxy): refine budget_exceeded_models_policy, add user-budget checks and unit tests --- litellm/proxy/proxy_server.py | 93 ++++++++++- .../test_budget_exceeded_models_policy.py | 158 ++++++++++++++++++ 2 files changed, 245 insertions(+), 6 deletions(-) create mode 100644 tests/test_litellm/proxy/test_budget_exceeded_models_policy.py diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 42718171a82..33b4ef2bed1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8667,6 +8667,84 @@ class ProxyStartupEvent: ) +def _apply_budget_exceeded_models_policy( + user_api_key_dict: Optional[UserAPIKeyAuth], + all_models: list, +) -> list: + """ + Applies the budget_exceeded_models_policy general setting to model listings. + Supports list of strings or list of model dictionaries. + """ + global general_settings, llm_router + if general_settings is None or llm_router is None or not user_api_key_dict: + return all_models + + settings = cast(dict[str, object], general_settings) + policy = settings.get("budget_exceeded_models_policy", "blocked") + + # If policy is not free_only or blocked, we just return all_models (the "all" behavior) + if policy not in ("free_only", "blocked"): + return all_models + + # Determine if any budget is exceeded + budget_exceeded = False + current_cost = 0.0 + max_budget = 0.0 + + if ( + user_api_key_dict.max_budget is not None + and (user_api_key_dict.spend or 0) >= user_api_key_dict.max_budget + ): + budget_exceeded = True + current_cost = user_api_key_dict.spend or 0.0 + max_budget = user_api_key_dict.max_budget + elif ( + user_api_key_dict.team_max_budget is not None + and (user_api_key_dict.team_spend or 0) >= user_api_key_dict.team_max_budget + ): + budget_exceeded = True + current_cost = user_api_key_dict.team_spend or 0.0 + max_budget = user_api_key_dict.team_max_budget + elif ( + user_api_key_dict.user_max_budget is not None + and (user_api_key_dict.user_spend or 0) >= user_api_key_dict.user_max_budget + ): + budget_exceeded = True + current_cost = user_api_key_dict.user_spend or 0.0 + max_budget = user_api_key_dict.user_max_budget + + if not budget_exceeded: + return all_models + + # If policy is "blocked" and budget is exceeded, raise BudgetExceededError (429) + if policy == "blocked": + import litellm + raise litellm.BudgetExceededError( + current_cost=current_cost, + max_budget=max_budget, + ) + + # If policy is "free_only" and budget is exceeded, filter to zero-cost models + if policy == "free_only": + from litellm.proxy.auth.auth_checks import _is_model_cost_zero + + filtered_models = [] + for m in all_models: + if isinstance(m, str): + model_name = m + elif isinstance(m, dict): + model_name = m.get("model_name") or m.get("model") or m.get("litellm_params", {}).get("model") + else: + model_name = getattr(m, "model_name", None) or getattr(m, "model", None) + + if model_name and _is_model_cost_zero(model_name, llm_router): + filtered_models.append(m) + + return filtered_models + + return all_models + + #### API ENDPOINTS #### @router.get( "/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"] @@ -8832,12 +8910,8 @@ async def model_list( if hidden_names: all_models = [m for m in all_models if m not in hidden_names] - # Budget-exceeded models policy: when "free_only", return only zero-cost models - policy = settings.get("budget_exceeded_models_policy", "blocked") - if policy == "free_only": - from litellm.proxy.auth.auth_checks import _is_model_cost_zero - - all_models = [m for m in all_models if _is_model_cost_zero(m, llm_router)] + # Budget-exceeded models policy + all_models = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) # Surface the public team name by default; legacy internal keys via flag. # The internal routing key drives the metadata/fallback lookup, while the @@ -12562,6 +12636,10 @@ async def model_info_v2( # Translate `model_name` to the public name for team-scoped rows. all_models = [_translate_model_name_for_response(m) for m in all_models] + # Budget-exceeded models policy + all_models = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + search_total_count = len(all_models) + return _paginate_models_response( all_models=all_models, page=page, @@ -13336,6 +13414,9 @@ async def model_info_v1( user_api_key_dict=user_api_key_dict, ) + # Budget-exceeded models policy + all_models = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + verbose_proxy_logger.debug("all_models: %s", all_models) return {"data": all_models} diff --git a/tests/test_litellm/proxy/test_budget_exceeded_models_policy.py b/tests/test_litellm/proxy/test_budget_exceeded_models_policy.py new file mode 100644 index 00000000000..6fa8ad8f241 --- /dev/null +++ b/tests/test_litellm/proxy/test_budget_exceeded_models_policy.py @@ -0,0 +1,158 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import litellm +from unittest.mock import MagicMock, patch +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.proxy_server import _apply_budget_exceeded_models_policy + + +@pytest.mark.asyncio +async def test_apply_budget_exceeded_models_policy_all(): + """ + When policy is "all", the full model list is returned regardless of budget. + """ + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + max_budget=10.0, + spend=15.0, # Over budget + ) + + all_models = ["gpt-4", "gpt-3.5-turbo"] + + mock_settings = {"budget_exceeded_models_policy": "all"} + mock_router = MagicMock() + + with patch("litellm.proxy.proxy_server.general_settings", mock_settings), \ + patch("litellm.proxy.proxy_server.llm_router", mock_router): + + result = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + assert result == all_models + + +@pytest.mark.asyncio +async def test_apply_budget_exceeded_models_policy_blocked(): + """ + When policy is "blocked" and budget is exceeded, BudgetExceededError is raised. + """ + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + max_budget=10.0, + spend=15.0, # Over budget + ) + + all_models = ["gpt-4", "gpt-3.5-turbo"] + + mock_settings = {"budget_exceeded_models_policy": "blocked"} + mock_router = MagicMock() + + with patch("litellm.proxy.proxy_server.general_settings", mock_settings), \ + patch("litellm.proxy.proxy_server.llm_router", mock_router): + + with pytest.raises(litellm.BudgetExceededError): + _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + + +@pytest.mark.asyncio +async def test_apply_budget_exceeded_models_policy_free_only(): + """ + When policy is "free_only" and budget is exceeded, only zero-cost models are returned. + """ + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + max_budget=10.0, + spend=15.0, # Over budget + ) + + all_models = ["gpt-4", "free-model"] + + mock_settings = {"budget_exceeded_models_policy": "free_only"} + mock_router = MagicMock() + + def mock_is_model_cost_zero(model, router): + return model == "free-model" + + with patch("litellm.proxy.proxy_server.general_settings", mock_settings), \ + patch("litellm.proxy.proxy_server.llm_router", mock_router), \ + patch("litellm.proxy.auth.auth_checks._is_model_cost_zero", side_effect=mock_is_model_cost_zero): + + result = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + assert result == ["free-model"] + + +@pytest.mark.asyncio +async def test_apply_budget_exceeded_models_policy_under_budget(): + """ + When the caller is under budget, all models are returned regardless of policy. + """ + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + max_budget=10.0, + spend=5.0, # Under budget + ) + + all_models = ["gpt-4", "free-model"] + + mock_settings = {"budget_exceeded_models_policy": "free_only"} + mock_router = MagicMock() + + with patch("litellm.proxy.proxy_server.general_settings", mock_settings), \ + patch("litellm.proxy.proxy_server.llm_router", mock_router): + + result = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + assert result == all_models + + +@pytest.mark.asyncio +async def test_apply_budget_exceeded_models_policy_team_budget_exceeded(): + """ + Tests that team budget exceedance is also honored. + """ + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + team_max_budget=100.0, + team_spend=120.0, # Over budget + ) + + all_models = ["gpt-4", "free-model"] + + mock_settings = {"budget_exceeded_models_policy": "free_only"} + mock_router = MagicMock() + + def mock_is_model_cost_zero(model, router): + return model == "free-model" + + with patch("litellm.proxy.proxy_server.general_settings", mock_settings), \ + patch("litellm.proxy.proxy_server.llm_router", mock_router), \ + patch("litellm.proxy.auth.auth_checks._is_model_cost_zero", side_effect=mock_is_model_cost_zero): + + result = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + assert result == ["free-model"] + + +@pytest.mark.asyncio +async def test_apply_budget_exceeded_models_policy_user_budget_exceeded(): + """ + Tests that user budget exceedance is also honored. + """ + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + user_max_budget=50.0, + user_spend=60.0, # Over budget + ) + + all_models = ["gpt-4", "free-model"] + + mock_settings = {"budget_exceeded_models_policy": "free_only"} + mock_router = MagicMock() + + def mock_is_model_cost_zero(model, router): + return model == "free-model" + + with patch("litellm.proxy.proxy_server.general_settings", mock_settings), \ + patch("litellm.proxy.proxy_server.llm_router", mock_router), \ + patch("litellm.proxy.auth.auth_checks._is_model_cost_zero", side_effect=mock_is_model_cost_zero): + + result = _apply_budget_exceeded_models_policy(user_api_key_dict, all_models) + assert result == ["free-model"]