fix(team): reject normalized-duplicate budget entries and keep caps on team-fetch fallback

A team model_max_budget containing both bare and provider-prefixed entries for
one model let each request spelling match its own entry and counter, so the
team could consume up to every duplicate's cap. /team/new and /team/update now
share one validator that rejects entries colliding after provider-prefix
normalization, keeping the entry name a unique canonical counter key.

The combined token view now selects the team's model_max_budget and
_team_obj_from_token carries it into the DB-fallback team object, so a failing
team fetch no longer silently disables team model caps and their spend
tracking.
This commit is contained in:
ryan-crabbe-berri 2026-08-03 14:02:18 -07:00
parent 5cc8c1b1a2
commit a411f0d75a
6 changed files with 117 additions and 30 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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