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
This commit is contained in:
Yucheng He 2026-09-28 15:49:50 -07:00
parent 032dc13d26
commit 79fcef6256
3 changed files with 117 additions and 49 deletions

View file

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

View file

@ -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:

View file

@ -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"
}