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 217ae0fb32f..f1a68112d03 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 @@ -37,6 +37,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_project_key_limits, _check_team_key_limits, _collect_key_team_limit_warnings, + _maybe_add_key_team_limit_warnings, _common_key_generation_helper, _enforce_upperbound_key_params, _get_and_validate_existing_key, @@ -53,6 +54,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( check_team_key_model_specific_limits, delete_verification_tokens, generate_key_fn, + generate_service_account_key_fn, generate_key_helper_fn, key_aliases, key_generation_check, @@ -3394,6 +3396,127 @@ async def test_validate_update_key_data_returns_team_limit_warnings(monkeypatch) ) +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( + team_id="warn-team", + team_alias="warn-team", + rpm_limit=60, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[], + ) + payload = {"key": "sk-test-key-123456", "rpm_limit": 600} + + assert ( + _maybe_add_key_team_limit_warnings( + payload, + UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=600), + None, + ) + is payload + ) + assert ( + _maybe_add_key_team_limit_warnings( + payload, + UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=30), + team_table, + ) + is payload + ) + + with_warnings = _maybe_add_key_team_limit_warnings( + payload, + UpdateKeyRequest(key="sk-test-key-123456", rpm_limit=600), + team_table, + ) + assert with_warnings is not payload + assert with_warnings["key"] == "sk-test-key-123456" + assert with_warnings["warnings"] == [ + { + "field": "rpm_limit", + "requested": 600, + "effective_team_cap": 60, + } + ] + + +@pytest.mark.asyncio +async def test_generate_service_account_key_fn_attaches_team_limit_warnings(monkeypatch): + """Service-account generate also surfaces team-limit warnings.""" + 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="sa-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.check_org_admin_can_generate_keys", + AsyncMock(), + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_team_id_used_in_service_account_request", + AsyncMock(), + ) + 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.key_generation_check", + MagicMock(), + ) + 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-service-account-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_service_account_key_fn( + data=data, user_api_key_dict=user_api_key_dict + ) + + assert result.key == "sk-service-account-123456" + assert result.warnings == [ + { + "field": "rpm_limit", + "requested": 600, + "effective_team_cap": 60, + } + ] + + @pytest.mark.asyncio async def test_check_team_key_limits_no_existing_keys(): """