feat(proxy): refine budget_exceeded_models_policy, add user-budget checks and unit tests

This commit is contained in:
perseus 2026-06-24 10:33:00 -05:00
parent f93b8c1e0b
commit 1dd18ff5e2
2 changed files with 245 additions and 6 deletions

View file

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

View file

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