diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 3ad5485dea1..269ce2d027b 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -132,6 +132,7 @@ async def _send_message_via_completion_bridge( custom_llm_provider: str, api_base: Optional[str], litellm_params: Dict[str, Any], + agent_extra_headers: Optional[Dict[str, str]] = None, ) -> LiteLLMSendMessageResponse: """ Route a send_message through the LiteLLM completion bridge (e.g. LangGraph, Bedrock AgentCore). @@ -152,6 +153,13 @@ async def _send_message_via_completion_bridge( else dict(request.params) ) + if agent_extra_headers: + merged = dict(litellm_params) + existing = dict(merged.get("extra_headers") or {}) + existing.update(agent_extra_headers) + merged["extra_headers"] = existing + litellm_params = merged + response_dict = await A2ACompletionBridgeHandler.handle_non_streaming( request_id=str(request.id), params=params, @@ -281,6 +289,7 @@ async def asend_message( custom_llm_provider=custom_llm_provider, api_base=api_base, litellm_params=litellm_params, + agent_extra_headers=agent_extra_headers, ) # Standard A2A client flow @@ -509,6 +518,13 @@ async def asend_message_streaming( # noqa: PLR0915 else dict(request.params) ) + if agent_extra_headers: + merged = dict(litellm_params) + existing = dict(merged.get("extra_headers") or {}) + existing.update(agent_extra_headers) + merged["extra_headers"] = existing + litellm_params = merged + async for chunk in A2ACompletionBridgeHandler.handle_streaming( request_id=str(request.id), params=params,