diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index 706beb7dc5e..fa4bdb7dc7a 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -56,7 +56,7 @@ class BasePassthroughUtils: request_headers: dict, headers: dict, forward_headers: Optional[bool] = False, - ): + ) -> dict: """ Helper to forward headers from original request. @@ -64,34 +64,36 @@ class BasePassthroughUtils: with the prefix stripped, regardless of forward_headers setting. e.g., 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value' """ + x_pass_entries = tuple( + (name, value) + for name, value in request_headers.items() + if name.lower().startswith(PASS_THROUGH_HEADER_PREFIX) + ) + if forward_headers is True: - # Header We Should NOT forward - request_headers.pop("content-length", None) - request_headers.pop("host", None) + for header_name in list(request_headers.keys()): + lower = header_name.lower() + if lower in _PASS_THROUGH_PROTECTED_HEADERS or lower.startswith(PASS_THROUGH_HEADER_PREFIX): + request_headers.pop(header_name, None) custom_header_names = {header_name.lower() for header_name in headers} for header_name in list(request_headers.keys()): if header_name.lower() in custom_header_names: request_headers.pop(header_name, None) - # Combine request headers with custom headers headers = {**request_headers, **headers} - # Process x-pass- prefixed headers (strip prefix and forward) - # 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 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 + for header_name, header_value in x_pass_entries: + 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 return headers diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 0ac4182ebd2..4de257d518f 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -1833,13 +1833,9 @@ async def websocket_passthrough_request( upstream_headers = custom_headers.copy() if forward_headers: - # Forward relevant headers from the incoming request incoming_headers = dict(websocket.headers) for header_name, header_value in incoming_headers.items(): - # Only forward certain headers to avoid conflicts if header_name.lower() in [ - "authorization", - "x-api-key", "x-goog-user-project", ]: upstream_headers[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 aaf1dad4910..ad2817ba632 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 @@ -61,7 +61,6 @@ async def test_vertex_passthrough_load_balancing(): new_callable=AsyncMock, ) as mock_auth, ): - # Setup additional mocks to avoid side effects mock_pt_router.get_vertex_credentials.return_value = MagicMock() mock_prep_headers.return_value = ( @@ -86,9 +85,7 @@ async def test_vertex_passthrough_load_balancing(): # Verify # 1. Check that get_available_deployment_for_pass_through was called with the correct model ID - mock_router.get_available_deployment_for_pass_through.assert_called_once_with( - model="gemini-pro" - ) + mock_router.get_available_deployment_for_pass_through.assert_called_once_with(model="gemini-pro") # 2. Check that get_model_list was NOT called (this ensures we aren't doing the old logic) mock_router.get_model_list.assert_not_called() @@ -210,9 +207,7 @@ def test_get_available_deployment_for_pass_through_load_balancing(): # Call multiple times and track selected deployments selections = {"project-1": 0, "project-2": 0} for _ in range(100): - deployment = router.get_available_deployment_for_pass_through( - model="gemini-pro" - ) + deployment = router.get_available_deployment_for_pass_through(model="gemini-pro") project = deployment["litellm_params"]["vertex_project"] selections[project] += 1 @@ -241,9 +236,7 @@ async def test_async_get_available_deployment_for_pass_through(): router = Router(model_list=model_list, routing_strategy="simple-shuffle") - deployment = await router.async_get_available_deployment_for_pass_through( - model="gemini-pro", request_kwargs={} - ) + deployment = await router.async_get_available_deployment_for_pass_through(model="gemini-pro", request_kwargs={}) assert deployment is not None assert deployment["litellm_params"]["use_in_pass_through"] is True @@ -307,7 +300,6 @@ async def test_vertex_passthrough_forwards_anthropic_beta_header(): return_value=("new-access-token", None), ) as mock_get_token, ): - # Call the function ( headers, @@ -402,7 +394,6 @@ async def test_vertex_passthrough_does_not_forward_litellm_auth_token(): return_value=("vertex-access-token", None), ), ): - ( headers, _base_target_url, @@ -621,7 +612,6 @@ async def test_vertex_passthrough_custom_model_name_replaced_in_url(): new_callable=AsyncMock, ) as mock_auth, ): - mock_pt_router.get_vertex_credentials.return_value = MagicMock() mock_prep_headers.return_value = ( {}, @@ -634,9 +624,7 @@ async def test_vertex_passthrough_custom_model_name_replaced_in_url(): mock_create_route.return_value = mock_endpoint_func mock_auth.return_value = {} - mock_handler.get_default_base_target_url.return_value = ( - "https://global-aiplatform.googleapis.com" - ) + mock_handler.get_default_base_target_url.return_value = "https://global-aiplatform.googleapis.com" await _base_vertex_proxy_route( endpoint=test_endpoint, @@ -646,17 +634,78 @@ async def test_vertex_passthrough_custom_model_name_replaced_in_url(): ) # Verify the router was called with the custom model name (extracted from URL) - mock_router.get_available_deployment_for_pass_through.assert_called_once_with( - model="gcp/google/gemini-3-pro" - ) + mock_router.get_available_deployment_for_pass_through.assert_called_once_with(model="gcp/google/gemini-3-pro") # Verify the target URL passed to create_pass_through_route contains # the REAL Vertex AI model name, not the custom one create_route_call = mock_create_route.call_args target_url = create_route_call.kwargs.get("target", "") - assert ( - "gcp/google/gemini-3-pro" not in target_url - ), f"Custom model name should have been replaced in target URL. Got: {target_url}" - assert ( - "gemini-3-pro" in target_url - ), f"Actual Vertex AI model name should be in target URL. Got: {target_url}" + assert "gcp/google/gemini-3-pro" not in target_url, ( + f"Custom model name should have been replaced in target URL. Got: {target_url}" + ) + assert "gemini-3-pro" in target_url, f"Actual Vertex AI model name should be in target URL. Got: {target_url}" + + +def test_forward_headers_strips_proxy_auth_when_forwarding(): + """ + Regression test for #32202: when forward_headers=True and custom headers + do NOT include an Authorization key, the proxy's own auth header must + still be stripped so it is never forwarded to the upstream provider. + """ + from litellm.passthrough.utils import BasePassthroughUtils + + request_headers = { + "authorization": "Bearer sk-litellm-master", + "api-key": "sk-litellm-azure", + "x-api-key": "sk-litellm-anthropic", + "x-goog-api-key": "sk-litellm-google", + "content-type": "application/json", + "x-request-id": "req-456", + } + custom_headers = { + "x-custom": "custom-value", + } + + result = BasePassthroughUtils.forward_headers_from_request( + request_headers=request_headers.copy(), + headers=custom_headers.copy(), + forward_headers=True, + ) + + assert "authorization" not in result + assert "api-key" not in result + assert "x-api-key" not in result + assert "x-goog-api-key" not in result + assert result["x-custom"] == "custom-value" + assert result["content-type"] == "application/json" + assert result["x-request-id"] == "req-456" + + +def test_forward_headers_no_raw_x_pass_prefix_in_output(): + """ + Regression test for #32202: when forward_headers=True, x-pass- prefixed + headers must not appear verbatim in the forwarded set. Only the + prefix-stripped version should be present. + """ + from litellm.passthrough.utils import BasePassthroughUtils + + request_headers = { + "authorization": "Bearer sk-litellm-key", + "x-pass-anthropic-beta": "prompt-caching-2025-04", + "x-pass-custom-header": "custom-value", + "content-type": "application/json", + } + custom_headers = {"x-api-key": "sk-real-provider-key"} + + result = BasePassthroughUtils.forward_headers_from_request( + request_headers=request_headers.copy(), + headers=custom_headers.copy(), + forward_headers=True, + ) + + assert "x-pass-anthropic-beta" not in result + assert "x-pass-custom-header" not in result + assert result["anthropic-beta"] == "prompt-caching-2025-04" + assert result["custom-header"] == "custom-value" + assert "authorization" not in result + assert result["x-api-key"] == "sk-real-provider-key"