mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
032dc13d26
commit
79fcef6256
3 changed files with 117 additions and 49 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue