mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge bd2090a85a into 6cf51383bf
This commit is contained in:
commit
8f4dcd0648
12 changed files with 539 additions and 14 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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue