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:
Yucheng He 2026-09-29 10:43:27 -07:00
parent 106f52e1c1
commit 9ca392868d
4 changed files with 60 additions and 6 deletions

View file

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

View file

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

View file

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

View file

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