diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index d1854276e90..96b1c844bf6 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -294,7 +294,8 @@ class AzureAIStudioConfig(OpenAIConfig): def _drop_tool_level_extra_fields(self, request_data: dict, error_text: str) -> dict: from litellm.llms.base_llm.base_utils import parse_rejected_tool_fields - fields_to_drop = frozenset().union(*parse_rejected_tool_fields(error_text).values() or (frozenset(),)) + rejected = parse_rejected_tool_fields(error_text) + fields_to_drop = frozenset(field for fields in rejected.values() for field in fields) tools = request_data.get("tools") if fields_to_drop and isinstance(tools, list): for tool in tools: diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 9800df38614..7fa77b1dba1 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -147,6 +147,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials: Credentials, aws_region_name: str, caller_headers: Mapping[str, str], + extra_headers: Mapping[str, str] | None, endpoint_url: str, api_key: str | None, ) -> tuple[_SendResultT, str]: @@ -168,7 +169,7 @@ class BedrockConverseLLM(BaseAWSLLM): error_text=_provider_error_text(err), credentials=credentials, aws_region_name=aws_region_name, - extra_headers=caller_headers, + extra_headers=extra_headers, endpoint_url=endpoint_url, headers=caller_headers, api_key=api_key, @@ -188,6 +189,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials: Credentials, aws_region_name: str, caller_headers: Mapping[str, str], + extra_headers: Mapping[str, str] | None, endpoint_url: str, api_key: str | None, ) -> tuple[_SendResultT, str]: @@ -200,7 +202,7 @@ class BedrockConverseLLM(BaseAWSLLM): error_text=_provider_error_text(err), credentials=credentials, aws_region_name=aws_region_name, - extra_headers=caller_headers, + extra_headers=extra_headers, endpoint_url=endpoint_url, headers=caller_headers, api_key=api_key, @@ -283,6 +285,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials=credentials, aws_region_name=litellm_params.get("aws_region_name") or "us-west-2", caller_headers=headers, + extra_headers=headers, endpoint_url=api_base, api_key=api_key, ) @@ -377,6 +380,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials=credentials, aws_region_name=litellm_params.get("aws_region_name") or "us-west-2", caller_headers=caller_headers, + extra_headers=caller_headers, endpoint_url=api_base, api_key=api_key, ) @@ -619,6 +623,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials=credentials, aws_region_name=aws_region_name, caller_headers=headers, + extra_headers=extra_headers, endpoint_url=proxy_endpoint_url, api_key=api_key, ) @@ -656,6 +661,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials=credentials, aws_region_name=aws_region_name, caller_headers=headers, + extra_headers=extra_headers, endpoint_url=proxy_endpoint_url, api_key=api_key, ) diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_rejected_tool_field_retry.py b/tests/test_litellm/llms/bedrock/chat/test_converse_rejected_tool_field_retry.py index b13671f9bfd..9505650134d 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_rejected_tool_field_retry.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_rejected_tool_field_retry.py @@ -49,6 +49,7 @@ def _retry_kwargs(): "credentials": _credentials(), "aws_region_name": "us-east-1", "caller_headers": {"Content-Type": "application/json"}, + "extra_headers": None, "endpoint_url": "https://bedrock-runtime.us-east-1.amazonaws.com/model/m/converse", "api_key": None, } @@ -191,6 +192,27 @@ async def test_async_retry_resends_without_the_rejected_field_and_resigns() -> N assert sent_body == attempts[1][0] +def test_retry_preserves_a_caller_supplied_authorization_header() -> None: + """``extra_headers`` is not a duplicate of ``caller_headers``: it is the only thing + that restores a caller's non-SigV4 ``Authorization`` after signing, so the retry has + to pass it through or a proxied bearer token is silently replaced by a SigV4 one.""" + bearer = {"Authorization": "Bearer caller-supplied-token"} + attempts: list[dict] = [] + + def send(body: str, headers) -> str: + attempts.append(dict(headers)) + if len(attempts) == 1: + raise BedrockError(status_code=400, message=_STRICT_REJECTION) + return "ok" + + BedrockConverseLLM()._send_retrying_rejected_tool_fields( + send=send, + **{**_retry_kwargs(), "caller_headers": {"Content-Type": "application/json", **bearer}, "extra_headers": bearer}, + ) + + assert attempts[1]["Authorization"] == "Bearer caller-supplied-token" + + def test_reported_body_is_the_original_when_no_retry_happens() -> None: """A request that succeeds first time reports exactly what it sent."""