test(guardrails): drive the real guardrail rotator from the master key rotation test

- inject an encrypted guardrail row through the prisma client instead of replacing the GuardrailRegistry method
- assert the written params decrypt under the new master key
This commit is contained in:
Yucheng He 2026-09-28 17:29:14 -07:00
parent 27ac33a1a6
commit bd2090a85a

View file

@ -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",
}