From da41bc40b38964c5466773ac8f5579789a0182c9 Mon Sep 17 00:00:00 2001 From: Milan Date: Sat, 9 May 2026 11:32:02 +0300 Subject: [PATCH] Fix Bedrock KB pass-through SigV4 headers and signed body Coerce botocore HeadersDict to a dict for pass-through routes. When forward_headers is true, drop request headers that collide case-insensitively with signed headers so client Bearer auth does not shadow AWS SigV4. Send prepped.body as raw content so the outbound payload matches the signature after logging hooks mutate the parsed dict. Co-authored-by: Cursor --- litellm/passthrough/utils.py | 5 ++ .../llm_passthrough_endpoints.py | 4 ++ .../pass_through_endpoints.py | 72 +++++++++++++++---- .../pass_through_endpoints.py | 4 ++ .../test_pass_through_endpoints.py | 15 +++- .../test_vertex_passthrough_load_balancing.py | 30 ++++++++ 6 files changed, 116 insertions(+), 14 deletions(-) diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index d39a0dda152..9484922833a 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -71,6 +71,11 @@ class BasePassthroughUtils: request_headers.pop("content-length", None) request_headers.pop("host", 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} diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index ce103f806e1..7ca28a5d4ac 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -44,6 +44,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( ) from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, + LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, ) from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( @@ -1123,6 +1124,9 @@ async def bedrock_proxy_route( _forward_headers=True, ) # dynamically construct pass-through endpoint based on incoming path setattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, data) + # SigV4 signs an exact payload; pass-through must send prepped.body, not json.dumps + # of a dict that hooks may mutate (logging_obj, metadata, etc.). + setattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, prepped.body) received_value = await endpoint_func( request, fastapi_response, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index cc6c26fdf90..d3f9973473c 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -6,7 +6,7 @@ import posixpath import traceback from base64 import b64encode from datetime import datetime -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast from urllib.parse import urlencode, urlparse import httpx @@ -62,6 +62,7 @@ from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.pass_through_endpoints import ( EndpointType, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, + LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, PassthroughStandardLoggingPayload, ) @@ -375,6 +376,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): headers: dict, requested_query_params: Optional[dict] = None, custom_body: Optional[dict] = None, + custom_raw_body: Optional[Union[str, bytes]] = None, ) -> httpx.Response: """ Make a non-streaming HTTP request @@ -388,6 +390,14 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): headers=headers, params=requested_query_params, ) + elif custom_raw_body is not None: + response = await async_client.request( + method=request.method, + url=url, + headers=headers, + params=requested_query_params, + content=custom_raw_body, + ) else: response = await async_client.request( method=request.method, @@ -406,6 +416,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): headers: dict, requested_query_params: Optional[dict] = None, _parsed_body: Optional[dict] = None, + custom_raw_body: Optional[Union[str, bytes]] = None, forward_multipart: bool = False, ) -> httpx.Response: """ @@ -420,6 +431,14 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): headers=headers, params=requested_query_params, ) + elif custom_raw_body is not None: + response = await async_client.request( + method=request.method, + url=url, + headers=headers, + params=requested_query_params, + content=custom_raw_body, + ) elif ( HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart @@ -659,6 +678,7 @@ async def pass_through_request( # noqa: PLR0915 custom_headers: dict, user_api_key_dict: UserAPIKeyAuth, custom_body: Optional[dict] = None, + custom_raw_body: Optional[Union[str, bytes]] = None, forward_headers: Optional[bool] = False, merge_query_params: Optional[bool] = False, query_params: Optional[dict] = None, @@ -677,6 +697,9 @@ async def pass_through_request( # noqa: PLR0915 custom_headers: The custom headers user_api_key_dict: The user API key dictionary custom_body: The custom body + custom_raw_body: Exact request body bytes/str for upstream (e.g. SigV4-signed). + When set, this is sent as ``content=...`` instead of re-encoding ``custom_body`` + as JSON, so signatures and Content-Length stay consistent. forward_headers: Whether to forward headers merge_query_params: Whether to merge query params query_params: The query params @@ -738,10 +761,12 @@ async def pass_through_request( # noqa: PLR0915 # Skip body parsing for multipart requests - make_multipart_http_request will handle it # But if custom_body is provided (e.g., JSON parsed despite multipart content-type), use it is_multipart = ( - HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body + HttpPassThroughEndpointHelpers.is_multipart(request) + and custom_body is None + and custom_raw_body is None ) - if custom_body: + if custom_body is not None: _parsed_body = custom_body elif is_multipart: # Don't parse multipart body here - it will be handled by make_multipart_http_request @@ -883,13 +908,22 @@ async def pass_through_request( # noqa: PLR0915 ) ) else: - req = async_client.build_request( - "POST", - url, - json=_parsed_body, - params=requested_query_params, - headers=headers, - ) + if custom_raw_body is not None: + req = async_client.build_request( + "POST", + url, + content=custom_raw_body, + params=requested_query_params, + headers=headers, + ) + else: + req = async_client.build_request( + "POST", + url, + json=_parsed_body, + params=requested_query_params, + headers=headers, + ) response = await async_client.send(req, stream=stream) @@ -925,6 +959,7 @@ async def pass_through_request( # noqa: PLR0915 headers=headers, requested_query_params=requested_query_params, _parsed_body=_parsed_body, + custom_raw_body=custom_raw_body, forward_multipart=is_multipart, ) ) @@ -1225,7 +1260,7 @@ async def _parse_request_data_by_content_type( def create_pass_through_route( endpoint, target: str, - custom_headers: Optional[dict] = None, + custom_headers: Optional[Mapping[str, Any]] = None, _forward_headers: Optional[bool] = False, _merge_query_params: Optional[bool] = False, dependencies: Optional[List] = None, @@ -1334,9 +1369,12 @@ def create_pass_through_route( ) ) - # Ensure custom_headers is a dict + # Ensure custom_headers is a dict. Botocore returns a HeadersDict + # for SigV4-prepared requests, which is a Mapping but not a dict. headers_dict = ( - param_custom_headers if isinstance(param_custom_headers, dict) else {} + dict(param_custom_headers) + if isinstance(param_custom_headers, Mapping) + else {} ) # Ensure query_params and custom_body are dicts or None @@ -1352,6 +1390,11 @@ def create_pass_through_route( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, None, ) + state_raw_body: Optional[Union[str, bytes]] = getattr( + request.state, + LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, + None, + ) final_custom_body: Optional[dict] = None if isinstance(state_custom_body, dict): final_custom_body = state_custom_body @@ -1372,6 +1415,7 @@ def create_pass_through_route( ), stream=is_streaming_request or stream, custom_body=final_custom_body, + custom_raw_body=state_raw_body, cost_per_request=cast(Optional[float], param_cost_per_request), custom_llm_provider=custom_llm_provider, guardrails_config=cast(Optional[dict], param_guardrails), @@ -1379,6 +1423,8 @@ def create_pass_through_route( finally: if hasattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY): delattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY) + if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY): + delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY) return endpoint_func diff --git a/litellm/types/passthrough_endpoints/pass_through_endpoints.py b/litellm/types/passthrough_endpoints/pass_through_endpoints.py index 4a07fa5e849..3524a7eb7f7 100644 --- a/litellm/types/passthrough_endpoints/pass_through_endpoints.py +++ b/litellm/types/passthrough_endpoints/pass_through_endpoints.py @@ -7,6 +7,10 @@ from typing_extensions import TypedDict # JSON without a FastAPI `custom_body` parameter (which would consume the HTTP body). LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY = "litellm_pass_through_custom_body" +# Request.state key for programmatic pass-through callers that must preserve an +# exact byte/string body, such as AWS SigV4-signed requests. +LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY = "litellm_pass_through_raw_body" + class EndpointType(str, Enum): VERTEX_AI = "vertex-ai" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 97a21136198..f9568478785 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -18,6 +18,7 @@ sys.path.insert( from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( HttpPassThroughEndpointHelpers, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, + LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, pass_through_request, ) from litellm.proxy.pass_through_endpoints.success_handler import ( @@ -2153,7 +2154,12 @@ async def test_create_pass_through_route_custom_body_url_target(): endpoint_func = create_pass_through_route( endpoint=unique_path, target="https://bedrock-agent-runtime.us-east-1.amazonaws.com", - custom_headers={"Content-Type": "application/json"}, + custom_headers=Headers( + { + "Authorization": "AWS4-HMAC-SHA256 signed", + "Content-Type": "application/json", + } + ), _forward_headers=True, ) @@ -2200,6 +2206,8 @@ async def test_create_pass_through_route_custom_body_url_target(): setattr( mock_request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, bedrock_body ) + signed_body = json.dumps(bedrock_body) + setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, signed_body) await endpoint_func( request=mock_request, @@ -2213,6 +2221,11 @@ async def test_create_pass_through_route_custom_body_url_target(): # The critical assertion: custom_body takes precedence over # the body parsed from the raw request assert call_kwargs["custom_body"] == bedrock_body + assert call_kwargs["custom_raw_body"] == signed_body + assert call_kwargs["custom_headers"] == { + "authorization": "AWS4-HMAC-SHA256 signed", + "content-type": "application/json", + } @pytest.mark.asyncio 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 7176cf455c8..aaf1dad4910 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 @@ -538,6 +538,36 @@ def test_forward_headers_from_request_protected_headers_not_overwritten(): assert "Anthropic-Beta" not in result +def test_forward_headers_custom_wins_case_insensitive_over_request_authorization(): + """ + When forwarding request headers, provider-signed/custom headers must win + even if the incoming request uses a different case for the same header name. + """ + from litellm.passthrough.utils import BasePassthroughUtils + + request_headers = { + "authorization": "Bearer sk-litellm-key", + "content-type": "application/json", + "x-request-id": "req-123", + } + signed_headers = { + "Authorization": "AWS4-HMAC-SHA256 signed", + "Content-Type": "application/json", + } + + result = BasePassthroughUtils.forward_headers_from_request( + request_headers=request_headers, + headers=signed_headers.copy(), + forward_headers=True, + ) + + assert result["Authorization"] == "AWS4-HMAC-SHA256 signed" + assert "authorization" not in result + assert result["Content-Type"] == "application/json" + assert "content-type" not in result + assert result["x-request-id"] == "req-123" + + @pytest.mark.asyncio async def test_vertex_passthrough_custom_model_name_replaced_in_url(): """