This commit is contained in:
yucheng-berri 2026-09-30 15:10:49 +08:00 • committed by GitHub
commit 8f4dcd0648
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 539 additions and 14 deletions

View file

@ -70,6 +70,7 @@ DEFAULT_MAX_RETRIES: Final = int(os.getenv("DEFAULT_MAX_RETRIES", 2))
# radius: each record fans out to spend logs + every callback integration.
MAX_CALLBACK_LOG_RECORDS: Final = 1000
DEFAULT_MAX_RECURSE_DEPTH: Final = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100))
GUARDRAIL_ROTATION_ATTEMPTS: Final = 3
DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER", 10))
DEFAULT_FAILURE_THRESHOLD_PERCENT: Final = float(
os.getenv("DEFAULT_FAILURE_THRESHOLD_PERCENT", 0.5)

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,16 @@ 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, GUARDRAIL_ROTATION_ATTEMPTS
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.auth.master_key_boot_check import SALT_KEY_ENV_VAR
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 +81,95 @@ 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
if isinstance(value, dict):
return {k: _encrypted_param(k, v, new_encryption_key, depth + 1) for k, v in value.items()}
if isinstance(value, list):
return [_encrypted_param(key, item, new_encryption_key, depth + 1) for item in value]
if not (
isinstance(value, str)
and value
and is_sensitive_callback_key(key)
and not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX)
):
return value
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
def _decrypted_param(key: str, value: object, depth: int = 0) -> object:
if depth > DEFAULT_MAX_RECURSE_DEPTH:
return value
if isinstance(value, dict):
return {k: _decrypted_param(k, v, depth + 1) for k, v in value.items()}
if isinstance(value, list):
return [_decrypted_param(key, item, depth + 1) for item in value]
if not (isinstance(value, str) and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX)):
return value
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
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 not isinstance(stored_params, Mapping):
return Guardrail(**fields)
return Guardrail(**{**fields, "litellm_params": decrypt_guardrail_litellm_params(stored_params)})
async def _rotate_guardrail_row(
prisma_client: PrismaClient,
row: "prisma_models.LiteLLM_GuardrailsTable | None",
encryption_key: str,
attempts_left: int = GUARDRAIL_ROTATION_ATTEMPTS,
) -> int:
"""Re-encrypt one row's params under encryption_key with a compare-and-set on updated_at.
A row edited since it was read is re-read and retried, up to attempts_left writes. Returns 1 when rewritten."""
if row is None or not isinstance(row.litellm_params, Mapping):
return 0
rotated_params: Final = encrypt_guardrail_litellm_params(
decrypt_guardrail_litellm_params(row.litellm_params), new_encryption_key=encryption_key
)
if rotated_params == row.litellm_params:
return 0
if 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 1
if attempts_left <= 1:
verbose_proxy_logger.warning(
"Guardrail %s kept changing during master key rotation; its secrets were not re-encrypted",
row.guardrail_id,
)
return 0
latest_row: Final = await _guardrail_table(prisma_client).find_unique(where={"guardrail_id": row.guardrail_id})
return await _rotate_guardrail_row(prisma_client, latest_row, encryption_key, attempts_left - 1)
guardrail_initializer_registry: Final = {
SupportedGuardrailIntegrations.BEDROCK.value: initialize_bedrock,
SupportedGuardrailIntegrations.LAKERA.value: initialize_lakera,
@ -295,7 +388,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 +434,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 +450,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 +470,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 +486,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 +502,18 @@ 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 every guardrail row's sensitive litellm_params under the key the proxy decrypts with after the
rotation (LITELLM_SALT_KEY when set, otherwise new_master_key). Returns the number of rows rewritten."""
encryption_key: Final = os.environ.get(SALT_KEY_ENV_VAR) or new_master_key
rows: Final = await _guardrail_table(prisma_client).find_many()
return sum([await _rotate_guardrail_row(prisma_client, row, encryption_key) for row in rows])
def _apply_configured_bool_overrides(instance: CustomGuardrail, litellm_params: LitellmParams) -> None:
"""Override the parallel/raw-scan flags only when ``litellm_params`` explicitly

View file

@ -5303,6 +5303,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,9 @@ 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.
"_rotate_guardrail_row", # bounded by attempts_left.
"_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

@ -1086,3 +1086,234 @@ 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"
@pytest.mark.asyncio
async def test_rotate_guardrail_params_keeps_salt_key_encryption_when_salt_key_is_set(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")
stored = encrypt_guardrail_litellm_params({"guardrail": "bedrock", "api_key": "vendor-key"})
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(
return_value=[_Row(guardrail_id="g-1", updated_at="t1", litellm_params=stored)]
)
prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1)
await GuardrailRegistry.rotate_guardrail_params_master_key(prisma_client=prisma_client, new_master_key="sk-new")
rotated = _stored_params(prisma_client.db.litellm_guardrailstable.update_many)
assert decrypt_guardrail_litellm_params(rotated)["api_key"] == "vendor-key"
@pytest.mark.asyncio
async def test_rotate_guardrail_params_retries_a_row_edited_during_rotation(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")
snapshot = _Row(
guardrail_id="g-1", updated_at="t1", litellm_params=encrypt_guardrail_litellm_params({"api_key": "old-key"})
)
edited = _Row(
guardrail_id="g-1", updated_at="t2", litellm_params=encrypt_guardrail_litellm_params({"api_key": "edited-key"})
)
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[snapshot])
prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=edited)
prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(side_effect=[0, 1])
rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key(
prisma_client=prisma_client, new_master_key="sk-new-master"
)
last_call = prisma_client.db.litellm_guardrailstable.update_many.call_args
assert rows_updated == 1
assert last_call.kwargs["where"] == {"guardrail_id": "g-1", "updated_at": "t2"}
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master")
assert decrypt_guardrail_litellm_params(_stored_params(prisma_client.db.litellm_guardrailstable.update_many)) == {
"api_key": "edited-key"
}
@pytest.mark.asyncio
async def test_rotate_guardrail_params_gives_up_on_a_row_that_keeps_changing(monkeypatch):
from litellm.constants import GUARDRAIL_ROTATION_ATTEMPTS
from litellm.proxy.guardrails.guardrail_registry import encrypt_guardrail_litellm_params
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master")
row = _Row(guardrail_id="g-1", updated_at="t1", litellm_params=encrypt_guardrail_litellm_params({"api_key": "k"}))
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[row])
prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row)
prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=0)
rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key(
prisma_client=prisma_client, new_master_key="sk-new-master"
)
assert rows_updated == 0
assert prisma_client.db.litellm_guardrailstable.update_many.await_count == GUARDRAIL_ROTATION_ATTEMPTS
assert prisma_client.db.litellm_guardrailstable.find_unique.await_count == GUARDRAIL_ROTATION_ATTEMPTS - 1

View file

@ -21102,3 +21102,60 @@ 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."""
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 (
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,
)
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())
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,
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",
)
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",
}