mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): keep master key rotation going when pass-through header re-encryption fails
Wrap the pass-through header step of _rotate_master_key in the same try/except the MCP, SSO and credentials steps use, so a failed write logs a warning instead of aborting the remaining steps with a 500. Treat a marked header value that decrypts to an empty string as undecryptable, so a literal such as litellm_enc::*** is forwarded unchanged instead of as an empty header.
This commit is contained in:
parent
106f52e1c1
commit
9ca392868d
4 changed files with 60 additions and 6 deletions
|
|
@ -5353,7 +5353,10 @@ async def _rotate_master_key(
|
|||
)
|
||||
|
||||
if os.getenv(SALT_KEY_ENV_VAR) is None:
|
||||
await _reencrypt_pass_through_endpoint_headers(prisma_client, new_master_key)
|
||||
try:
|
||||
await _reencrypt_pass_through_endpoint_headers(prisma_client, new_master_key)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Failed to rotate pass-through endpoint headers: %s", str(e))
|
||||
|
||||
# 4. process MCP server table
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -43,11 +43,15 @@ def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str |
|
|||
|
||||
|
||||
def _decrypted(name: str, value: str) -> str | None:
|
||||
return decrypt_value_helper(
|
||||
value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX),
|
||||
key=name,
|
||||
exception_type="debug",
|
||||
return_original_value=False,
|
||||
"""Plaintext of a marked value; None when it does not decrypt or decrypts to an empty string."""
|
||||
return (
|
||||
decrypt_value_helper(
|
||||
value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX),
|
||||
key=name,
|
||||
exception_type="debug",
|
||||
return_original_value=False,
|
||||
)
|
||||
or None
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -21202,6 +21202,47 @@ async def test_rotate_master_key_leaves_pass_through_headers_under_salt_key(monk
|
|||
mock_prisma_client.db.execute_raw.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotate_master_key_continues_when_pass_through_header_reencryption_fails(monkeypatch):
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints import key_management_endpoints
|
||||
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
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=[MagicMock(param_name="general_settings", param_value={})]
|
||||
)
|
||||
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(
|
||||
key_management_endpoints,
|
||||
"_reencrypt_pass_through_endpoint_headers",
|
||||
AsyncMock(side_effect=RuntimeError("general_settings write fault")),
|
||||
)
|
||||
later_steps = {
|
||||
name: AsyncMock()
|
||||
for name 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",
|
||||
)
|
||||
}
|
||||
for name, step in later_steps.items():
|
||||
monkeypatch.setattr(key_management_endpoints, name, step)
|
||||
|
||||
await key_management_endpoints._rotate_master_key(
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"),
|
||||
current_master_key="sk-old-master-key",
|
||||
new_master_key="sk-new-master-key",
|
||||
)
|
||||
|
||||
assert all(step.await_count == 1 for step in later_steps.values())
|
||||
mock_prisma_client.db.litellm_credentialstable.find_many.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotate_master_key_reencrypts_pass_through_headers_on_the_writer(monkeypatch):
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||||
|
|
|
|||
|
|
@ -230,3 +230,9 @@ def test_undecryptable_pass_through_header_names(salt_key):
|
|||
def test_decrypt_pass_through_headers_keeps_a_bare_marker_literal(salt_key):
|
||||
assert decrypt_pass_through_headers({"x-tag": _ENC}) == {"x-tag": _ENC}
|
||||
assert undecryptable_pass_through_header_names({"x-tag": _ENC}) == frozenset()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [_ENC + "***", _ENC + "!!"])
|
||||
def test_decrypt_pass_through_headers_keeps_a_marked_value_that_decrypts_to_nothing(salt_key, value):
|
||||
assert decrypt_pass_through_headers({"x-tag": value}) == {"x-tag": value}
|
||||
assert undecryptable_pass_through_header_names({"x-tag": value}) == {"x-tag"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue