From 5d42cb7cfa8098e84663420b3909d0ce71def4fe Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Fri, 2 Oct 2026 22:49:41 -0700 Subject: [PATCH] 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> --- litellm/constants.py | 1 + .../common_utils/registry_read_through.py | 4 +- litellm/proxy/db/master_key_migration.py | 1 + .../proxy/guardrails/guardrail_endpoints.py | 41 ++- .../proxy/guardrails/guardrail_registry.py | 192 +++++++++- .../key_management_endpoints.py | 9 + .../code_coverage_tests/recursive_detector.py | 4 + .../test_registry_read_through.py | 36 ++ .../proxy/db/test_master_key_migration.py | 26 ++ .../guardrails/test_guardrail_endpoints.py | 186 ++++++++++ .../guardrails/test_guardrail_registry.py | 327 +++++++++++++++++- .../test_key_management_endpoints.py | 56 +++ 12 files changed, 862 insertions(+), 21 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 18fc6aa7e74..49514fc4d0e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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) diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py index 32beb0b7447..0a47de96645 100644 --- a/litellm/proxy/common_utils/registry_read_through.py +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -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 diff --git a/litellm/proxy/db/master_key_migration.py b/litellm/proxy/db/master_key_migration.py index d100554201a..803f67e142d 100644 --- a/litellm/proxy/db/master_key_migration.py +++ b/litellm/proxy/db/master_key_migration.py @@ -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"), diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index aee6260b5e2..4195f319f16 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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( diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 0dc50cd6196..1374a88cbfe 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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", ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d417ec1479f..459868ed78d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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( diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 863934d76f7..863e8befcf9 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -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. diff --git a/tests/unit/proxy/common_utils/test_registry_read_through.py b/tests/unit/proxy/common_utils/test_registry_read_through.py index 5713f7dfaa7..35f448c4fcf 100644 --- a/tests/unit/proxy/common_utils/test_registry_read_through.py +++ b/tests/unit/proxy/common_utils/test_registry_read_through.py @@ -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): diff --git a/tests/unit/proxy/db/test_master_key_migration.py b/tests/unit/proxy/db/test_master_key_migration.py index 9c0fc163b9f..948e8ac8518 100644 --- a/tests/unit/proxy/db/test_master_key_migration.py +++ b/tests/unit/proxy/db/test_master_key_migration.py @@ -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")] diff --git a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py index 4339febb0e3..a3d4786f7d1 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py @@ -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() diff --git a/tests/unit/proxy/guardrails/test_guardrail_registry.py b/tests/unit/proxy/guardrails/test_guardrail_registry.py index 022fe85c779..2f29f964e63 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_registry.py +++ b/tests/unit/proxy/guardrails/test_guardrail_registry.py @@ -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 diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 5ea38ce23d5..55fbe04d1bc 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -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", + }