fix(guardrails): encrypt guardrail litellm_params secrets at rest (#43627)

* fix(guardrails): encrypt guardrail litellm_params secrets at rest

* fix(guardrails): keep salt-key encryption on master key rotation and retry rows edited mid-rotation

- rotate guardrail params under LITELLM_SALT_KEY when set, matching the key reads decrypt with
- re-read and retry a row whose updated_at moved during rotation, up to GUARDRAIL_ROTATION_ATTEMPTS
- build decrypted Guardrail rows and the rotation count without mutating locals

* refactor(guardrails): retry guardrail rotation by bounded recursion instead of a rebound cursor

- each attempt re-reads the row and recurses with attempts_left - 1, so no loop variable is rebound
- cover the give-up path after GUARDRAIL_ROTATION_ATTEMPTS writes

* 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

* Annotate guardrail param encryption collections for type-discipline gate

* Type guardrail param recursion through validated JSON containers

* Type guardrail registry test helpers and drop section comment

* Reject client-supplied encrypted values in guardrail litellm_params

* Allow depth-bounded contains_encrypted_marker in the recursion detector

* Keep a loaded guardrail when its DB params do not decrypt with the current key

* Apply other DB edits while keeping loaded values that do not decrypt, including PATCH models

* Keep the loaded guardrail when an undecryptable param has no loaded value

* Drop suppressions the type discipline gate on main now reports as unused

* Assert what the reinitialized guardrail holds after an edit to an undecryptable one

* Drive the rotation sync tests through a registered guardrail instead of patching reinitialize

* Type the rotation test helpers and drop the new test docstrings

* fix(guardrails): refuse to approve a submission whose params do not decrypt

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng-berri 2026-10-02 22:49:41 -07:00 • committed by GitHub
parent 6d8434f940
commit 5d42cb7cfa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 862 additions and 21 deletions

View file

@ -83,6 +83,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

@ -141,9 +141,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
@ -157,7 +157,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

@ -21,6 +21,7 @@ from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX
from litellm.proxy.common_utils.path_utils import is_within, safe_join
from litellm.proxy.guardrails.content_filter_data import CATEGORIES_DIR, DATA_ROOTS, category_dirs, find_category_file
from litellm.proxy.guardrails.guardrail_hooks.custom_code.bounded_execution import (
@ -33,7 +34,12 @@ 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,
contains_encrypted_marker,
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
@ -81,6 +87,16 @@ def _as_str_object_mapping(mapping: Mapping[str, object]) -> Mapping[str, object
return mapping
def _reject_encrypted_litellm_params(litellm_params: object) -> None:
"""Raise 400 if a client-supplied litellm_params value carries the encrypted-value prefix."""
params: Final = litellm_params.model_dump() if isinstance(litellm_params, BaseModel) else litellm_params
if contains_encrypted_marker(params):
raise HTTPException(
status_code=400,
detail=f"litellm_params values must not start with {CALLBACK_VAR_ENCRYPTED_PREFIX!r}",
)
def _guardrails_table(prisma_client: "PrismaClient") -> "TableActions[LiteLLM_GuardrailsTable]":
return GuardrailsRepository(prisma_client).table
@ -397,6 +413,8 @@ async def create_guardrail(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
_reject_encrypted_litellm_params(request.guardrail.get("litellm_params"))
try:
result = await GUARDRAIL_REGISTRY.add_guardrail_to_db(guardrail=request.guardrail, prisma_client=prisma_client)
@ -507,6 +525,8 @@ async def update_guardrail(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
_reject_encrypted_litellm_params(request.guardrail.get("litellm_params"))
try:
# Check if guardrail exists
existing_guardrail: Final = await GUARDRAIL_REGISTRY.get_guardrail_by_id_from_db(
@ -731,6 +751,7 @@ async def register_guardrail(
)
params: Final = request.get_litellm_params_dict()
_reject_encrypted_litellm_params(params)
if params.get("guardrail") != GENERIC_GUARDRAIL_API:
raise HTTPException(
status_code=400,
@ -774,7 +795,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
@ -848,7 +869,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,
@ -1027,13 +1048,21 @@ async def approve_guardrail_submission(
detail=f"Guardrail is not pending review (status={row.status})",
)
litellm_params: Final = _parse_json_field(row.litellm_params)
decrypted_params: Final = decrypt_guardrail_litellm_params(litellm_params or {})
if contains_encrypted_marker(decrypted_params):
raise HTTPException(
status_code=409,
detail="Guardrail litellm_params do not decrypt with the current key. "
"Restart the proxy if the master key was rotated, then approve again.",
)
now: Final = datetime.now(timezone.utc)
await _guardrails_table(prisma_client).update(
where={"guardrail_id": guardrail_id},
data={"status": "active", "reviewed_at": now, "updated_at": now},
)
litellm_params: Final = _parse_json_field(row.litellm_params)
guardrail_info: Final = _parse_json_field(row.guardrail_info)
if not litellm_params:
raise HTTPException(
@ -1043,7 +1072,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": decrypted_params,
"guardrail_info": guardrail_info or {},
"team_id": row.team_id,
}
@ -1190,6 +1219,8 @@ async def patch_guardrail(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
_reject_encrypted_litellm_params(request.litellm_params)
try:
# Check if guardrail exists and get current data
existing_guardrail: Final = await GUARDRAIL_REGISTRY.get_guardrail_by_id_from_db(

View file

@ -3,23 +3,27 @@
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
from pydantic import ValidationError
from pydantic import BaseModel, TypeAdapter, ValidationError
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,129 @@ def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models
return GuardrailsRepository(prisma_client).table
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
_JSON_ARRAY: Final = TypeAdapter(list[object])
def _as_json_object(value: object) -> dict[str, object] | None:
if not isinstance(value, Mapping):
return None
try:
return _JSON_OBJECT.validate_python(value)
except ValidationError:
return None
def _as_json_array(value: object) -> list[object] | None:
return _JSON_ARRAY.validate_python(value) if isinstance(value, list) else None
def contains_encrypted_marker(value: object, depth: int = 0) -> bool:
"""True if any string in value, at any JSON depth, starts with the encrypted-value prefix."""
if depth > DEFAULT_MAX_RECURSE_DEPTH:
return False
if isinstance(value, str):
return value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX)
json_object: Final = _as_json_object(value)
if json_object is not None:
return any(contains_encrypted_marker(v, depth + 1) for v in json_object.values())
json_array: Final = _as_json_array(value)
return json_array is not None and any(contains_encrypted_marker(item, depth + 1) for item in json_array)
def _encrypted_param(key: str, value: object, new_encryption_key: str | None, depth: int = 0) -> object:
if depth > DEFAULT_MAX_RECURSE_DEPTH:
return value
json_object: Final = _as_json_object(value)
if json_object is not None:
return {k: _encrypted_param(k, v, new_encryption_key, depth + 1) for k, v in json_object.items()}
json_array: Final = _as_json_array(value)
if json_array is not None:
return [_encrypted_param(key, item, new_encryption_key, depth + 1) for item in json_array]
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
json_object: Final = _as_json_object(value)
if json_object is not None:
return {k: _decrypted_param(k, v, depth + 1) for k, v in json_object.items()}
json_array: Final = _as_json_array(value)
if json_array is not None:
return [_decrypted_param(key, item, depth + 1) for item in json_array]
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 = _as_json_object(fields.get("litellm_params"))
if stored_params is None:
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 +422,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 +468,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 +484,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 +504,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 +520,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 +536,20 @@ 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."""
salt_key: Final = os.environ.get(SALT_KEY_ENV_VAR)
encryption_key: Final = new_master_key if salt_key is None else salt_key
rows: Final = await _guardrail_table(prisma_client).find_many()
rotated = [await _rotate_guardrail_row(prisma_client, row, encryption_key) for row in rows]
return sum(rotated)
def _apply_configured_bool_overrides(instance: CustomGuardrail, litellm_params: LitellmParams) -> None:
"""Override the parallel/raw-scan flags only when ``litellm_params`` explicitly
@ -857,9 +993,40 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id)
raise ValueError(f"Guardrail initialization failed: {init_error}") from init_error
def _with_loaded_values_where_undecryptable(self, guardrail_id: str, guardrail: Guardrail) -> Guardrail:
"""Swap each DB litellm_params value that did not decrypt with the current key for the loaded guardrail's value,
or keep the loaded guardrail whole when it has no value for one of them."""
existing: Final = self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
stored_params: Final = guardrail.get("litellm_params")
db_params: Final = _as_json_object(
stored_params.model_dump() if isinstance(stored_params, BaseModel) else stored_params
)
if existing is None or db_params is None or not contains_encrypted_marker(db_params):
return guardrail
loaded_params: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params"))
verbose_proxy_logger.warning(
"Guardrail %s has litellm_params that do not decrypt with the current key; keeping the loaded values for "
"them. Restart the proxy if the master key was rotated.",
guardrail_id,
)
if loaded_params is None or any(
contains_encrypted_marker(value) and loaded_params.get(key) is None for key, value in db_params.items()
):
return existing
return Guardrail(
**{
**guardrail,
"litellm_params": {
key: loaded_params.get(key) if contains_encrypted_marker(value) else value
for key, value in db_params.items()
},
}
)
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
"""
Sync a guardrail from DB - initializes if new, re-initializes if changed.
DB values that do not decrypt with the current key keep the loaded guardrail's values.
This is the method to call during DB polling.
"""
guardrail_id: Final = guardrail.get("guardrail_id")
@ -867,13 +1034,14 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.error("Cannot sync guardrail without guardrail_id")
return None
if self._has_guardrail_params_changed(guardrail_id, guardrail):
guardrail_name: Final = guardrail.get("guardrail_name", "Unknown")
synced: Final = self._with_loaded_values_where_undecryptable(guardrail_id, guardrail)
if self._has_guardrail_params_changed(guardrail_id, synced):
guardrail_name: Final = synced.get("guardrail_name", "Unknown")
verbose_proxy_logger.info(
"Guardrail '%s' (ID: %s) params changed, re-initializing...", guardrail_name, guardrail_id
)
return self.reinitialize_guardrail(
guardrail=guardrail,
guardrail=synced,
config_file_path=config_file_path,
source="db",
)

View file

@ -5293,6 +5293,15 @@ async def _rotate_master_key(
data={"param_value": prisma.Json(encrypted_env_vars)},
)
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,10 @@ 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.
"contains_encrypted_marker", # 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

@ -572,6 +572,42 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra
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"
@pytest.mark.asyncio
@pytest.mark.parametrize("lookup", ["agent-id", "Agent name"])
async def test_agent_read_through_hydrates_identity_binding(lookup, clean_agent_registry, fresh_agent_read_through, monkeypatch):

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

@ -41,6 +41,7 @@ MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
from litellm.proxy.guardrails.guardrail_registry import (
IN_MEMORY_GUARDRAIL_HANDLER,
InMemoryGuardrailHandler,
encrypt_guardrail_litellm_params,
)
from litellm.types.guardrails import (
ApplyGuardrailRequest,
@ -2675,6 +2676,91 @@ async def test_test_custom_code_endpoint_reports_a_system_exit_as_an_execution_e
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):
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"
@pytest.mark.asyncio
async def test_approve_guardrail_submission_rejects_params_that_do_not_decrypt(mocker, monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-worker-key")
stored_params = encrypt_guardrail_litellm_params(
{"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "team-vendor-secret-1234"},
new_encryption_key="sk-rotated-key-the-worker-lacks",
)
row = mocker.Mock(
guardrail_id="reg-rotated",
guardrail_name="team-rotated",
status="pending_review",
team_id="team-1",
litellm_params=stored_params,
guardrail_info={},
)
mock_prisma = mocker.Mock()
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row)
mock_prisma.db.litellm_guardrailstable.update = AsyncMock()
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
mock_handler = mocker.Mock()
mocker.patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", mock_handler)
with pytest.raises(HTTPException) as exc_info:
await approve_guardrail_submission("reg-rotated", MOCK_ADMIN_USER)
assert exc_info.value.status_code == 409
mock_prisma.db.litellm_guardrailstable.update.assert_not_called()
mock_handler.initialize_guardrail.assert_not_called()
@pytest.mark.asyncio
async def test_get_category_yaml_returns_bundled_category_and_its_file_type():
result = await get_category_yaml("harmful_self_harm", roots=DATA_ROOTS)
@ -2728,3 +2814,103 @@ async def test_get_category_yaml_serves_a_symlink_that_stays_inside_a_category_f
result = await get_category_yaml("alias", roots=(*DATA_ROOTS, str(tmp_path / "legacy")))
assert result["file_type"] == "yaml"
assert yaml.safe_load(result["yaml_content"])["category_name"] == "real"
_ENCRYPTED_MARKER_VALUE = "litellm_enc::opaque-value"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"extra_params",
[
{"description": _ENCRYPTED_MARKER_VALUE},
{"api_key": _ENCRYPTED_MARKER_VALUE},
{"extra_headers": {"x-team": "a", "x-secret": _ENCRYPTED_MARKER_VALUE}},
{"extra_headers": ["plain", _ENCRYPTED_MARKER_VALUE]},
],
ids=["top_level_description", "top_level_api_key", "nested_object", "array_second_element"],
)
async def test_register_guardrail_rejects_encrypted_marker_values(mocker, extra_params):
mock_prisma = mocker.Mock()
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_guardrailstable.create = AsyncMock()
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
req = RegisterGuardrailRequest(
guardrail_name="marker-guard",
litellm_params={
"guardrail": "generic_guardrail_api",
"mode": "pre_call",
"api_base": "https://guardrails.example.com/validate",
**extra_params,
},
)
user = UserAPIKeyAuth(user_id="u1", user_email="a@b.com", team_id="team-1")
with pytest.raises(HTTPException) as exc_info:
await register_guardrail(req, user)
assert exc_info.value.status_code == 400
assert "litellm_enc::" in exc_info.value.detail
mock_prisma.db.litellm_guardrailstable.create.assert_not_called()
def _guardrail_with_encrypted_api_key() -> Guardrail:
return Guardrail(
guardrail_name="marker-guard",
litellm_params=LitellmParams(
guardrail="generic_guardrail_api",
mode="pre_call",
api_base="https://guardrails.example.com/validate",
api_key=_ENCRYPTED_MARKER_VALUE,
),
)
@pytest.mark.asyncio
async def test_create_guardrail_rejects_encrypted_marker_values(mocker, mock_guardrail_registry):
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) # test-quality-ok: endpoint has no DI seam
mocker.patch( # test-quality-ok: endpoint has no DI seam
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry
)
with pytest.raises(HTTPException) as exc_info:
await create_guardrail(
CreateGuardrailRequest(guardrail=_guardrail_with_encrypted_api_key()),
user_api_key_dict=MOCK_ADMIN_USER,
)
assert exc_info.value.status_code == 400
mock_guardrail_registry.add_guardrail_to_db.assert_not_called()
@pytest.mark.asyncio
async def test_update_guardrail_rejects_encrypted_marker_values(mocker, mock_guardrail_registry):
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) # test-quality-ok: endpoint has no DI seam
mocker.patch( # test-quality-ok: endpoint has no DI seam
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry
)
with pytest.raises(HTTPException) as exc_info:
await update_guardrail(
"test-guardrail-id",
UpdateGuardrailRequest(guardrail=_guardrail_with_encrypted_api_key()),
user_api_key_dict=MOCK_ADMIN_USER,
)
assert exc_info.value.status_code == 400
mock_guardrail_registry.update_guardrail_in_db.assert_not_called()
@pytest.mark.asyncio
async def test_patch_guardrail_rejects_encrypted_marker_values(mocker, mock_guardrail_registry):
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) # test-quality-ok: endpoint has no DI seam
mocker.patch( # test-quality-ok: endpoint has no DI seam
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(api_key=_ENCRYPTED_MARKER_VALUE))
with pytest.raises(HTTPException) as exc_info:
await patch_guardrail("test-guardrail-id", request, user_api_key_dict=MOCK_ADMIN_USER)
assert exc_info.value.status_code == 400
mock_guardrail_registry.update_guardrail_in_db.assert_not_called()

View file

@ -1,5 +1,5 @@
from collections.abc import Iterable
from unittest.mock import AsyncMock, MagicMock
from collections.abc import Iterable, Iterator
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -400,6 +400,99 @@ def test_sync_guardrail_from_db_marks_source_db_when_unchanged():
assert handler.get_source("collide") == "db"
@pytest.fixture
def rotation_handler() -> Iterator[InMemoryGuardrailHandler]:
registry_module = _register_mode_following_initializer("rotation_test")
lists = _all_callback_lists()
snapshots = [list(cb_list) for cb_list in lists]
try:
yield InMemoryGuardrailHandler()
finally:
registry_module.guardrail_initializer_registry.pop("rotation_test", None)
for cb_list, snapshot in zip(lists, snapshots):
cb_list[:] = snapshot
def _rotation_row(litellm_params: dict[str, object] | LitellmParams) -> Guardrail:
return Guardrail(guardrail_id="rotated", guardrail_name="mode-following", litellm_params=litellm_params)
_LOADED_PARAMS = {"guardrail": "rotation_test", "mode": "pre_call", "default_on": True, "api_key": "gk-loaded"}
def test_sync_guardrail_from_db_keeps_the_loaded_guardrail_when_db_params_do_not_decrypt(rotation_handler):
rotation_handler.initialize_guardrail(guardrail=_rotation_row(dict(_LOADED_PARAMS)), source="db")
live_instance = rotation_handler.guardrail_id_to_custom_guardrail["rotated"]
rotation_handler.sync_guardrail_from_db(
_rotation_row({**_LOADED_PARAMS, "api_key": "litellm_enc::sealed-under-the-new-key"})
)
assert rotation_handler.guardrail_id_to_custom_guardrail["rotated"] is live_instance
assert rotation_handler.IN_MEMORY_GUARDRAILS["rotated"]["litellm_params"].api_key == "gk-loaded"
def test_sync_guardrail_from_db_applies_other_edits_and_keeps_the_loaded_value_that_does_not_decrypt(
rotation_handler,
):
rotation_handler.initialize_guardrail(guardrail=_rotation_row(dict(_LOADED_PARAMS)), source="db")
rotation_handler.sync_guardrail_from_db(
_rotation_row({**_LOADED_PARAMS, "mode": "post_call", "api_key": "litellm_enc::sealed-under-the-new-key"})
)
synced_params = rotation_handler.IN_MEMORY_GUARDRAILS["rotated"]["litellm_params"]
assert synced_params.mode == "post_call"
assert synced_params.api_key == "gk-loaded"
live_instance = rotation_handler.guardrail_id_to_custom_guardrail["rotated"]
assert live_instance.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_call) is True
def test_sync_guardrail_from_db_keeps_the_loaded_guardrail_when_an_undecryptable_param_has_no_loaded_value(
rotation_handler,
):
loaded_params = {key: value for key, value in _LOADED_PARAMS.items() if key != "api_key"}
rotation_handler.initialize_guardrail(guardrail=_rotation_row(dict(loaded_params)), source="db")
live_instance = rotation_handler.guardrail_id_to_custom_guardrail["rotated"]
rotation_handler.sync_guardrail_from_db(
_rotation_row({**loaded_params, "mode": "post_call", "api_key": "litellm_enc::sealed-under-the-new-key"})
)
assert rotation_handler.guardrail_id_to_custom_guardrail["rotated"] is live_instance
synced_params = rotation_handler.IN_MEMORY_GUARDRAILS["rotated"]["litellm_params"]
assert synced_params.mode == "pre_call"
assert synced_params.api_key is None
def test_sync_guardrail_from_db_keeps_the_loaded_value_when_a_patch_passes_litellm_params_as_a_model(
rotation_handler,
):
rotation_handler.initialize_guardrail(guardrail=_rotation_row(dict(_LOADED_PARAMS)), source="db")
rotation_handler.sync_guardrail_from_db(
_rotation_row(LitellmParams(**{**_LOADED_PARAMS, "default_on": False, "api_key": "litellm_enc::sealed"}))
)
synced_params = rotation_handler.IN_MEMORY_GUARDRAILS["rotated"]["litellm_params"]
assert synced_params.default_on is False
assert synced_params.api_key == "gk-loaded"
def test_sync_guardrail_from_db_applies_an_edit_to_a_guardrail_loaded_with_an_undecryptable_value(
rotation_handler,
):
stale_params = {**_LOADED_PARAMS, "api_key": "litellm_enc::stale"}
rotation_handler.initialize_guardrail(guardrail=_rotation_row(dict(stale_params)), source="db")
rotation_handler.sync_guardrail_from_db(_rotation_row({**stale_params, "mode": "post_call", "default_on": False}))
synced_params = rotation_handler.IN_MEMORY_GUARDRAILS["rotated"]["litellm_params"]
assert synced_params.mode == "post_call"
assert synced_params.default_on is False
assert synced_params.api_key == "litellm_enc::stale"
def _db_litellm_params() -> dict:
"""
Shape produced by GuardrailRegistry.get_all_guardrails_from_db: litellm_params
@ -1086,3 +1179,233 @@ 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[str, object]):
def __getattr__(self, name: str) -> object:
return self[name]
def _stored_params(create_or_update_mock: AsyncMock) -> dict[str, object]:
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

@ -21097,3 +21097,59 @@ 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):
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",
}