From 79fcef6256600bf570bbe26b6fe858e443a21236 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 15:49:50 -0700 Subject: [PATCH] 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 --- litellm/constants.py | 1 + .../proxy/guardrails/guardrail_registry.py | 111 ++++++++++-------- .../guardrails/test_guardrail_registry.py | 54 +++++++++ 3 files changed, 117 insertions(+), 49 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 5675a7e1ea5..9b0a00040dc 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/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 89d9cee31d3..5ba71a7fa12 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -14,13 +14,14 @@ import litellm from litellm import Router from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid -from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH +from litellm.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 ( @@ -83,40 +84,39 @@ def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models def _encrypted_param(key: str, value: object, new_encryption_key: str | None, depth: int = 0) -> object: if depth > DEFAULT_MAX_RECURSE_DEPTH: return value - match value: - case dict(): - return {k: _encrypted_param(k, v, new_encryption_key, depth + 1) for k, v in value.items()} - case list(): - return [_encrypted_param(key, item, new_encryption_key, depth + 1) for item in value] - case str() if value and is_sensitive_callback_key(key) and not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): - try: - return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper( - value, new_encryption_key=new_encryption_key - ) - except Exception: # noqa: BLE001 # no salt key or master key configured: store the value as written - return value - case _: - return value + 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 - match value: - case dict(): - return {k: _decrypted_param(k, v, depth + 1) for k, v in value.items()} - case list(): - return [_decrypted_param(key, item, depth + 1) for item in value] - case str() if value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): - decrypted: Final = decrypt_value_helper( - value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), - key=key, - exception_type="debug", - return_original_value=False, - ) - return value if decrypted is None else decrypted - case _: - return value + 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( @@ -135,9 +135,33 @@ def guardrail_from_db_row(row: Iterable[tuple[str, object]]) -> Guardrail: """Build a Guardrail from a guardrails table row with its litellm_params decrypted.""" fields: Final = dict(row) stored_params: Final = fields.get("litellm_params") - if isinstance(stored_params, Mapping): - fields["litellm_params"] = decrypt_guardrail_litellm_params(stored_params) - return Guardrail(**fields) + 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", encryption_key: str +) -> int: + current: prisma_models.LiteLLM_GuardrailsTable | None = row + for _ in range(GUARDRAIL_ROTATION_ATTEMPTS): + if current is None or not isinstance(current.litellm_params, Mapping): + return 0 + rotated_params: dict[str, object] = encrypt_guardrail_litellm_params( + decrypt_guardrail_litellm_params(current.litellm_params), new_encryption_key=encryption_key + ) + if rotated_params == current.litellm_params: + return 0 + if await _guardrail_table(prisma_client).update_many( + where={"guardrail_id": current.guardrail_id, "updated_at": current.updated_at}, + data={"litellm_params": safe_dumps(rotated_params)}, + ): + return 1 + current = await _guardrail_table(prisma_client).find_unique(where={"guardrail_id": row.guardrail_id}) + verbose_proxy_logger.warning( + "Guardrail %s kept changing during master key rotation; its secrets were not re-encrypted", row.guardrail_id + ) + return 0 guardrail_initializer_registry: Final = { @@ -478,22 +502,11 @@ class GuardrailRegistry: @staticmethod async def rotate_guardrail_params_master_key(prisma_client: PrismaClient, new_master_key: str) -> int: - """Re-encrypt the sensitive litellm_params of every guardrail row under new_master_key. Returns rows updated.""" - rows_updated = 0 - for row in await _guardrail_table(prisma_client).find_many(): - stored_params = row.litellm_params - if not isinstance(stored_params, Mapping): - continue - rotated_params = encrypt_guardrail_litellm_params( - decrypt_guardrail_litellm_params(stored_params), new_encryption_key=new_master_key - ) - if rotated_params == stored_params: - continue - rows_updated += await _guardrail_table(prisma_client).update_many( - where={"guardrail_id": row.guardrail_id, "updated_at": row.updated_at}, - data={"litellm_params": safe_dumps(rotated_params)}, - ) - return rows_updated + """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: diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 04fed805e38..45955270d5d 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -1239,3 +1239,57 @@ async def test_rotate_guardrail_params_master_key_reencrypts_under_the_new_key(m 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" + }