From 9c890f658d3e63a554e0dfae341705e7a261c68d Mon Sep 17 00:00:00 2001 From: lei_lei Date: Sun, 13 Sep 2026 10:06:40 +0000 Subject: [PATCH] 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 --- litellm/proxy/_types.py | 12 + .../key_management_endpoints.py | 66 +++++- .../test_key_management_endpoints.py | 217 ++++++++++++++++++ 3 files changed, 288 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b315a2beac9..767e5d7b5ff 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 95ccb7bbe0b..a66fc35d443 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 2ac52da57df..217ae0fb32f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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(): """