From f46e9959db59be9a2079f233038a4c55cdda8be2 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 16 Apr 2026 15:35:43 -0700 Subject: [PATCH 1/4] fix: restrict x-pass- header forwarding for protected header names --- litellm/passthrough/utils.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index ef4357d1ca2..8ae32d37a93 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -56,11 +56,15 @@ class BasePassthroughUtils: # Combine request headers with custom headers headers = {**request_headers, **headers} - # Always process x-pass- prefixed headers (strip prefix and forward) + # Process x-pass- prefixed headers (strip prefix and forward) + # Certain protocol-level and credential headers are excluded from this mechanism. + _PROTECTED_HEADERS = {"authorization", "api-key", "host", "content-length"} for header_name, header_value in request_headers.items(): if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX): # Strip the 'x-pass-' prefix to get the actual header name actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :] + if actual_header_name.lower() in _PROTECTED_HEADERS: + continue headers[actual_header_name] = header_value return headers From 5df7c21c9a56529b907e2ddb30a241b288e39bca Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 16 Apr 2026 15:41:34 -0700 Subject: [PATCH 2/4] fix: extend x-pass- header protection to cover additional credential headers and add tests - Move protected-headers set to module level as a frozenset - Add x-api-key, x-goog-api-key to protected set (provider credential headers) - Block x-amz- prefix to cover AWS SigV4 signing headers - Normalize forwarded header names to lowercase on write - Log at debug level when a protected header is skipped - Add unit test covering protected-header drop and non-protected forwarding --- litellm/passthrough/utils.py | 34 ++++++++++-- .../test_vertex_passthrough_load_balancing.py | 54 +++++++++++++++++++ 2 files changed, 83 insertions(+), 5 deletions(-) diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index 8ae32d37a93..5dde13f0078 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -3,8 +3,26 @@ from urllib.parse import parse_qs import httpx +from litellm._logging import verbose_logger from litellm.constants import PASS_THROUGH_HEADER_PREFIX +# Headers that must not be overwritten via the x-pass- forwarding mechanism. +# Includes standard credential/auth headers and protocol-level headers that +# affect routing or message framing. +_PASS_THROUGH_PROTECTED_HEADERS: frozenset = frozenset( + { + "authorization", + "api-key", + "x-api-key", + "x-goog-api-key", + "host", + "content-length", + } +) + +# Header name prefix used to block AWS SigV4 signing headers from being overridden. +_PASS_THROUGH_PROTECTED_HEADER_PREFIXES: tuple = ("x-amz-",) + class BasePassthroughUtils: @staticmethod @@ -57,13 +75,19 @@ class BasePassthroughUtils: headers = {**request_headers, **headers} # Process x-pass- prefixed headers (strip prefix and forward) - # Certain protocol-level and credential headers are excluded from this mechanism. - _PROTECTED_HEADERS = {"authorization", "api-key", "host", "content-length"} + # Credential and protocol-level headers are excluded from this mechanism. for header_name, header_value in request_headers.items(): if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX): - # Strip the 'x-pass-' prefix to get the actual header name - actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :] - if actual_header_name.lower() in _PROTECTED_HEADERS: + # Strip the 'x-pass-' prefix and normalize to lowercase + actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :].lower() + if actual_header_name in _PASS_THROUGH_PROTECTED_HEADERS or any( + actual_header_name.startswith(p) + for p in _PASS_THROUGH_PROTECTED_HEADER_PREFIXES + ): + verbose_logger.debug( + "x-pass- header %s maps to a protected header name; skipping", + header_name, + ) continue headers[actual_header_name] = header_value diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index eb4749549c2..b6dd714a232 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -451,6 +451,60 @@ def test_forward_headers_from_request_x_pass_prefix(): assert "x-pass-custom-header" not in result +def test_forward_headers_from_request_protected_headers_not_overwritten(): + """ + Test that x-pass- headers whose stripped names resolve to credential or + protocol-level header names are silently dropped and do not overwrite + values already present in the outbound headers dict. + """ + from litellm.passthrough.utils import BasePassthroughUtils + + proxy_headers = { + "authorization": "Bearer proxy-upstream-key", + "api-key": "proxy-azure-key", + "x-api-key": "proxy-anthropic-key", + "x-goog-api-key": "proxy-google-key", + } + + request_headers = { + "x-pass-authorization": "Bearer attacker-key", + "x-pass-api-key": "attacker-azure-key", + "x-pass-x-api-key": "attacker-anthropic-key", + "x-pass-x-goog-api-key": "attacker-google-key", + "x-pass-host": "evil.example.com", + "x-pass-content-length": "0", + "x-pass-x-amz-security-token": "attacker-aws-token", + # Legitimate x-pass- header that should still be forwarded + "x-pass-anthropic-beta": "context-1m-2025-08-07", + "content-type": "application/json", + } + + result = BasePassthroughUtils.forward_headers_from_request( + request_headers=request_headers, + headers=proxy_headers.copy(), + forward_headers=False, + ) + + # Protected headers must retain the proxy-configured values + assert result["authorization"] == "Bearer proxy-upstream-key" + assert result["api-key"] == "proxy-azure-key" + assert result["x-api-key"] == "proxy-anthropic-key" + assert result["x-goog-api-key"] == "proxy-google-key" + + # Protocol headers must not be injected + assert "host" not in result + assert "content-length" not in result + + # AWS SigV4 headers must not be injected + assert "x-amz-security-token" not in result + + # Legitimate non-protected x-pass- header still forwarded + assert result["anthropic-beta"] == "context-1m-2025-08-07" + + # Header name must be normalized to lowercase in output + assert "Anthropic-Beta" not in result + + @pytest.mark.asyncio async def test_vertex_passthrough_custom_model_name_replaced_in_url(): """ From 2ea3fafb68e1424673ea5eaf452016d2a2f8f0e5 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 16 Apr 2026 16:01:07 -0700 Subject: [PATCH 3/4] fix: tighten api_key value check in credential validation --- litellm/proxy/auth/auth_utils.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 64766bbaadd..33f4134865d 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -69,7 +69,8 @@ def check_complete_credentials(request_body: dict) -> bool: # complex credentials - easier to make a malicious request return False - if "api_key" in request_body: + api_key_value = request_body.get("api_key") + if api_key_value and isinstance(api_key_value, str) and api_key_value.strip(): return True return False From f8f356fac3bdfb8a52b51d4b9aec7a27c3ea5741 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 16 Apr 2026 16:08:47 -0700 Subject: [PATCH 4/4] test: add parametrized tests for api_key value handling in credential check --- tests/proxy_unit_tests/test_proxy_utils.py | 58 +++++++++++++++------- 1 file changed, 41 insertions(+), 17 deletions(-) diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 9f5f14457e8..1de0ab450d7 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -19,7 +19,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import litellm from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth -from litellm.proxy.auth.auth_utils import is_request_body_safe +from litellm.proxy.auth.auth_utils import ( + check_complete_credentials, + is_request_body_safe, +) from litellm.proxy.litellm_pre_call_utils import ( _get_dynamic_logging_metadata, add_litellm_data_to_request, @@ -33,7 +36,9 @@ def mock_request(monkeypatch): mock_request = Mock(spec=Request) mock_request.query_params = {} # Set mock query_params to an empty dictionary mock_request.headers = {"traceparent": "test_traceparent"} - mock_request.state = State() # Real State so _safe_get_request_headers caching works + mock_request.state = ( + State() + ) # Real State so _safe_get_request_headers caching works monkeypatch.setattr( "litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request", mock_request ) @@ -465,6 +470,21 @@ def test_is_request_body_safe_model_enabled( assert expect_error == error_raised +@pytest.mark.parametrize( + "api_key_value, expect_complete", + [ + ("sk-real-key", True), + ("", False), + (None, False), + (" ", False), + ], +) +def test_check_complete_credentials_api_key_values(api_key_value, expect_complete): + request_body = {"model": "gpt-3.5-turbo", "api_key": api_key_value} + result = check_complete_credentials(request_body=request_body) + assert result == expect_complete + + def test_reading_openai_org_id_from_headers(): from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup @@ -735,6 +755,7 @@ def test_get_docs_url(env_vars, expected_url): result = _get_docs_url() assert result == expected_url + @pytest.mark.parametrize( "env_vars, expected_url", [ @@ -1516,7 +1537,7 @@ class MockPrismaClientDB: mock_key_data, ): self.db = MockDb(mock_team_data, mock_key_data) - + async def get_data( self, token: Optional[Union[str, list]] = None, @@ -1534,7 +1555,7 @@ class MockPrismaClientDB: ): """Mock get_data method to return user info for admin""" from litellm.proxy._types import LiteLLM_UserTable - + # Return a proper LiteLLM_UserTable object when querying by user_id if user_id: return LiteLLM_UserTable( @@ -2072,7 +2093,7 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch): monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False) import litellm.proxy.litellm_pre_call_utils as pre_call_utils from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists - + # Reset module-level cache to ensure test isolation pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None @@ -2097,7 +2118,7 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch): def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch): import litellm.proxy.litellm_pre_call_utils as pre_call_utils from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists - + # Reset module-level cache to ensure test isolation pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None @@ -2394,16 +2415,17 @@ async def test_handle_logging_proxy_only_error_syncs_normalized_call_type( captured_logging_obj["logging_obj"] = logging_obj return logging_obj, data - with patch( - "litellm.proxy.utils.litellm.utils.function_setup", - side_effect=_capture_function_setup, - ), patch.object( - Logging, "async_failure_handler", new=AsyncMock(return_value=None) - ), patch.object( - Logging, "failure_handler", return_value=None - ), patch( - "litellm.proxy.utils.threading.Thread" - ) as mock_thread: + with ( + patch( + "litellm.proxy.utils.litellm.utils.function_setup", + side_effect=_capture_function_setup, + ), + patch.object( + Logging, "async_failure_handler", new=AsyncMock(return_value=None) + ), + patch.object(Logging, "failure_handler", return_value=None), + patch("litellm.proxy.utils.threading.Thread") as mock_thread, + ): mock_thread.return_value.start = Mock() await proxy_logging._handle_logging_proxy_only_error( @@ -2647,7 +2669,9 @@ async def test_handle_logging_proxy_only_error_skips_handlers_for_pass_through() "model": "claude-3-5-sonnet", } - with patch.object(logging_obj, "async_failure_handler", new_callable=AsyncMock) as mock_async: + with patch.object( + logging_obj, "async_failure_handler", new_callable=AsyncMock + ) as mock_async: with patch.object(logging_obj, "failure_handler") as mock_sync: await proxy_logging._handle_logging_proxy_only_error( request_data=request_data,