fix(guardrails): encrypt guardrail litellm_params secrets at rest

This commit is contained in:
Yucheng He 2026-09-28 14:17:50 -07:00
parent 6d4ccf7e97
commit 032dc13d26
11 changed files with 422 additions and 14 deletions

View file

@ -136,9 +136,9 @@ async def _resync_guardrails(guardrail_name: str) -> bool:
from litellm.proxy.guardrails.guardrail_registry import (
GUARDRAIL_RECONCILE_LOCK,
IN_MEMORY_GUARDRAIL_HANDLER,
guardrail_from_db_row,
)
from litellm.repositories.table_repositories import GuardrailsRepository
from litellm.types.guardrails import Guardrail
if not _db_backed_registries_enabled("guardrails"):
return False
@ -152,7 +152,7 @@ async def _resync_guardrails(guardrail_name: str) -> bool:
if row is None:
return False
async with GUARDRAIL_RECONCILE_LOCK:
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=Guardrail(**dict(row)))
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=guardrail_from_db_row(row))
return _initialized_guardrail(guardrail_name) is not None

View file

@ -27,6 +27,7 @@ _SECRET_COLUMNS: Final = (
_SecretColumn("LiteLLM_ProxyModelTable", "model_id", "litellm_params"),
_SecretColumn("LiteLLM_CredentialsTable", "credential_id", "credential_values"),
_SecretColumn("LiteLLM_Config", "param_name", "param_value"),
_SecretColumn("LiteLLM_GuardrailsTable", "guardrail_id", "litellm_params", only_rows_with_marked_ciphertexts=True),
_SecretColumn("LiteLLM_SSOConfig", "id", "sso_settings"),
_SecretColumn("LiteLLM_CacheConfig", "id", "cache_settings"),
_SecretColumn("LiteLLM_ConfigOverrides", "config_type", "config_value"),

View file

@ -32,7 +32,11 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
build_sandbox_globals,
compile_sandboxed,
)
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
from litellm.proxy.guardrails.guardrail_registry import (
GuardrailRegistry,
decrypt_guardrail_litellm_params,
encrypt_guardrail_litellm_params,
)
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.repositories.prisma_protocols import TableActions
@ -773,7 +777,7 @@ async def register_guardrail(
raise HTTPException(status_code=500, detail=str(e))
now: Final = datetime.now(timezone.utc)
litellm_params_str: Final = safe_dumps(params)
litellm_params_str: Final = safe_dumps(encrypt_guardrail_litellm_params(params))
guardrail_info: Final = dict(request.guardrail_info or {})
guardrail_info["submitted_by_user_id"] = user_api_key_dict.user_id
guardrail_info["submitted_by_email"] = user_api_key_dict.user_email
@ -847,7 +851,7 @@ def _row_to_submission_item(row: "LiteLLM_GuardrailsTable") -> GuardrailSubmissi
guardrail_info: Final = _parse_json_field(row.guardrail_info) or {}
team_guardrail: Final = row.team_id is not None
raw_params: Final = _parse_json_field(row.litellm_params) or {}
raw_params: Final = decrypt_guardrail_litellm_params(_parse_json_field(row.litellm_params) or {})
masked_params: Final = _get_masked_values(raw_params, unmasked_length=4, number_of_asterisks=4)
return GuardrailSubmissionItem(
guardrail_id=row.guardrail_id,
@ -1042,7 +1046,7 @@ async def approve_guardrail_submission(
guardrail_dict: Final = {
"guardrail_id": row.guardrail_id,
"guardrail_name": row.guardrail_name,
"litellm_params": litellm_params,
"litellm_params": decrypt_guardrail_litellm_params(litellm_params),
"guardrail_info": guardrail_info or {},
"team_id": row.team_id,
}

View file

@ -3,7 +3,7 @@
import asyncio
import importlib
import os
from collections.abc import Callable, Iterator, Mapping, Sequence
from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence
from datetime import datetime, timezone
from itertools import chain, count
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, cast
@ -14,12 +14,15 @@ import litellm
from litellm import Router
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_scan_only_tool_results_for_guardrail,
effective_skip_tool_message_for_guardrail,
)
from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX, is_sensitive_callback_key
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockGuardrail,
)
@ -77,6 +80,66 @@ def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models
return GuardrailsRepository(prisma_client).table
def _encrypted_param(key: str, value: object, new_encryption_key: str | None, depth: int = 0) -> object:
if depth > DEFAULT_MAX_RECURSE_DEPTH:
return value
match value:
case dict():
return {k: _encrypted_param(k, v, new_encryption_key, depth + 1) for k, v in value.items()}
case list():
return [_encrypted_param(key, item, new_encryption_key, depth + 1) for item in value]
case str() if value and is_sensitive_callback_key(key) and not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX):
try:
return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(
value, new_encryption_key=new_encryption_key
)
except Exception: # noqa: BLE001 # no salt key or master key configured: store the value as written
return value
case _:
return value
def _decrypted_param(key: str, value: object, depth: int = 0) -> object:
if depth > DEFAULT_MAX_RECURSE_DEPTH:
return value
match value:
case dict():
return {k: _decrypted_param(k, v, depth + 1) for k, v in value.items()}
case list():
return [_decrypted_param(key, item, depth + 1) for item in value]
case str() if value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX):
decrypted: Final = decrypt_value_helper(
value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX),
key=key,
exception_type="debug",
return_original_value=False,
)
return value if decrypted is None else decrypted
case _:
return value
def encrypt_guardrail_litellm_params(
litellm_params: Mapping[str, object], new_encryption_key: str | None = None
) -> dict[str, object]:
"""Encrypt every string stored under a sensitive key (at any dict depth) for the guardrails table."""
return {key: _encrypted_param(key, value, new_encryption_key) for key, value in litellm_params.items()}
def decrypt_guardrail_litellm_params(litellm_params: Mapping[str, object]) -> dict[str, object]:
"""Decrypt values written by encrypt_guardrail_litellm_params; plaintext values pass through unchanged."""
return {key: _decrypted_param(key, value) for key, value in litellm_params.items()}
def guardrail_from_db_row(row: Iterable[tuple[str, object]]) -> Guardrail:
"""Build a Guardrail from a guardrails table row with its litellm_params decrypted."""
fields: Final = dict(row)
stored_params: Final = fields.get("litellm_params")
if isinstance(stored_params, Mapping):
fields["litellm_params"] = decrypt_guardrail_litellm_params(stored_params)
return Guardrail(**fields)
guardrail_initializer_registry: Final = {
SupportedGuardrailIntegrations.BEDROCK.value: initialize_bedrock,
SupportedGuardrailIntegrations.LAKERA.value: initialize_lakera,
@ -295,7 +358,7 @@ class GuardrailRegistry:
litellm_params_dict = litellm_params_obj.model_dump()
else:
litellm_params_dict = dict(litellm_params_obj) if litellm_params_obj else {}
litellm_params: Final[str] = safe_dumps(litellm_params_dict)
litellm_params: Final[str] = safe_dumps(encrypt_guardrail_litellm_params(litellm_params_dict))
guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {}))
# Create guardrail in DB
@ -341,7 +404,7 @@ class GuardrailRegistry:
litellm_params_dict = litellm_params_obj.model_dump()
else:
litellm_params_dict = dict(litellm_params_obj) if litellm_params_obj else {}
litellm_params: Final[str] = safe_dumps(litellm_params_dict)
litellm_params: Final[str] = safe_dumps(encrypt_guardrail_litellm_params(litellm_params_dict))
guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {}))
# Update in DB
@ -357,8 +420,7 @@ class GuardrailRegistry:
if updated_guardrail is None:
raise ValueError(f"Guardrail not found, passed guardrail_id={guardrail_id}")
# Convert to dict and return
return dict(updated_guardrail)
return dict(guardrail_from_db_row(updated_guardrail))
except Exception as e:
raise Exception(f"Error updating guardrail in DB: {e}")
@ -378,7 +440,7 @@ class GuardrailRegistry:
guardrails: Final[list[Guardrail]] = []
for guardrail in guardrails_from_db:
guardrails.append(Guardrail(**(dict(guardrail))))
guardrails.append(guardrail_from_db_row(guardrail))
return guardrails
except Exception as e:
@ -394,7 +456,7 @@ class GuardrailRegistry:
if not guardrail:
return None
return Guardrail(**(dict(guardrail)))
return guardrail_from_db_row(guardrail)
except Exception as e:
raise Exception(f"Error getting guardrail from DB: {e}")
@ -410,10 +472,29 @@ class GuardrailRegistry:
if not guardrail:
return None
return Guardrail(**(dict(guardrail)))
return guardrail_from_db_row(guardrail)
except Exception as e:
raise Exception(f"Error getting guardrail from DB: {e}")
@staticmethod
async def rotate_guardrail_params_master_key(prisma_client: PrismaClient, new_master_key: str) -> int:
"""Re-encrypt the sensitive litellm_params of every guardrail row under new_master_key. Returns rows updated."""
rows_updated = 0
for row in await _guardrail_table(prisma_client).find_many():
stored_params = row.litellm_params
if not isinstance(stored_params, Mapping):
continue
rotated_params = encrypt_guardrail_litellm_params(
decrypt_guardrail_litellm_params(stored_params), new_encryption_key=new_master_key
)
if rotated_params == stored_params:
continue
rows_updated += await _guardrail_table(prisma_client).update_many(
where={"guardrail_id": row.guardrail_id, "updated_at": row.updated_at},
data={"litellm_params": safe_dumps(rotated_params)},
)
return rows_updated
def _apply_configured_bool_overrides(instance: CustomGuardrail, litellm_params: LitellmParams) -> None:
"""Override the parallel/raw-scan flags only when ``litellm_params`` explicitly

View file

@ -5301,6 +5301,16 @@ async def _rotate_master_key(
data={"param_value": prisma.Json(encrypted_env_vars)},
)
# 3b. process guardrails table
try:
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
await GuardrailRegistry.rotate_guardrail_params_master_key(
prisma_client=prisma_client, new_master_key=new_master_key
)
except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation
verbose_proxy_logger.warning("Failed to rotate guardrail params: %s", str(e))
# 4. process MCP server table
try:
await rotate_mcp_server_credentials_master_key(

View file

@ -36,6 +36,8 @@ IGNORE_FUNCTIONS = [
"_collect_argument_paths", # max depth set.
"_split_text", # max depth set.
"_mask_sequence", # max depth set.
"_encrypted_param", # max depth set.
"_decrypted_param", # max depth set.
"_delete_nested_value_custom", # max depth set (bounded by number of path segments).
"filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion.
"__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion.

View file

@ -521,3 +521,39 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra
assert await resync_task is True
assert len(clean_agent_registry.agent_list) == 1
@pytest.mark.asyncio
async def test_resync_guardrails_syncs_decrypted_litellm_params(monkeypatch):
from unittest.mock import AsyncMock, MagicMock
import litellm.proxy.common_utils.registry_read_through as read_through_module
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.common_utils.registry_read_through import _resync_guardrails
from litellm.proxy.guardrails.guardrail_registry import (
IN_MEMORY_GUARDRAIL_HANDLER,
encrypt_guardrail_litellm_params,
)
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test")
encrypted_params: Final = encrypt_guardrail_litellm_params(
{"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "vendor-key"}
)
prisma_client: Final = MagicMock()
prisma_client.db.litellm_guardrailstable.find_first = AsyncMock(
return_value={
"guardrail_id": "enc-id",
"guardrail_name": "enc-guardrail",
"litellm_params": encrypted_params,
"guardrail_info": {},
"status": "active",
}
)
synced: list[dict] = []
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
monkeypatch.setattr(IN_MEMORY_GUARDRAIL_HANDLER, "sync_guardrail_from_db", lambda guardrail: synced.append(guardrail))
monkeypatch.setattr(read_through_module, "_initialized_guardrail", lambda guardrail_name: MagicMock())
assert await _resync_guardrails("enc-guardrail") is True
assert synced[0]["litellm_params"]["api_key"] == "vendor-key"

View file

@ -561,3 +561,29 @@ async def test_boot_leaves_the_database_alone_unless_a_migration_was_requested_a
assert result is outcome
assert len(database_handles_taken) == (0 if outcome is None else 1)
assert len(logged) == (0 if outcome is None else 1)
@pytest.mark.asyncio
async def test_guardrail_params_move_to_the_new_key_and_legacy_plaintext_rows_are_left_alone():
legacy_params = {"guardrail": "generic_guardrail_api", "api_key": "legacy-plaintext-key"}
tables: Tables = {
"LiteLLM_GuardrailsTable": [
{
"guardrail_id": "guardrail-1",
"litellm_params": {
"guardrail": "generic_guardrail_api",
"api_key": "litellm_enc::" + _encrypted("guardrail-vendor-key"),
},
},
{"guardrail_id": "guardrail-legacy", "litellm_params": dict(legacy_params)},
]
}
database = _FakeDatabase(tables)
assert await reencrypt_stored_values(database, from_key=PREVIOUS_KEY, to_key=NEW_KEY) == 1
migrated_key = tables["LiteLLM_GuardrailsTable"][0]["litellm_params"]["api_key"]
assert migrated_key.startswith("litellm_enc::")
assert decrypt_if_encrypted_with(migrated_key.removeprefix("litellm_enc::"), NEW_KEY) == "guardrail-vendor-key"
assert tables["LiteLLM_GuardrailsTable"][1]["litellm_params"] == legacy_params
assert database.writes == [("LiteLLM_GuardrailsTable", "litellm_params", "guardrail-1")]

View file

@ -2670,3 +2670,59 @@ async def test_test_custom_code_endpoint_reports_a_system_exit_as_an_execution_e
assert response.error == "Execution error: SystemExit: bye"
assert response.error_type == "execution"
assert time.monotonic() - started < 2.0
@pytest.mark.asyncio
async def test_team_guardrail_api_key_is_encrypted_at_rest_and_decrypted_on_review(mocker, monkeypatch):
"""register stores the vendor api_key encrypted; submission detail masks the plaintext and approve loads it."""
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test")
mock_prisma = mocker.Mock()
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_guardrailstable.create = AsyncMock(
return_value=mocker.Mock(
guardrail_id="reg-enc",
guardrail_name="team-enc",
status="pending_review",
submitted_at=datetime.now(),
)
)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
request = RegisterGuardrailRequest(
guardrail_name="team-enc",
litellm_params={
"guardrail": "generic_guardrail_api",
"mode": "pre_call",
"api_base": "https://guardrails.example.com/validate",
"api_key": "team-vendor-secret-1234",
},
)
await register_guardrail(request, UserAPIKeyAuth(user_id="u1", team_id="team-1"))
stored_params = json.loads(mock_prisma.db.litellm_guardrailstable.create.call_args[1]["data"]["litellm_params"])
assert stored_params["api_key"].startswith("litellm_enc::")
assert "team-vendor-secret-1234" not in json.dumps(stored_params)
row = mocker.Mock(
guardrail_id="reg-enc",
guardrail_name="team-enc",
status="pending_review",
team_id="team-1",
litellm_params=stored_params,
guardrail_info={},
submitted_at=None,
reviewed_at=None,
created_at=datetime.now(),
updated_at=datetime.now(),
)
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row)
mock_prisma.db.litellm_guardrailstable.update = AsyncMock()
admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
submission = await get_guardrail_submission("reg-enc", admin)
assert submission.litellm_params["api_key"] == "te****34"
mock_handler = mocker.Mock()
mocker.patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", mock_handler)
await approve_guardrail_submission("reg-enc", admin)
loaded = mock_handler.initialize_guardrail.call_args.kwargs["guardrail"]
assert loaded["litellm_params"]["api_key"] == "team-vendor-secret-1234"

View file

@ -1084,3 +1084,158 @@ def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance():
finally:
for cb_list, snapshot in zip(lists, snapshots):
cb_list[:] = snapshot
_ENCRYPTED_PREFIX = "litellm_enc::"
class _Row(dict):
"""Prisma-row stand-in: dict(row) yields the columns and attributes read like a model."""
def __getattr__(self, name):
return self[name]
def _stored_params(create_or_update_mock: AsyncMock) -> dict:
import json
return json.loads(create_or_update_mock.call_args.kwargs["data"]["litellm_params"])
@pytest.mark.asyncio
async def test_add_guardrail_to_db_encrypts_sensitive_params_at_rest(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test")
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.create = AsyncMock(return_value=_Row(guardrail_id="g-1"))
await GuardrailRegistry().add_guardrail_to_db(
guardrail=Guardrail(
guardrail_name="vendor",
litellm_params=LitellmParams(
guardrail="generic_guardrail_api",
mode="pre_call",
api_key="vendor-secret-key",
api_base="http://vendor.example",
aws_secret_access_key="aws-secret",
custom_headers={"Authorization": "Bearer header-secret", "x-tenant": "t1"},
),
),
prisma_client=prisma_client,
)
stored = _stored_params(prisma_client.db.litellm_guardrailstable.create)
for leaked in ("vendor-secret-key", "aws-secret", "header-secret"):
assert leaked not in str(stored)
assert stored["api_key"].startswith(_ENCRYPTED_PREFIX)
assert stored["aws_secret_access_key"].startswith(_ENCRYPTED_PREFIX)
assert stored["custom_headers"]["Authorization"].startswith(_ENCRYPTED_PREFIX)
assert stored["custom_headers"]["x-tenant"] == "t1"
assert stored["guardrail"] == "generic_guardrail_api"
assert stored["mode"] == "pre_call"
assert stored["api_base"] == "http://vendor.example"
@pytest.mark.asyncio
async def test_get_all_guardrails_from_db_decrypts_new_rows_and_reads_legacy_plaintext(monkeypatch):
from litellm.proxy.guardrails.guardrail_registry import encrypt_guardrail_litellm_params
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test")
encrypted_row = _Row(
guardrail_id="g-new",
guardrail_name="new",
litellm_params=encrypt_guardrail_litellm_params(
{"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "new-key"}
),
)
legacy_row = _Row(
guardrail_id="g-legacy",
guardrail_name="legacy",
litellm_params={"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "legacy-key"},
)
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[encrypted_row, legacy_row])
guardrails = await GuardrailRegistry.get_all_guardrails_from_db(prisma_client=prisma_client)
assert [g["litellm_params"]["api_key"] for g in guardrails] == ["new-key", "legacy-key"]
@pytest.mark.asyncio
async def test_update_guardrail_in_db_encrypts_and_returns_decrypted_row(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test")
prisma_client = MagicMock()
async def _update(where, data):
import json
return _Row(
guardrail_id=where["guardrail_id"],
guardrail_name="vendor",
litellm_params=json.loads(data["litellm_params"]),
)
prisma_client.db.litellm_guardrailstable.update = AsyncMock(side_effect=_update)
result = await GuardrailRegistry().update_guardrail_in_db(
guardrail_id="g-1",
guardrail=Guardrail(
guardrail_name="vendor",
litellm_params={"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "rotated-key"},
),
prisma_client=prisma_client,
)
assert _stored_params(prisma_client.db.litellm_guardrailstable.update)["api_key"].startswith(_ENCRYPTED_PREFIX)
assert result["litellm_params"]["api_key"] == "rotated-key"
def test_encrypt_guardrail_litellm_params_does_not_double_encrypt(monkeypatch):
from litellm.proxy.guardrails.guardrail_registry import (
decrypt_guardrail_litellm_params,
encrypt_guardrail_litellm_params,
)
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test")
params = {
"api_key": "k",
"default_on": True,
"auth_token": None,
"extra_headers": [{"x-api-key": "list-secret", "x-tenant": "t1"}],
}
encrypted = encrypt_guardrail_litellm_params(params)
assert encrypted["extra_headers"][0]["x-api-key"].startswith(_ENCRYPTED_PREFIX)
assert encrypted["extra_headers"][0]["x-tenant"] == "t1"
assert encrypt_guardrail_litellm_params(encrypted) == encrypted
assert decrypt_guardrail_litellm_params(encrypted) == params
@pytest.mark.asyncio
async def test_rotate_guardrail_params_master_key_reencrypts_under_the_new_key(monkeypatch):
from litellm.proxy.guardrails.guardrail_registry import (
decrypt_guardrail_litellm_params,
encrypt_guardrail_litellm_params,
)
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master")
stored = encrypt_guardrail_litellm_params({"guardrail": "bedrock", "mode": "pre_call", "api_key": "vendor-key"})
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(
return_value=[_Row(guardrail_id="g-1", updated_at="2026-09-28T00:00:00Z", litellm_params=stored)]
)
prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1)
rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key(
prisma_client=prisma_client, new_master_key="sk-new-master"
)
rotated = _stored_params(prisma_client.db.litellm_guardrailstable.update_many)
assert rows_updated == 1
assert prisma_client.db.litellm_guardrailstable.update_many.call_args.kwargs["where"] == {
"guardrail_id": "g-1",
"updated_at": "2026-09-28T00:00:00Z",
}
assert rotated["api_key"] != stored["api_key"]
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master")
assert decrypt_guardrail_litellm_params(rotated)["api_key"] == "vendor-key"

View file

@ -21093,3 +21093,40 @@ class TestTeamAdminMemberKeyBudgetUpdate:
)
assert exc.value.status_code == 403
assert "member_key_budgets" not in str(exc.value.detail)
@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."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
from litellm.proxy.management_endpoints import key_management_endpoints
from litellm.proxy.management_endpoints.key_management_endpoints import (
_rotate_master_key,
)
for rotator in (
"rotate_mcp_server_credentials_master_key",
"rotate_mcp_user_credentials_master_key",
"rotate_mcp_user_env_vars_master_key",
"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)
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=[])
await _rotate_master_key(
prisma_client=mock_prisma_client,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"),
current_master_key="sk-old-master-key",
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")