fix(proxy): warn when key limits exceed team caps on generate/update

Compare rpm/tpm/max_parallel_requests/max_budget against the team on
/key/generate and /key/update, and return field-level warnings when the
key asks higher than the team cap, without rejecting the write yet
This commit is contained in:
lei_lei 2026-09-13 10:06:40 +00:00
parent 30f33a949b
commit 9c890f658d
No known key found for this signature in database
3 changed files with 288 additions and 7 deletions

View file

@ -1190,6 +1190,17 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
return v
KeyTeamLimitField = Literal["rpm_limit", "tpm_limit", "max_parallel_requests", "max_budget"]
class KeyTeamLimitWarning(TypedDict):
"""Non-blocking warning when a key limit exceeds its team's effective cap."""
field: ReadOnly[KeyTeamLimitField]
requested: ReadOnly[float | int]
effective_team_cap: ReadOnly[float | int]
class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase):
index_name: str
index_permissions: list[Literal["read", "write"]]
@ -1267,6 +1278,7 @@ class GenerateKeyResponse(KeyRequestBase):
updated_by: str | None = None
created_at: datetime | None = None
updated_at: datetime | None = None
warnings: list[KeyTeamLimitWarning] | None = None
@model_validator(mode="before")
@classmethod

View file

@ -1501,6 +1501,33 @@ def check_team_key_rpm_tpm_limits(
)
def _collect_key_team_limit_warnings(
data: GenerateKeyRequest | UpdateKeyRequest,
team_table: LiteLLM_TeamTable | LiteLLM_TeamTableCachedObj,
) -> tuple[KeyTeamLimitWarning, ...]:
"""
Compare key rpm/tpm/max_parallel_requests/max_budget against the team's caps.
Returns warnings when the key requests a higher value than the team allows.
Does not reject; runtime still applies the stricter team limit.
"""
comparisons: Final[tuple[tuple[KeyTeamLimitField, float | int | None, float | int | None], ...]] = (
("rpm_limit", data.rpm_limit, team_table.rpm_limit),
("tpm_limit", data.tpm_limit, team_table.tpm_limit),
("max_parallel_requests", data.max_parallel_requests, team_table.max_parallel_requests),
("max_budget", data.max_budget, team_table.max_budget),
)
return tuple(
KeyTeamLimitWarning(
field=field_name,
requested=requested,
effective_team_cap=team_cap,
)
for field_name, requested, team_cap in comparisons
if requested is not None and team_cap is not None and requested > team_cap
)
async def _check_team_key_limits(
team_table: LiteLLM_TeamTableCachedObj,
data: GenerateKeyRequest | UpdateKeyRequest,
@ -1924,12 +1951,18 @@ async def generate_key_fn(
user_api_key_cache=user_api_key_cache,
)
return await _common_key_generation_helper(
team_limit_warnings: Final = (
_collect_key_team_limit_warnings(data=data, team_table=team_table) if team_table is not None else ()
)
response: Final = await _common_key_generation_helper(
data=data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
team_table=team_table,
)
if team_limit_warnings:
response.warnings = list(team_limit_warnings)
return response
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.generate_key_fn(): Exception occured - %s", e)
@ -2090,12 +2123,18 @@ async def generate_service_account_key_fn(
data.user_id = None # do not allow user_id to be set for service account keys
return await _common_key_generation_helper(
team_limit_warnings: Final = (
_collect_key_team_limit_warnings(data=data, team_table=team_table) if team_table is not None else ()
)
response: Final = await _common_key_generation_helper(
data=data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
team_table=team_table,
)
if team_limit_warnings:
response.warnings = list(team_limit_warnings)
return response
def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_metadata: dict) -> dict:
@ -2538,9 +2577,10 @@ async def _process_single_key_update(
# Get team object and check team limits if team_id is provided
team_obj: LiteLLM_TeamTableCachedObj | None = None
if update_key_request.team_id is not None:
_team_id_to_check: Final = update_key_request.team_id or getattr(existing_key_row, "team_id", None)
if _team_id_to_check is not None:
team_obj = await get_team_object(
team_id=update_key_request.team_id,
team_id=_team_id_to_check,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
@ -2631,6 +2671,11 @@ async def _process_single_key_update(
updated_key_info.pop("token", None)
team_limit_warnings: Final = (
_collect_key_team_limit_warnings(data=update_key_request, team_table=team_obj) if team_obj is not None else ()
)
if team_limit_warnings:
return {**updated_key_info, "warnings": list(team_limit_warnings)}
return updated_key_info
@ -2688,7 +2733,7 @@ async def _validate_update_key_data(
premium_user: bool,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
) -> None:
) -> tuple[KeyTeamLimitWarning, ...]:
"""Validate permissions and constraints for key update."""
checked_prisma_client: Final = _require_prisma_client(prisma_client)
@ -2946,6 +2991,10 @@ async def _validate_update_key_data(
if normalized_object_permission is not None:
data.object_permission = LiteLLM_ObjectPermissionBase(**normalized_object_permission)
if team_obj is None:
return ()
return _collect_key_team_limit_warnings(data=data, team_table=team_obj)
@router.post("/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)])
@management_endpoint_wrapper
@ -3067,7 +3116,7 @@ async def update_key_fn(
key: Final = _resolve_token_to_update(data=data, existing_key_row=existing_key_row)
data.key = key
await _validate_update_key_data(
team_limit_warnings: Final = await _validate_update_key_data(
data=data,
existing_key_row=existing_key_row,
user_api_key_dict=user_api_key_dict,
@ -3177,7 +3226,10 @@ async def update_key_fn(
if response is None:
raise ValueError("Failed to update key got response = None")
return {"key": key, **response["data"]}
updated_key_info: Final[dict[str, object]] = {"key": key, **response["data"]}
if team_limit_warnings:
return {**updated_key_info, "warnings": list(team_limit_warnings)}
return updated_key_info
# update based on remaining passed in values
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_key_fn(): Exception occured - %s", e)

View file

@ -36,6 +36,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_org_key_limits,
_check_project_key_limits,
_check_team_key_limits,
_collect_key_team_limit_warnings,
_common_key_generation_helper,
_enforce_upperbound_key_params,
_get_and_validate_existing_key,
@ -3177,6 +3178,222 @@ async def test_update_key_fn_auto_rotate_disable():
assert result["auto_rotate"] is False
def test_collect_key_team_limit_warnings_above_team_caps():
"""Key limits above team caps produce field-level warnings without rejection."""
team_table = LiteLLM_TeamTableCachedObj(
team_id="capped-team",
team_alias="capped-team",
rpm_limit=60,
tpm_limit=1000,
max_parallel_requests=5,
max_budget=10.0,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
data = GenerateKeyRequest(
team_id="capped-team",
rpm_limit=600,
tpm_limit=5000,
max_parallel_requests=20,
max_budget=100.0,
)
warnings = _collect_key_team_limit_warnings(data=data, team_table=team_table)
assert warnings == (
{
"field": "rpm_limit",
"requested": 600,
"effective_team_cap": 60,
},
{
"field": "tpm_limit",
"requested": 5000,
"effective_team_cap": 1000,
},
{
"field": "max_parallel_requests",
"requested": 20,
"effective_team_cap": 5,
},
{
"field": "max_budget",
"requested": 100.0,
"effective_team_cap": 10.0,
},
)
def test_collect_key_team_limit_warnings_within_or_unset_caps():
"""No warning when key is within team caps, team has no cap, or key omits the field."""
team_table = LiteLLM_TeamTableCachedObj(
team_id="partial-caps",
team_alias="partial-caps",
rpm_limit=60,
tpm_limit=None,
max_parallel_requests=5,
max_budget=None,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
data = UpdateKeyRequest(
key="sk-test-key-123456",
rpm_limit=60,
tpm_limit=999999,
max_parallel_requests=3,
max_budget=50.0,
)
warnings = _collect_key_team_limit_warnings(data=data, team_table=team_table)
assert warnings == ()
@pytest.mark.asyncio
async def test_generate_key_fn_attaches_team_limit_warnings(monkeypatch):
"""/key/generate succeeds and returns warnings when key rpm exceeds team rpm."""
team_table = LiteLLM_TeamTableCachedObj(
team_id="warn-team",
team_alias="warn-team",
rpm_limit=60,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
data = GenerateKeyRequest(team_id="warn-team", rpm_limit=600, key_alias="over-cap")
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-1234",
user_id="admin",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
AsyncMock(return_value=team_table),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.key_generation_check",
MagicMock(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._check_team_key_limits",
AsyncMock(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.check_org_admin_can_generate_keys",
AsyncMock(),
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
MagicMock(),
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache",
MagicMock(),
)
from litellm.proxy._types import GenerateKeyResponse
generated = GenerateKeyResponse(
key="sk-generated-key-123456",
token_id="hashed",
team_id="warn-team",
rpm_limit=600,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._common_key_generation_helper",
AsyncMock(return_value=generated),
)
result = await generate_key_fn(data=data, user_api_key_dict=user_api_key_dict)
assert result.key == "sk-generated-key-123456"
assert result.rpm_limit == 600
assert result.warnings == [
{
"field": "rpm_limit",
"requested": 600,
"effective_team_cap": 60,
}
]
@pytest.mark.asyncio
async def test_validate_update_key_data_returns_team_limit_warnings(monkeypatch):
"""/key/update validation returns warnings when updated limits exceed team caps."""
existing_key = LiteLLM_VerificationToken(
token="hashed-token",
team_id="warn-team",
user_id="user-1",
models=[],
)
team_table = LiteLLM_TeamTableCachedObj(
team_id="warn-team",
team_alias="warn-team",
rpm_limit=60,
max_budget=10.0,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
data = UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=600, max_budget=100.0)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-1234",
user_id="admin",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
AsyncMock(return_value=team_table),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._check_team_key_limits",
AsyncMock(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint",
AsyncMock(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.common_key_access_checks",
MagicMock(),
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.enforce_member_can_assign_access_groups",
MagicMock(),
)
warnings = await _validate_update_key_data(
data=data,
existing_key_row=existing_key,
user_api_key_dict=user_api_key_dict,
llm_router=None,
premium_user=True,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert warnings == (
{
"field": "rpm_limit",
"requested": 600,
"effective_team_cap": 60,
},
{
"field": "max_budget",
"requested": 100.0,
"effective_team_cap": 10.0,
},
)
@pytest.mark.asyncio
async def test_check_team_key_limits_no_existing_keys():
"""