diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index de2c6b51380..db3b03b0550 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2072,6 +2072,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached return LiteLLM_TeamTableCachedObj( team_id=valid_token.team_id, max_budget=valid_token.team_max_budget, + model_max_budget=valid_token.team_model_max_budget, soft_budget=valid_token.team_soft_budget, spend=valid_token.team_spend, tpm_limit=valid_token.team_tpm_limit, diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 7a279b20d94..de10d3134d6 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -386,7 +386,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): self._get_model_without_custom_llm_provider(model), None ) - def _get_model_without_custom_llm_provider(self, model: str) -> str: + @staticmethod + def _get_model_without_custom_llm_provider(model: str) -> str: if "/" in model: return model.split("/")[-1] return model diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index cc9f120216e..0a1fad6be16 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1077,6 +1077,59 @@ def _check_team_model_max_budget_update_authority( ) +def _validate_team_model_max_budget( + model_max_budget: Mapping[str, Mapping[str, str | float]] | None, +) -> None: + """ + Shared /team/new + /team/update validation: budget shapes must parse, and no + two entries may refer to the same model once the provider prefix is stripped, + since the entry name is the canonical spend-counter key and duplicates would + split one model's spend across counters. + """ + if model_max_budget is None: + return + + from litellm.proxy.hooks.model_max_budget_limiter import ( + _PROXY_VirtualKeyModelMaxBudgetLimiter, + ) + from litellm.proxy.management_endpoints.key_management_endpoints import ( + validate_model_max_budget, + ) + + try: + validate_model_max_budget(model_max_budget) + except ValueError as e: + raise ProxyException( + message=str(e), + type=ProxyErrorTypes.bad_request_error, + param="model_max_budget", + code="400", + ) + + normalized_names = tuple( + _PROXY_VirtualKeyModelMaxBudgetLimiter._get_model_without_custom_llm_provider(entry_name) + for entry_name in model_max_budget + ) + colliding_entries = tuple( + sorted( + entry_name + for entry_name, normalized_name in zip(model_max_budget, normalized_names) + if normalized_names.count(normalized_name) > 1 + ) + ) + if colliding_entries: + raise ProxyException( + message=( + f"model_max_budget entries {colliding_entries} refer to the same model " + "after the provider prefix is stripped; keep one entry per model so spend " + "accrues to a single counter" + ), + type=ProxyErrorTypes.bad_request_error, + param="model_max_budget", + code="400", + ) + + def _should_auto_add_team_creator( user_api_key_dict: UserAPIKeyAuth, general_settings: Mapping[str, object], @@ -1233,20 +1286,7 @@ async def new_team( }, ) - if data.model_max_budget is not None: - from litellm.proxy.management_endpoints.key_management_endpoints import ( - validate_model_max_budget, - ) - - try: - validate_model_max_budget(data.model_max_budget) - except ValueError as e: - raise ProxyException( - message=str(e), - type=ProxyErrorTypes.bad_request_error, - param="model_max_budget", - code="400", - ) + _validate_team_model_max_budget(data.model_max_budget) # Check if license is over limit total_teams = await _team_db(prisma_client).count() @@ -2021,20 +2061,7 @@ async def update_team( existing_model_max_budget=existing_team_row.model_max_budget, ) - if data.model_max_budget is not None: - from litellm.proxy.management_endpoints.key_management_endpoints import ( - validate_model_max_budget, - ) - - try: - validate_model_max_budget(data.model_max_budget) - except ValueError as e: - raise ProxyException( - message=str(e), - type=ProxyErrorTypes.bad_request_error, - param="model_max_budget", - code="400", - ) + _validate_team_model_max_budget(data.model_max_budget) updated_kv = data.json(exclude_unset=True) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 68d75384452..1a77fd32d0d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3698,8 +3698,9 @@ class PrismaClient: sql_query = """ SELECT v.*, - t.spend AS team_spend, + t.spend AS team_spend, t.max_budget AS team_max_budget, + t.model_max_budget AS team_model_max_budget, t.soft_budget AS team_soft_budget, t.tpm_limit AS team_tpm_limit, t.rpm_limit AS team_rpm_limit, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index affaaa3fbf4..2a6dcbe2c38 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -3726,6 +3726,17 @@ async def test_centralized_common_checks_http_exception_without_team_id(): setattr(_proxy_server_mod, k, v) +def test_team_obj_from_token_preserves_model_max_budget(): + """The DB-fallback team object must carry model_max_budget: dropping it made + team model caps silently unenforced (and untracked) for the whole window a + team fetch kept failing (Veria finding).""" + from litellm.proxy.auth.user_api_key_auth import _team_obj_from_token + + budget = {"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}} + token = UserAPIKeyAuth(api_key="sk-test", team_id="team-1", team_model_max_budget=budget) + assert _team_obj_from_token(token).model_max_budget == budget + + @pytest.mark.asyncio async def test_centralized_common_checks_team_404_does_not_zero_other_contexts(): """Per-fetch isolation: an HTTPException(404) from get_team_object diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index dad650bfd67..10f2e932d95 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -10742,3 +10742,49 @@ class TestTeamModelMaxBudgetUpdateAuthority: user_api_key_dict=self._team_admin(), existing_model_max_budget=self._existing(), ) + + +class TestValidateTeamModelMaxBudget: + """_validate_team_model_max_budget gates /team/new and /team/update: budget + shapes must parse and no two entries may collapse to one model after the + provider prefix is stripped, else spend splits across counters and the team + can consume up to every duplicate's cap (Greptile finding).""" + + def _validate(self, model_max_budget): + from litellm.proxy.management_endpoints.team_endpoints import ( + _validate_team_model_max_budget, + ) + + _validate_team_model_max_budget(model_max_budget) + + def test_normalized_duplicate_entries_rejected(self): + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc_info: + self._validate( + { + "openai/gpt-4": {"budget_limit": 5.0, "time_period": "1d"}, + "gpt-4": {"budget_limit": 50.0, "time_period": "1d"}, + } + ) + assert exc_info.value.code == "400" + assert "openai/gpt-4" in exc_info.value.message + assert "gpt-4" in exc_info.value.message + + def test_distinct_models_accepted(self): + self._validate( + { + "openai/gpt-4": {"budget_limit": 5.0, "time_period": "1d"}, + "claude-3": {"budget_limit": 50.0, "time_period": "1d"}, + } + ) + + def test_none_accepted(self): + self._validate(None) + + def test_invalid_shape_still_rejected(self): + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc_info: + self._validate({"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}}) + assert exc_info.value.code == "400"