diff --git a/litellm/constants.py b/litellm/constants.py index 530d678457d..f3658d79773 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -70,6 +70,7 @@ DEFAULT_MAX_RETRIES: Final = int(os.getenv("DEFAULT_MAX_RETRIES", 2)) # radius: each record fans out to spend logs + every callback integration. MAX_CALLBACK_LOG_RECORDS: Final = 1000 DEFAULT_MAX_RECURSE_DEPTH: Final = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100)) +GUARDRAIL_ROTATION_ATTEMPTS: Final = 3 DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER", 10)) DEFAULT_FAILURE_THRESHOLD_PERCENT: Final = float( os.getenv("DEFAULT_FAILURE_THRESHOLD_PERCENT", 0.5) diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py index 63da5d15207..5833458492d 100644 --- a/litellm/proxy/common_utils/registry_read_through.py +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -136,9 +136,9 @@ async def _resync_guardrails(guardrail_name: str) -> bool: from litellm.proxy.guardrails.guardrail_registry import ( GUARDRAIL_RECONCILE_LOCK, IN_MEMORY_GUARDRAIL_HANDLER, + guardrail_from_db_row, ) from litellm.repositories.table_repositories import GuardrailsRepository - from litellm.types.guardrails import Guardrail if not _db_backed_registries_enabled("guardrails"): return False @@ -152,7 +152,7 @@ async def _resync_guardrails(guardrail_name: str) -> bool: if row is None: return False async with GUARDRAIL_RECONCILE_LOCK: - IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=Guardrail(**dict(row))) + IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=guardrail_from_db_row(row)) return _initialized_guardrail(guardrail_name) is not None 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 6053ab26726..377aef72de6 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -32,7 +32,11 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( build_sandbox_globals, compile_sandboxed, ) -from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry +from litellm.proxy.guardrails.guardrail_registry import ( + GuardrailRegistry, + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, +) from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.repositories.prisma_protocols import TableActions @@ -773,7 +777,7 @@ async def register_guardrail( raise HTTPException(status_code=500, detail=str(e)) now: Final = datetime.now(timezone.utc) - litellm_params_str: Final = safe_dumps(params) + litellm_params_str: Final = safe_dumps(encrypt_guardrail_litellm_params(params)) guardrail_info: Final = dict(request.guardrail_info or {}) guardrail_info["submitted_by_user_id"] = user_api_key_dict.user_id guardrail_info["submitted_by_email"] = user_api_key_dict.user_email @@ -847,7 +851,7 @@ def _row_to_submission_item(row: "LiteLLM_GuardrailsTable") -> GuardrailSubmissi guardrail_info: Final = _parse_json_field(row.guardrail_info) or {} team_guardrail: Final = row.team_id is not None - raw_params: Final = _parse_json_field(row.litellm_params) or {} + raw_params: Final = decrypt_guardrail_litellm_params(_parse_json_field(row.litellm_params) or {}) masked_params: Final = _get_masked_values(raw_params, unmasked_length=4, number_of_asterisks=4) return GuardrailSubmissionItem( guardrail_id=row.guardrail_id, @@ -1042,7 +1046,7 @@ async def approve_guardrail_submission( guardrail_dict: Final = { "guardrail_id": row.guardrail_id, "guardrail_name": row.guardrail_name, - "litellm_params": litellm_params, + "litellm_params": decrypt_guardrail_litellm_params(litellm_params), "guardrail_info": guardrail_info or {}, "team_id": row.team_id, } diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 0dc50cd6196..13bfc675a01 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -3,7 +3,7 @@ import asyncio import importlib import os -from collections.abc import Callable, Iterator, Mapping, Sequence +from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence from datetime import datetime, timezone from itertools import chain, count from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, cast @@ -14,12 +14,16 @@ import litellm from litellm import Router from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH, GUARDRAIL_ROTATION_ATTEMPTS from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.base_llm.guardrail_translation.utils import ( effective_scan_only_tool_results_for_guardrail, effective_skip_tool_message_for_guardrail, ) +from litellm.proxy.auth.master_key_boot_check import SALT_KEY_ENV_VAR +from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX, is_sensitive_callback_key +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockGuardrail, ) @@ -77,6 +81,95 @@ def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models return GuardrailsRepository(prisma_client).table +def _encrypted_param(key: str, value: object, new_encryption_key: str | None, depth: int = 0) -> object: + if depth > DEFAULT_MAX_RECURSE_DEPTH: + return value + if isinstance(value, dict): + return {k: _encrypted_param(k, v, new_encryption_key, depth + 1) for k, v in value.items()} + if isinstance(value, list): + return [_encrypted_param(key, item, new_encryption_key, depth + 1) for item in value] + if not ( + isinstance(value, str) + and value + and is_sensitive_callback_key(key) + and not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) + ): + return value + try: + return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value, new_encryption_key=new_encryption_key) + except Exception: # noqa: BLE001 # no salt key or master key configured: store the value as written + return value + + +def _decrypted_param(key: str, value: object, depth: int = 0) -> object: + if depth > DEFAULT_MAX_RECURSE_DEPTH: + return value + if isinstance(value, dict): + return {k: _decrypted_param(k, v, depth + 1) for k, v in value.items()} + if isinstance(value, list): + return [_decrypted_param(key, item, depth + 1) for item in value] + if not (isinstance(value, str) and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX)): + return value + decrypted: Final = decrypt_value_helper( + value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), + key=key, + exception_type="debug", + return_original_value=False, + ) + return value if decrypted is None else decrypted + + +def encrypt_guardrail_litellm_params( + litellm_params: Mapping[str, object], new_encryption_key: str | None = None +) -> dict[str, object]: + """Encrypt every string stored under a sensitive key (at any dict depth) for the guardrails table.""" + return {key: _encrypted_param(key, value, new_encryption_key) for key, value in litellm_params.items()} + + +def decrypt_guardrail_litellm_params(litellm_params: Mapping[str, object]) -> dict[str, object]: + """Decrypt values written by encrypt_guardrail_litellm_params; plaintext values pass through unchanged.""" + return {key: _decrypted_param(key, value) for key, value in litellm_params.items()} + + +def guardrail_from_db_row(row: Iterable[tuple[str, object]]) -> Guardrail: + """Build a Guardrail from a guardrails table row with its litellm_params decrypted.""" + fields: Final = dict(row) + stored_params: Final = fields.get("litellm_params") + if not isinstance(stored_params, Mapping): + return Guardrail(**fields) + return Guardrail(**{**fields, "litellm_params": decrypt_guardrail_litellm_params(stored_params)}) + + +async def _rotate_guardrail_row( + prisma_client: PrismaClient, + row: "prisma_models.LiteLLM_GuardrailsTable | None", + encryption_key: str, + attempts_left: int = GUARDRAIL_ROTATION_ATTEMPTS, +) -> int: + """Re-encrypt one row's params under encryption_key with a compare-and-set on updated_at. + A row edited since it was read is re-read and retried, up to attempts_left writes. Returns 1 when rewritten.""" + if row is None or not isinstance(row.litellm_params, Mapping): + return 0 + rotated_params: Final = encrypt_guardrail_litellm_params( + decrypt_guardrail_litellm_params(row.litellm_params), new_encryption_key=encryption_key + ) + if rotated_params == row.litellm_params: + return 0 + if await _guardrail_table(prisma_client).update_many( + where={"guardrail_id": row.guardrail_id, "updated_at": row.updated_at}, + data={"litellm_params": safe_dumps(rotated_params)}, + ): + return 1 + if attempts_left <= 1: + verbose_proxy_logger.warning( + "Guardrail %s kept changing during master key rotation; its secrets were not re-encrypted", + row.guardrail_id, + ) + return 0 + latest_row: Final = await _guardrail_table(prisma_client).find_unique(where={"guardrail_id": row.guardrail_id}) + return await _rotate_guardrail_row(prisma_client, latest_row, encryption_key, attempts_left - 1) + + guardrail_initializer_registry: Final = { SupportedGuardrailIntegrations.BEDROCK.value: initialize_bedrock, SupportedGuardrailIntegrations.LAKERA.value: initialize_lakera, @@ -295,7 +388,7 @@ class GuardrailRegistry: litellm_params_dict = litellm_params_obj.model_dump() else: litellm_params_dict = dict(litellm_params_obj) if litellm_params_obj else {} - litellm_params: Final[str] = safe_dumps(litellm_params_dict) + litellm_params: Final[str] = safe_dumps(encrypt_guardrail_litellm_params(litellm_params_dict)) guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {})) # Create guardrail in DB @@ -341,7 +434,7 @@ class GuardrailRegistry: litellm_params_dict = litellm_params_obj.model_dump() else: litellm_params_dict = dict(litellm_params_obj) if litellm_params_obj else {} - litellm_params: Final[str] = safe_dumps(litellm_params_dict) + litellm_params: Final[str] = safe_dumps(encrypt_guardrail_litellm_params(litellm_params_dict)) guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {})) # Update in DB @@ -357,8 +450,7 @@ class GuardrailRegistry: if updated_guardrail is None: raise ValueError(f"Guardrail not found, passed guardrail_id={guardrail_id}") - # Convert to dict and return - return dict(updated_guardrail) + return dict(guardrail_from_db_row(updated_guardrail)) except Exception as e: raise Exception(f"Error updating guardrail in DB: {e}") @@ -378,7 +470,7 @@ class GuardrailRegistry: guardrails: Final[list[Guardrail]] = [] for guardrail in guardrails_from_db: - guardrails.append(Guardrail(**(dict(guardrail)))) + guardrails.append(guardrail_from_db_row(guardrail)) return guardrails except Exception as e: @@ -394,7 +486,7 @@ class GuardrailRegistry: if not guardrail: return None - return Guardrail(**(dict(guardrail))) + return guardrail_from_db_row(guardrail) except Exception as e: raise Exception(f"Error getting guardrail from DB: {e}") @@ -410,10 +502,18 @@ class GuardrailRegistry: if not guardrail: return None - return Guardrail(**(dict(guardrail))) + return guardrail_from_db_row(guardrail) except Exception as e: raise Exception(f"Error getting guardrail from DB: {e}") + @staticmethod + async def rotate_guardrail_params_master_key(prisma_client: PrismaClient, new_master_key: str) -> int: + """Re-encrypt every guardrail row's sensitive litellm_params under the key the proxy decrypts with after the + rotation (LITELLM_SALT_KEY when set, otherwise new_master_key). Returns the number of rows rewritten.""" + encryption_key: Final = os.environ.get(SALT_KEY_ENV_VAR) or new_master_key + rows: Final = await _guardrail_table(prisma_client).find_many() + return sum([await _rotate_guardrail_row(prisma_client, row, encryption_key) for row in rows]) + def _apply_configured_bool_overrides(instance: CustomGuardrail, litellm_params: LitellmParams) -> None: """Override the parallel/raw-scan flags only when ``litellm_params`` explicitly diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d37dfe87ad5..cc8a47aa394 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -5303,6 +5303,16 @@ async def _rotate_master_key( data={"param_value": prisma.Json(encrypted_env_vars)}, ) + # 3b. process guardrails table + try: + from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry + + await GuardrailRegistry.rotate_guardrail_params_master_key( + prisma_client=prisma_client, new_master_key=new_master_key + ) + except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation + verbose_proxy_logger.warning("Failed to rotate guardrail params: %s", str(e)) + # 4. process MCP server table try: await rotate_mcp_server_credentials_master_key( diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 659dc438f2d..9667156d3a8 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -36,6 +36,9 @@ IGNORE_FUNCTIONS = [ "_collect_argument_paths", # max depth set. "_split_text", # max depth set. "_mask_sequence", # max depth set. + "_encrypted_param", # max depth set. + "_decrypted_param", # max depth set. + "_rotate_guardrail_row", # bounded by attempts_left. "_delete_nested_value_custom", # max depth set (bounded by number of path segments). "filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion. "__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion. diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py index ca2ff8bcce1..07287500121 100644 --- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py +++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py @@ -521,3 +521,39 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra assert await resync_task is True assert len(clean_agent_registry.agent_list) == 1 + + +@pytest.mark.asyncio +async def test_resync_guardrails_syncs_decrypted_litellm_params(monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.common_utils.registry_read_through as read_through_module + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import _resync_guardrails + from litellm.proxy.guardrails.guardrail_registry import ( + IN_MEMORY_GUARDRAIL_HANDLER, + encrypt_guardrail_litellm_params, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + encrypted_params: Final = encrypt_guardrail_litellm_params( + {"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "vendor-key"} + ) + prisma_client: Final = MagicMock() + prisma_client.db.litellm_guardrailstable.find_first = AsyncMock( + return_value={ + "guardrail_id": "enc-id", + "guardrail_name": "enc-guardrail", + "litellm_params": encrypted_params, + "guardrail_info": {}, + "status": "active", + } + ) + synced: list[dict] = [] + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + monkeypatch.setattr(IN_MEMORY_GUARDRAIL_HANDLER, "sync_guardrail_from_db", lambda guardrail: synced.append(guardrail)) + monkeypatch.setattr(read_through_module, "_initialized_guardrail", lambda guardrail_name: MagicMock()) + + assert await _resync_guardrails("enc-guardrail") is True + assert synced[0]["litellm_params"]["api_key"] == "vendor-key" diff --git a/tests/test_litellm/proxy/db/test_master_key_migration.py b/tests/test_litellm/proxy/db/test_master_key_migration.py index 9c0fc163b9f..948e8ac8518 100644 --- a/tests/test_litellm/proxy/db/test_master_key_migration.py +++ b/tests/test_litellm/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/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 508736fb78e..fd31c0d499d 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -2670,3 +2670,59 @@ async def test_test_custom_code_endpoint_reports_a_system_exit_as_an_execution_e assert response.error == "Execution error: SystemExit: bye" assert response.error_type == "execution" assert time.monotonic() - started < 2.0 + + +@pytest.mark.asyncio +async def test_team_guardrail_api_key_is_encrypted_at_rest_and_decrypted_on_review(mocker, monkeypatch): + """register stores the vendor api_key encrypted; submission detail masks the plaintext and approve loads it.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + mock_prisma = mocker.Mock() + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_guardrailstable.create = AsyncMock( + return_value=mocker.Mock( + guardrail_id="reg-enc", + guardrail_name="team-enc", + status="pending_review", + submitted_at=datetime.now(), + ) + ) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + request = RegisterGuardrailRequest( + guardrail_name="team-enc", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": "https://guardrails.example.com/validate", + "api_key": "team-vendor-secret-1234", + }, + ) + await register_guardrail(request, UserAPIKeyAuth(user_id="u1", team_id="team-1")) + + stored_params = json.loads(mock_prisma.db.litellm_guardrailstable.create.call_args[1]["data"]["litellm_params"]) + assert stored_params["api_key"].startswith("litellm_enc::") + assert "team-vendor-secret-1234" not in json.dumps(stored_params) + + row = mocker.Mock( + guardrail_id="reg-enc", + guardrail_name="team-enc", + status="pending_review", + team_id="team-1", + litellm_params=stored_params, + guardrail_info={}, + submitted_at=None, + reviewed_at=None, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.db.litellm_guardrailstable.update = AsyncMock() + admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + submission = await get_guardrail_submission("reg-enc", admin) + assert submission.litellm_params["api_key"] == "te****34" + + mock_handler = mocker.Mock() + mocker.patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", mock_handler) + await approve_guardrail_submission("reg-enc", admin) + loaded = mock_handler.initialize_guardrail.call_args.kwargs["guardrail"] + assert loaded["litellm_params"]["api_key"] == "team-vendor-secret-1234" diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 022fe85c779..8da050cfc84 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -1086,3 +1086,234 @@ def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance(): finally: for cb_list, snapshot in zip(lists, snapshots): cb_list[:] = snapshot + + +_ENCRYPTED_PREFIX = "litellm_enc::" + + +class _Row(dict): + """Prisma-row stand-in: dict(row) yields the columns and attributes read like a model.""" + + def __getattr__(self, name): + return self[name] + + +def _stored_params(create_or_update_mock: AsyncMock) -> dict: + import json + + return json.loads(create_or_update_mock.call_args.kwargs["data"]["litellm_params"]) + + +@pytest.mark.asyncio +async def test_add_guardrail_to_db_encrypts_sensitive_params_at_rest(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.create = AsyncMock(return_value=_Row(guardrail_id="g-1")) + + await GuardrailRegistry().add_guardrail_to_db( + guardrail=Guardrail( + guardrail_name="vendor", + litellm_params=LitellmParams( + guardrail="generic_guardrail_api", + mode="pre_call", + api_key="vendor-secret-key", + api_base="http://vendor.example", + aws_secret_access_key="aws-secret", + custom_headers={"Authorization": "Bearer header-secret", "x-tenant": "t1"}, + ), + ), + prisma_client=prisma_client, + ) + + stored = _stored_params(prisma_client.db.litellm_guardrailstable.create) + for leaked in ("vendor-secret-key", "aws-secret", "header-secret"): + assert leaked not in str(stored) + assert stored["api_key"].startswith(_ENCRYPTED_PREFIX) + assert stored["aws_secret_access_key"].startswith(_ENCRYPTED_PREFIX) + assert stored["custom_headers"]["Authorization"].startswith(_ENCRYPTED_PREFIX) + assert stored["custom_headers"]["x-tenant"] == "t1" + assert stored["guardrail"] == "generic_guardrail_api" + assert stored["mode"] == "pre_call" + assert stored["api_base"] == "http://vendor.example" + + +@pytest.mark.asyncio +async def test_get_all_guardrails_from_db_decrypts_new_rows_and_reads_legacy_plaintext(monkeypatch): + from litellm.proxy.guardrails.guardrail_registry import encrypt_guardrail_litellm_params + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + encrypted_row = _Row( + guardrail_id="g-new", + guardrail_name="new", + litellm_params=encrypt_guardrail_litellm_params( + {"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "new-key"} + ), + ) + legacy_row = _Row( + guardrail_id="g-legacy", + guardrail_name="legacy", + litellm_params={"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "legacy-key"}, + ) + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[encrypted_row, legacy_row]) + + guardrails = await GuardrailRegistry.get_all_guardrails_from_db(prisma_client=prisma_client) + + assert [g["litellm_params"]["api_key"] for g in guardrails] == ["new-key", "legacy-key"] + + +@pytest.mark.asyncio +async def test_update_guardrail_in_db_encrypts_and_returns_decrypted_row(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + prisma_client = MagicMock() + + async def _update(where, data): + import json + + return _Row( + guardrail_id=where["guardrail_id"], + guardrail_name="vendor", + litellm_params=json.loads(data["litellm_params"]), + ) + + prisma_client.db.litellm_guardrailstable.update = AsyncMock(side_effect=_update) + + result = await GuardrailRegistry().update_guardrail_in_db( + guardrail_id="g-1", + guardrail=Guardrail( + guardrail_name="vendor", + litellm_params={"guardrail": "generic_guardrail_api", "mode": "pre_call", "api_key": "rotated-key"}, + ), + prisma_client=prisma_client, + ) + + assert _stored_params(prisma_client.db.litellm_guardrailstable.update)["api_key"].startswith(_ENCRYPTED_PREFIX) + assert result["litellm_params"]["api_key"] == "rotated-key" + + +def test_encrypt_guardrail_litellm_params_does_not_double_encrypt(monkeypatch): + from litellm.proxy.guardrails.guardrail_registry import ( + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + params = { + "api_key": "k", + "default_on": True, + "auth_token": None, + "extra_headers": [{"x-api-key": "list-secret", "x-tenant": "t1"}], + } + encrypted = encrypt_guardrail_litellm_params(params) + + assert encrypted["extra_headers"][0]["x-api-key"].startswith(_ENCRYPTED_PREFIX) + assert encrypted["extra_headers"][0]["x-tenant"] == "t1" + assert encrypt_guardrail_litellm_params(encrypted) == encrypted + assert decrypt_guardrail_litellm_params(encrypted) == params + + +@pytest.mark.asyncio +async def test_rotate_guardrail_params_master_key_reencrypts_under_the_new_key(monkeypatch): + from litellm.proxy.guardrails.guardrail_registry import ( + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, + ) + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master") + stored = encrypt_guardrail_litellm_params({"guardrail": "bedrock", "mode": "pre_call", "api_key": "vendor-key"}) + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( + return_value=[_Row(guardrail_id="g-1", updated_at="2026-09-28T00:00:00Z", litellm_params=stored)] + ) + prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1) + + rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key( + prisma_client=prisma_client, new_master_key="sk-new-master" + ) + + rotated = _stored_params(prisma_client.db.litellm_guardrailstable.update_many) + assert rows_updated == 1 + assert prisma_client.db.litellm_guardrailstable.update_many.call_args.kwargs["where"] == { + "guardrail_id": "g-1", + "updated_at": "2026-09-28T00:00:00Z", + } + assert rotated["api_key"] != stored["api_key"] + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master") + assert decrypt_guardrail_litellm_params(rotated)["api_key"] == "vendor-key" + + +@pytest.mark.asyncio +async def test_rotate_guardrail_params_keeps_salt_key_encryption_when_salt_key_is_set(monkeypatch): + from litellm.proxy.guardrails.guardrail_registry import ( + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-guardrail-test") + stored = encrypt_guardrail_litellm_params({"guardrail": "bedrock", "api_key": "vendor-key"}) + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( + return_value=[_Row(guardrail_id="g-1", updated_at="t1", litellm_params=stored)] + ) + prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1) + + await GuardrailRegistry.rotate_guardrail_params_master_key(prisma_client=prisma_client, new_master_key="sk-new") + + rotated = _stored_params(prisma_client.db.litellm_guardrailstable.update_many) + assert decrypt_guardrail_litellm_params(rotated)["api_key"] == "vendor-key" + + +@pytest.mark.asyncio +async def test_rotate_guardrail_params_retries_a_row_edited_during_rotation(monkeypatch): + from litellm.proxy.guardrails.guardrail_registry import ( + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, + ) + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master") + snapshot = _Row( + guardrail_id="g-1", updated_at="t1", litellm_params=encrypt_guardrail_litellm_params({"api_key": "old-key"}) + ) + edited = _Row( + guardrail_id="g-1", updated_at="t2", litellm_params=encrypt_guardrail_litellm_params({"api_key": "edited-key"}) + ) + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[snapshot]) + prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=edited) + prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(side_effect=[0, 1]) + + rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key( + prisma_client=prisma_client, new_master_key="sk-new-master" + ) + + last_call = prisma_client.db.litellm_guardrailstable.update_many.call_args + assert rows_updated == 1 + assert last_call.kwargs["where"] == {"guardrail_id": "g-1", "updated_at": "t2"} + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master") + assert decrypt_guardrail_litellm_params(_stored_params(prisma_client.db.litellm_guardrailstable.update_many)) == { + "api_key": "edited-key" + } + + +@pytest.mark.asyncio +async def test_rotate_guardrail_params_gives_up_on_a_row_that_keeps_changing(monkeypatch): + from litellm.constants import GUARDRAIL_ROTATION_ATTEMPTS + from litellm.proxy.guardrails.guardrail_registry import encrypt_guardrail_litellm_params + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master") + row = _Row(guardrail_id="g-1", updated_at="t1", litellm_params=encrypt_guardrail_litellm_params({"api_key": "k"})) + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[row]) + prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=0) + + rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key( + prisma_client=prisma_client, new_master_key="sk-new-master" + ) + + assert rows_updated == 0 + assert prisma_client.db.litellm_guardrailstable.update_many.await_count == GUARDRAIL_ROTATION_ATTEMPTS + assert prisma_client.db.litellm_guardrailstable.find_unique.await_count == GUARDRAIL_ROTATION_ATTEMPTS - 1 diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index a5d2828dd9c..baa329e56da 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21102,3 +21102,60 @@ class TestTeamAdminMemberKeyBudgetUpdate: ) assert exc.value.status_code == 403 assert "member_key_budgets" not in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): + """Master-key rotation re-encrypts the guardrails table's litellm_params under the new key.""" + import json + from types import SimpleNamespace + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.guardrails.guardrail_registry import ( + decrypt_guardrail_litellm_params, + encrypt_guardrail_litellm_params, + ) + from litellm.proxy.management_endpoints import key_management_endpoints + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _rotate_master_key, + ) + + for rotator in ( + "rotate_mcp_server_credentials_master_key", + "rotate_mcp_user_credentials_master_key", + "rotate_mcp_user_env_vars_master_key", + "rotate_sso_identity_assertions_master_key", + ): + monkeypatch.setattr(key_management_endpoints, rotator, AsyncMock()) + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master-key") + guardrail_row = SimpleNamespace( + guardrail_id="g-1", + updated_at="t1", + litellm_params=encrypt_guardrail_litellm_params({"guardrail": "bedrock", "aws_secret_access_key": "aws-secret"}), + ) + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[guardrail_row]) + mock_prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1) + + await _rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"), + current_master_key="sk-old-master-key", + new_master_key="sk-new-master-key", + ) + + write = mock_prisma_client.db.litellm_guardrailstable.update_many.call_args.kwargs + stored_params = json.loads(write["data"]["litellm_params"]) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master-key") + assert write["where"] == {"guardrail_id": "g-1", "updated_at": "t1"} + assert stored_params["aws_secret_access_key"].startswith("litellm_enc::") + assert decrypt_guardrail_litellm_params(stored_params) == { + "guardrail": "bedrock", + "aws_secret_access_key": "aws-secret", + }