mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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>
This commit is contained in:
parent
6d8434f940
commit
5d42cb7cfa
12 changed files with 862 additions and 21 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue