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 a302443c3eb..64140cdc951 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 @@ -21098,10 +21098,15 @@ class TestTeamAdminMemberKeyBudgetUpdate: @pytest.mark.asyncio async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): """Master-key rotation re-encrypts the guardrails table's litellm_params under the new key.""" + import json + from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry + from litellm.proxy.guardrails.guardrail_registry import ( + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, + ) from litellm.proxy.management_endpoints import key_management_endpoints from litellm.proxy.management_endpoints.key_management_endpoints import ( _rotate_master_key, @@ -21114,13 +21119,20 @@ async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): "rotate_sso_identity_assertions_master_key", ): monkeypatch.setattr(key_management_endpoints, rotator, AsyncMock()) - rotate_guardrails = AsyncMock(return_value=1) - monkeypatch.setattr(GuardrailRegistry, "rotate_guardrail_params_master_key", rotate_guardrails) + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master-key") + guardrail_row = SimpleNamespace( + guardrail_id="g-1", + updated_at="t1", + litellm_params=encrypt_guardrail_litellm_params({"guardrail": "bedrock", "aws_secret_access_key": "aws-secret"}), + ) mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[guardrail_row]) + mock_prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1) await _rotate_master_key( prisma_client=mock_prisma_client, @@ -21129,4 +21141,12 @@ async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): new_master_key="sk-new-master-key", ) - rotate_guardrails.assert_awaited_once_with(prisma_client=mock_prisma_client, new_master_key="sk-new-master-key") + write = mock_prisma_client.db.litellm_guardrailstable.update_many.call_args.kwargs + stored_params = json.loads(write["data"]["litellm_params"]) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master-key") + assert write["where"] == {"guardrail_id": "g-1", "updated_at": "t1"} + assert stored_params["aws_secret_access_key"].startswith("litellm_enc::") + assert decrypt_guardrail_litellm_params(stored_params) == { + "guardrail": "bedrock", + "aws_secret_access_key": "aws-secret", + }