fix(proxy): warn on retained key limits during team reassignment

/key/update often omits rpm/tpm/concurrency/budget when only team_id
changes. Compare those retained values against the effective team caps
so reassignment cannot silently keep an over-cap limit.

Addresses greptile P1 on #40998.
This commit is contained in:
leilei3167 2026-09-16 15:09:21 +00:00
parent 15e8e58394
commit 1c69fe21ca
2 changed files with 107 additions and 1 deletions

View file

@ -1501,6 +1501,34 @@ def check_team_key_rpm_tpm_limits(
)
_KEY_TEAM_LIMIT_WARNING_FIELDS: Final[tuple[str, ...]] = (
"rpm_limit",
"tpm_limit",
"max_parallel_requests",
"max_budget",
)
def _update_request_with_retained_team_limits(
data: UpdateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
) -> UpdateKeyRequest:
"""Fill omitted limit fields from the existing key for team-cap warnings.
/key/update often changes only team_id (or a subset of limits). Retained
rpm/tpm/concurrency/budget must still be compared against the effective
team caps so reassignment cannot silently keep an over-cap value.
"""
retained: dict[str, object] = {}
for field_name in _KEY_TEAM_LIMIT_WARNING_FIELDS:
if field_name in data.model_fields_set:
continue
retained[field_name] = getattr(existing_key_row, field_name, None)
if not retained:
return data
return data.model_copy(update=retained)
def _collect_key_team_limit_warnings(
data: GenerateKeyRequest | UpdateKeyRequest,
team_table: LiteLLM_TeamTable | LiteLLM_TeamTableCachedObj,
@ -3006,7 +3034,10 @@ async def _validate_update_key_data(
if team_obj is None:
return ()
return _collect_key_team_limit_warnings(data=data, team_table=team_obj)
return _collect_key_team_limit_warnings(
data=_update_request_with_retained_team_limits(data=data, existing_key_row=existing_key_row),
team_table=team_obj,
)
@router.post("/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)])

View file

@ -3396,6 +3396,81 @@ async def test_validate_update_key_data_returns_team_limit_warnings(monkeypatch)
)
async def test_validate_update_key_data_warns_on_retained_limits_team_change(monkeypatch):
"""Team reassignment without limit fields still warns on retained over-cap values."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import (
KeyTeamLimitWarning,
LiteLLM_TeamTableCachedObj,
LiteLLM_VerificationToken,
LitellmUserRoles,
UpdateKeyRequest,
UserAPIKeyAuth,
)
existing = LiteLLM_VerificationToken(
token="hashed-token",
team_id="team-old",
user_id="user-1",
max_budget=100.0,
max_parallel_requests=50,
rpm_limit=None,
tpm_limit=None,
)
new_team = LiteLLM_TeamTableCachedObj(
team_id="team-new",
max_budget=10.0,
max_parallel_requests=5,
rpm_limit=None,
tpm_limit=None,
)
data = UpdateKeyRequest(key="sk-test", team_id="team-new")
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
AsyncMock(return_value=new_team),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._check_team_key_limits",
AsyncMock(return_value=None),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
AsyncMock(return_value=None),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.enforce_member_can_assign_access_groups",
MagicMock(return_value=None),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_team_change",
AsyncMock(return_value=None),
)
# Avoid router requirement inside team-change validation path when mocked above
user = UserAPIKeyAuth(
user_id="user-1",
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
)
warnings = await _validate_update_key_data(
data=data,
existing_key_row=existing,
user_api_key_dict=user,
llm_router=MagicMock(),
premium_user=True,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert KeyTeamLimitWarning(
field="max_budget", requested=100.0, effective_team_cap=10.0
) in warnings
assert KeyTeamLimitWarning(
field="max_parallel_requests", requested=50, effective_team_cap=5
) in warnings
def test_maybe_add_key_team_limit_warnings_passthrough_and_attach():
"""Attach warnings to update payloads only when caps are exceeded."""
team_table = LiteLLM_TeamTableCachedObj(