mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(proxy): refine budget_exceeded_models_policy, add user-budget checks and unit tests
This commit is contained in:
parent
f93b8c1e0b
commit
1dd18ff5e2
2 changed files with 245 additions and 6 deletions
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
158
tests/test_litellm/proxy/test_budget_exceeded_models_policy.py
Normal file
158
tests/test_litellm/proxy/test_budget_exceeded_models_policy.py
Normal 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"]
|
||||
Loading…
Add table
Reference in a new issue