From 9ca392868d677cb9e8d9ee81e7ffbafe7188ae87 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 10:43:27 -0700 Subject: [PATCH] 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. --- .../key_management_endpoints.py | 5 ++- .../pass_through_endpoints/common_utils.py | 14 ++++--- .../test_key_management_endpoints.py | 41 +++++++++++++++++++ ...test_passthrough_endpoints_common_utils.py | 6 +++ 4 files changed, 60 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d7f9e644bba..5ee35643c54 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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: diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index 251e2fdac8a..a27bfd4abe7 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -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 ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 4b3fdeac9d7..3a9a1537f15 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py index a96ef88b5aa..5e579e4d2b1 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -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"}