diff --git a/test-quality-budget.json b/test-quality-budget.json index 6428a55ba787..91e8f39d195c 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -6,13 +6,13 @@ "limit": 742 }, "TQ003": { - "limit": 1075 + "limit": 1074 }, "TQ004": { "limit": 469 }, "TQ005": { - "limit": 2661 + "limit": 2621 }, "TQ006": { "limit": 34 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 fff6368cfc64..0c615cbaa327 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 @@ -1,16 +1,10 @@ import json -import os -import sys import litellm import pytest import yaml from fastapi.testclient import TestClient -sys.path.insert( - 0, os.path.abspath("../../../..") -) # Adds the parent directory to the system path - from unittest.mock import AsyncMock, MagicMock, patch from fastapi import HTTPException @@ -8125,26 +8119,22 @@ async def test_default_key_generate_params_duration(monkeypatch): monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Set default_key_generate_params with duration - original_value = litellm.default_key_generate_params - litellm.default_key_generate_params = {"duration": "180d"} + monkeypatch.setattr(litellm, "default_key_generate_params", {"duration": "180d"}) - try: - request = GenerateKeyRequest() # No duration specified - response = await _common_key_generation_helper( - data=request, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key="sk-1234", - user_id="1234", - ), - litellm_changed_by=None, - team_table=None, - ) + request = GenerateKeyRequest() # No duration specified + response = await _common_key_generation_helper( + data=request, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="1234", + ), + litellm_changed_by=None, + team_table=None, + ) - # Verify duration was applied from defaults - assert request.duration == "180d" - finally: - litellm.default_key_generate_params = original_value + # Verify duration was applied from defaults + assert request.duration == "180d" async def test_default_key_generate_params_object_permission_applied_when_absent( @@ -8184,28 +8174,28 @@ async def test_default_key_generate_params_object_permission_applied_when_absent monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - original_value = litellm.default_key_generate_params - litellm.default_key_generate_params = { - "object_permission": {"vector_stores": ["default-vs"]} - } + monkeypatch.setattr( + litellm, + "default_key_generate_params", + { + "object_permission": {"vector_stores": ["default-vs"]} + }, + ) - try: - request = GenerateKeyRequest() # No object_permission specified - await _common_key_generation_helper( - data=request, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key="sk-1234", - user_id="1234", - ), - litellm_changed_by=None, - team_table=None, - ) + request = GenerateKeyRequest() # No object_permission specified + await _common_key_generation_helper( + data=request, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="1234", + ), + litellm_changed_by=None, + team_table=None, + ) - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] - assert created_data["vector_stores"] == ["default-vs"] - finally: - litellm.default_key_generate_params = original_value + created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + assert created_data["vector_stores"] == ["default-vs"] async def test_default_key_generate_params_object_permission_merges_partial( @@ -8247,31 +8237,31 @@ async def test_default_key_generate_params_object_permission_merges_partial( monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - original_value = litellm.default_key_generate_params - litellm.default_key_generate_params = { - "object_permission": {"vector_stores": ["default-vs"]} - } + monkeypatch.setattr( + litellm, + "default_key_generate_params", + { + "object_permission": {"vector_stores": ["default-vs"]} + }, + ) - try: - request = GenerateKeyRequest( - object_permission=LiteLLM_ObjectPermissionBase(agents=["agent-1"]) - ) - await _common_key_generation_helper( - data=request, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key="sk-1234", - user_id="1234", - ), - litellm_changed_by=None, - team_table=None, - ) + request = GenerateKeyRequest( + object_permission=LiteLLM_ObjectPermissionBase(agents=["agent-1"]) + ) + await _common_key_generation_helper( + data=request, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="1234", + ), + litellm_changed_by=None, + team_table=None, + ) - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] - assert created_data["agents"] == ["agent-1"] - assert created_data["vector_stores"] == ["default-vs"] - finally: - litellm.default_key_generate_params = original_value + created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + assert created_data["agents"] == ["agent-1"] + assert created_data["vector_stores"] == ["default-vs"] async def test_default_key_generate_params_object_permission_does_not_override_explicit( @@ -8312,32 +8302,32 @@ async def test_default_key_generate_params_object_permission_does_not_override_e monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - original_value = litellm.default_key_generate_params - litellm.default_key_generate_params = { - "object_permission": {"vector_stores": ["default-vs"]} - } + monkeypatch.setattr( + litellm, + "default_key_generate_params", + { + "object_permission": {"vector_stores": ["default-vs"]} + }, + ) - try: - request = GenerateKeyRequest( - object_permission=LiteLLM_ObjectPermissionBase( - vector_stores=["explicit-vs"] - ) - ) - await _common_key_generation_helper( - data=request, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key="sk-1234", - user_id="1234", - ), - litellm_changed_by=None, - team_table=None, + request = GenerateKeyRequest( + object_permission=LiteLLM_ObjectPermissionBase( + vector_stores=["explicit-vs"] ) + ) + await _common_key_generation_helper( + data=request, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="1234", + ), + litellm_changed_by=None, + team_table=None, + ) - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] - assert created_data["vector_stores"] == ["explicit-vs"] - finally: - litellm.default_key_generate_params = original_value + created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + assert created_data["vector_stores"] == ["explicit-vs"] async def test_default_key_generate_params_object_permission_not_rejected_for_non_admin_personal_key( @@ -8380,29 +8370,29 @@ async def test_default_key_generate_params_object_permission_not_rejected_for_no monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - original_value = litellm.default_key_generate_params - litellm.default_key_generate_params = { - "object_permission": {"vector_stores": ["default-vs"]} - } + monkeypatch.setattr( + litellm, + "default_key_generate_params", + { + "object_permission": {"vector_stores": ["default-vs"]} + }, + ) - try: - request = GenerateKeyRequest(user_id="alice") # No object_permission specified - response = await _common_key_generation_helper( - data=request, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, - api_key="sk-alice", - user_id="alice", - ), - litellm_changed_by=None, - team_table=None, - ) + request = GenerateKeyRequest(user_id="alice") # No object_permission specified + response = await _common_key_generation_helper( + data=request, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + litellm_changed_by=None, + team_table=None, + ) - assert response is not None - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] - assert created_data["vector_stores"] == ["default-vs"] - finally: - litellm.default_key_generate_params = original_value + assert response is not None + created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + assert created_data["vector_stores"] == ["default-vs"] @pytest.mark.asyncio @@ -9261,10 +9251,8 @@ async def test_key_aliases_admin_sees_all(): class TestValidateKeyAliasFormat: @pytest.fixture(autouse=True) - def reset_key_alias_flag(self): - litellm.enable_key_alias_format_validation = False - yield - litellm.enable_key_alias_format_validation = False + def reset_key_alias_flag(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "enable_key_alias_format_validation", False) def test_validation_skipped_when_flag_disabled(self): """When enable_key_alias_format_validation is False (default), no charset/length validation occurs.""" @@ -9305,12 +9293,12 @@ def test_validate_key_alias_format_rejects_traversal_and_control_chars_even_when assert str(exc.value.code) == "400" assert "Invalid key_alias" in str(exc.value.message) - def test_validate_key_alias_format_valid(self): + def test_validate_key_alias_format_valid(self, monkeypatch): from litellm.proxy.management_endpoints.key_management_endpoints import ( _validate_key_alias_format, ) - litellm.enable_key_alias_format_validation = True + monkeypatch.setattr(litellm, "enable_key_alias_format_validation", True) # Valid cases _validate_key_alias_format(None) # OK _validate_key_alias_format("valid-alias") @@ -9322,13 +9310,13 @@ def test_validate_key_alias_format_valid(self): _validate_key_alias_format("user/user@example.com") _validate_key_alias_format("team/user@example.com") - def test_validate_key_alias_format_invalid(self): + def test_validate_key_alias_format_invalid(self, monkeypatch): from litellm.proxy.management_endpoints.key_management_endpoints import ( _validate_key_alias_format, ) from litellm.proxy._types import ProxyException - litellm.enable_key_alias_format_validation = True + monkeypatch.setattr(litellm, "enable_key_alias_format_validation", True) invalid_aliases = [ "", # empty " ", # whitespace @@ -10956,10 +10944,8 @@ class TestKeyAliasSkipValidationOnUnchanged: """ @pytest.fixture(autouse=True) - def enable_validation(self): - litellm.enable_key_alias_format_validation = True - yield - litellm.enable_key_alias_format_validation = False + def enable_validation(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "enable_key_alias_format_validation", True) @pytest.fixture def mock_prisma(self): @@ -11075,146 +11061,142 @@ async def test_update_key_alias_none_skips_validation(self): # --- Tests: _enforce_upperbound_key_params --- -def test_enforce_upperbound_rejects_over_limit_on_generate(): +def test_enforce_upperbound_rejects_over_limit_on_generate(monkeypatch): """Test that key generation is rejected when values exceed upperbound.""" import litellm from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - tpm_limit=1000, rpm_limit=100, max_budget=10.0 - ) - data = GenerateKeyRequest(tpm_limit=5000) - with pytest.raises(HTTPException) as exc_info: - _enforce_upperbound_key_params(data, fill_defaults=True) - assert exc_info.value.status_code == 400 - assert "tpm_limit" in str(exc_info.value.detail) - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + tpm_limit=1000, rpm_limit=100, max_budget=10.0 + ), + ) + data = GenerateKeyRequest(tpm_limit=5000) + with pytest.raises(HTTPException) as exc_info: + _enforce_upperbound_key_params(data, fill_defaults=True) + assert exc_info.value.status_code == 400 + assert "tpm_limit" in str(exc_info.value.detail) -def test_enforce_upperbound_fills_defaults_on_generate(): +def test_enforce_upperbound_fills_defaults_on_generate(monkeypatch): """Test that None values are filled with upperbound defaults during generation.""" import litellm from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - tpm_limit=1000, rpm_limit=100 - ) - data = GenerateKeyRequest() # tpm_limit=None, rpm_limit=None - _enforce_upperbound_key_params(data, fill_defaults=True) - assert data.tpm_limit == 1000 - assert data.rpm_limit == 100 - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + tpm_limit=1000, rpm_limit=100 + ), + ) + data = GenerateKeyRequest() # tpm_limit=None, rpm_limit=None + _enforce_upperbound_key_params(data, fill_defaults=True) + assert data.tpm_limit == 1000 + assert data.rpm_limit == 100 -def test_enforce_upperbound_skips_none_on_update(): +def test_enforce_upperbound_skips_none_on_update(monkeypatch): """Test that None values are NOT filled during update (fill_defaults=False).""" import litellm from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - tpm_limit=1000, rpm_limit=100 - ) - data = UpdateKeyRequest(key="sk-test") # tpm_limit=None, rpm_limit=None - _enforce_upperbound_key_params(data, fill_defaults=False) - assert data.tpm_limit is None # should NOT be filled - assert data.rpm_limit is None # should NOT be filled - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + tpm_limit=1000, rpm_limit=100 + ), + ) + data = UpdateKeyRequest(key="sk-test") # tpm_limit=None, rpm_limit=None + _enforce_upperbound_key_params(data, fill_defaults=False) + assert data.tpm_limit is None # should NOT be filled + assert data.rpm_limit is None # should NOT be filled -def test_enforce_upperbound_rejects_over_limit_on_update(): +def test_enforce_upperbound_rejects_over_limit_on_update(monkeypatch): """Test that key update is rejected when values exceed upperbound.""" import litellm from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - tpm_limit=1000, rpm_limit=100, max_budget=10.0 - ) - data = UpdateKeyRequest(key="sk-test", tpm_limit=5000) - with pytest.raises(HTTPException) as exc_info: - _enforce_upperbound_key_params(data, fill_defaults=False) - assert exc_info.value.status_code == 400 - assert "tpm_limit" in str(exc_info.value.detail) - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + tpm_limit=1000, rpm_limit=100, max_budget=10.0 + ), + ) + data = UpdateKeyRequest(key="sk-test", tpm_limit=5000) + with pytest.raises(HTTPException) as exc_info: + _enforce_upperbound_key_params(data, fill_defaults=False) + assert exc_info.value.status_code == 400 + assert "tpm_limit" in str(exc_info.value.detail) -def test_enforce_upperbound_allows_within_limit_on_update(): +def test_enforce_upperbound_allows_within_limit_on_update(monkeypatch): """Test that key update passes when values are within upperbound.""" import litellm from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - tpm_limit=1000, rpm_limit=100, max_budget=10.0 - ) - data = UpdateKeyRequest( - key="sk-test", tpm_limit=500, rpm_limit=50, max_budget=5.0 - ) - _enforce_upperbound_key_params(data, fill_defaults=False) - # Should not raise - assert data.tpm_limit == 500 - assert data.rpm_limit == 50 - assert data.max_budget == 5.0 - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + tpm_limit=1000, rpm_limit=100, max_budget=10.0 + ), + ) + data = UpdateKeyRequest( + key="sk-test", tpm_limit=500, rpm_limit=50, max_budget=5.0 + ) + _enforce_upperbound_key_params(data, fill_defaults=False) + # Should not raise + assert data.tpm_limit == 500 + assert data.rpm_limit == 50 + assert data.max_budget == 5.0 -def test_enforce_upperbound_duration_over_limit(): +def test_enforce_upperbound_duration_over_limit(monkeypatch): """Test that duration exceeding upperbound is rejected.""" import litellm from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - duration="7d" - ) - data = UpdateKeyRequest(key="sk-test", duration="30d") - with pytest.raises(HTTPException) as exc_info: - _enforce_upperbound_key_params(data, fill_defaults=False) - assert exc_info.value.status_code == 400 - assert "duration" in str(exc_info.value.detail) - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + duration="7d" + ), + ) + data = UpdateKeyRequest(key="sk-test", duration="30d") + with pytest.raises(HTTPException) as exc_info: + _enforce_upperbound_key_params(data, fill_defaults=False) + assert exc_info.value.status_code == 400 + assert "duration" in str(exc_info.value.detail) -def test_enforce_upperbound_no_config_is_noop(): +def test_enforce_upperbound_no_config_is_noop(monkeypatch): """Test that no enforcement happens when upperbound params are not configured.""" import litellm - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = None - data = UpdateKeyRequest(key="sk-test", tpm_limit=999999) - _enforce_upperbound_key_params(data, fill_defaults=False) - # Should not raise — no enforcement configured - assert data.tpm_limit == 999999 - finally: - litellm.upperbound_key_generate_params = original + monkeypatch.setattr(litellm, "upperbound_key_generate_params", None) + data = UpdateKeyRequest(key="sk-test", tpm_limit=999999) + _enforce_upperbound_key_params(data, fill_defaults=False) + # Should not raise — no enforcement configured + assert data.tpm_limit == 999999 # --- Tests: _execute_virtual_key_regeneration enforces upperbound --- @@ -11267,7 +11249,7 @@ def _make_regenerate_existing_key(): @pytest.mark.asyncio -async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(): +async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monkeypatch): """Regenerate must reject durations exceeding upperbound_key_generate_params.duration.""" from litellm.proxy._types import RegenerateKeyRequest from litellm.proxy.management_endpoints.key_management_endpoints import ( @@ -11277,91 +11259,34 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(): LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - duration="1h" - ) - existing_key = _make_regenerate_existing_key() - data = RegenerateKeyRequest(duration="2h") - user_api_key_dict = _make_regenerate_user_api_key_dict() - mock_prisma_client = _make_regenerate_mock_prisma() - - with ( - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", - new_callable=AsyncMock, - return_value="sk-newtoken1234ab12", - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", - new_callable=AsyncMock, - ), - ): - with pytest.raises(HTTPException) as exc_info: - await _execute_virtual_key_regeneration( - prisma_client=mock_prisma_client, - key_in_db=existing_key, - hashed_api_key="abc123", - key="abc123", - data=data, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=None, - user_api_key_cache=MagicMock(), - proxy_logging_obj=MagicMock(), - ) - assert exc_info.value.status_code == 400 - assert "duration" in str(exc_info.value.detail) - # Rejected regenerate must not reach the DB update. - assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 - finally: - litellm.upperbound_key_generate_params = original - - -@pytest.mark.asyncio -async def test_execute_virtual_key_regeneration_allows_within_limit_duration(): - """Regenerate must accept durations within upperbound_key_generate_params.duration.""" - from litellm.proxy._types import RegenerateKeyRequest - from litellm.proxy.management_endpoints.key_management_endpoints import ( - _execute_virtual_key_regeneration, - ) - from litellm.types.proxy.management_endpoints.ui_sso import ( - LiteLLM_UpperboundKeyGenerateParams, + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + duration="1h" + ), ) + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest(duration="2h") + user_api_key_dict = _make_regenerate_user_api_key_dict() + mock_prisma_client = _make_regenerate_mock_prisma() - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - duration="1h" - ) - existing_key = _make_regenerate_existing_key() - data = RegenerateKeyRequest(duration="30m") - user_api_key_dict = _make_regenerate_user_api_key_dict() - mock_prisma_client = _make_regenerate_mock_prisma() - - with ( - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", - new_callable=AsyncMock, - return_value="sk-newtoken1234ab12", - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", - new_callable=AsyncMock, - ), - ): + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + ): + with pytest.raises(HTTPException) as exc_info: await _execute_virtual_key_regeneration( prisma_client=mock_prisma_client, key_in_db=existing_key, @@ -11373,14 +11298,15 @@ async def test_execute_virtual_key_regeneration_allows_within_limit_duration(): user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), ) - assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 - finally: - litellm.upperbound_key_generate_params = original + assert exc_info.value.status_code == 400 + assert "duration" in str(exc_info.value.detail) + # Rejected regenerate must not reach the DB update. + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 @pytest.mark.asyncio -async def test_execute_virtual_key_regeneration_rejects_over_limit_max_budget(): - """Regenerate must reject max_budget exceeding upperbound — proves the fix covers non-duration fields.""" +async def test_execute_virtual_key_regeneration_allows_within_limit_duration(monkeypatch): + """Regenerate must accept durations within upperbound_key_generate_params.duration.""" from litellm.proxy._types import RegenerateKeyRequest from litellm.proxy.management_endpoints.key_management_endpoints import ( _execute_virtual_key_regeneration, @@ -11389,54 +11315,54 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_max_budget(): LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - max_budget=10.0 - ) - existing_key = _make_regenerate_existing_key() - data = RegenerateKeyRequest(max_budget=500.0) - user_api_key_dict = _make_regenerate_user_api_key_dict() - mock_prisma_client = _make_regenerate_mock_prisma() + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + duration="1h" + ), + ) + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest(duration="30m") + user_api_key_dict = _make_regenerate_user_api_key_dict() + mock_prisma_client = _make_regenerate_mock_prisma() - with ( - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", - new_callable=AsyncMock, - return_value="sk-newtoken1234ab12", - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", - new_callable=AsyncMock, - ), - ): - with pytest.raises(HTTPException) as exc_info: - await _execute_virtual_key_regeneration( - prisma_client=mock_prisma_client, - key_in_db=existing_key, - hashed_api_key="abc123", - key="abc123", - data=data, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=None, - user_api_key_cache=MagicMock(), - proxy_logging_obj=MagicMock(), - ) - assert exc_info.value.status_code == 400 - assert "max_budget" in str(exc_info.value.detail) - assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 - finally: - litellm.upperbound_key_generate_params = original + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 @pytest.mark.asyncio -async def test_execute_virtual_key_regeneration_skips_none_values(): - """Regenerate with data.duration=None must not raise, even when upperbound is set - (fill_defaults=False semantic — None means 'inherit from existing key').""" +async def test_execute_virtual_key_regeneration_rejects_over_limit_max_budget(monkeypatch): + """Regenerate must reject max_budget exceeding upperbound — proves the fix covers non-duration fields.""" from litellm.proxy._types import RegenerateKeyRequest from litellm.proxy.management_endpoints.key_management_endpoints import ( _execute_virtual_key_regeneration, @@ -11445,35 +11371,34 @@ async def test_execute_virtual_key_regeneration_skips_none_values(): LiteLLM_UpperboundKeyGenerateParams, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = LiteLLM_UpperboundKeyGenerateParams( - duration="1h" - ) - existing_key = _make_regenerate_existing_key() - data = RegenerateKeyRequest() # all fields None - user_api_key_dict = _make_regenerate_user_api_key_dict() - mock_prisma_client = _make_regenerate_mock_prisma() + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + max_budget=10.0 + ), + ) + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest(max_budget=500.0) + user_api_key_dict = _make_regenerate_user_api_key_dict() + mock_prisma_client = _make_regenerate_mock_prisma() - with ( - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", - new_callable=AsyncMock, - return_value="sk-newtoken1234ab12", - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", - new_callable=AsyncMock, - ), - ): + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + ): + with pytest.raises(HTTPException) as exc_info: await _execute_virtual_key_regeneration( prisma_client=mock_prisma_client, key_in_db=existing_key, @@ -11485,60 +11410,113 @@ async def test_execute_virtual_key_regeneration_skips_none_values(): user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), ) - assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 - finally: - litellm.upperbound_key_generate_params = original + assert exc_info.value.status_code == 400 + assert "max_budget" in str(exc_info.value.detail) + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 + + +@pytest.mark.asyncio +async def test_execute_virtual_key_regeneration_skips_none_values(monkeypatch): + """Regenerate with data.duration=None must not raise, even when upperbound is set + (fill_defaults=False semantic — None means 'inherit from existing key').""" + from litellm.proxy._types import RegenerateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _execute_virtual_key_regeneration, + ) + from litellm.types.proxy.management_endpoints.ui_sso import ( + LiteLLM_UpperboundKeyGenerateParams, + ) + + monkeypatch.setattr( + litellm, + "upperbound_key_generate_params", + LiteLLM_UpperboundKeyGenerateParams( + duration="1h" + ), + ) + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest() # all fields None + user_api_key_dict = _make_regenerate_user_api_key_dict() + mock_prisma_client = _make_regenerate_mock_prisma() + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 @pytest.mark.asyncio -async def test_execute_virtual_key_regeneration_no_upperbound_config_is_noop(): +async def test_execute_virtual_key_regeneration_no_upperbound_config_is_noop(monkeypatch): """Regenerate with no upperbound config set must accept any duration.""" from litellm.proxy._types import RegenerateKeyRequest from litellm.proxy.management_endpoints.key_management_endpoints import ( _execute_virtual_key_regeneration, ) - original = litellm.upperbound_key_generate_params - try: - litellm.upperbound_key_generate_params = None - existing_key = _make_regenerate_existing_key() - data = RegenerateKeyRequest(duration="30d") - user_api_key_dict = _make_regenerate_user_api_key_dict() - mock_prisma_client = _make_regenerate_mock_prisma() + monkeypatch.setattr(litellm, "upperbound_key_generate_params", None) + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest(duration="30d") + user_api_key_dict = _make_regenerate_user_api_key_dict() + mock_prisma_client = _make_regenerate_mock_prisma() - with ( - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", - new_callable=AsyncMock, - return_value="sk-newtoken1234ab12", - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", - new_callable=AsyncMock, - ), - ): - await _execute_virtual_key_regeneration( - prisma_client=mock_prisma_client, - key_in_db=existing_key, - hashed_api_key="abc123", - key="abc123", - data=data, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=None, - user_api_key_cache=MagicMock(), - proxy_logging_obj=MagicMock(), - ) - assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 - finally: - litellm.upperbound_key_generate_params = original + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 class TestAllowedRoutesCallerPermission: