fix(a2a): forward agent_extra_headers through completion bridge (#28277)

* fix(a2a): forward agent_extra_headers through completion bridge

A2A agents backed by a custom_llm_provider (e.g. langgraph,
bedrock_agentcore) silently dropped any per-request headers rewritten
from the inbound `x-a2a-{agent}-*` convention or admin-configured
`extra_headers`. The headers were correctly extracted in
`a2a_endpoints.py` but never passed into
`_send_message_via_completion_bridge` or the bridge handler, so the
upstream HTTP request reached the agent backend without them.

Thread `agent_extra_headers` through:
  - asend_message / asend_message_streaming -> bridge call sites
  - _send_message_via_completion_bridge
  - A2ACompletionBridgeHandler.handle_non_streaming / handle_streaming
  - Inject as `extra_headers` into the underlying litellm.acompletion()
    call, and forward to provider configs via kwargs (their **kwargs
    signature absorbs it harmlessly today).

* fix(a2a): forward agent_extra_headers through bridge convenience wrappers

Address greptile review on PR #28277:

- handle_a2a_completion / handle_a2a_completion_streaming (the public,
  exported convenience wrappers) now accept agent_extra_headers and
  forward it to the underlying class methods. Without this, callers
  going through the public API would still silently drop per-request
  headers — the exact regression this PR fixes for the class-method
  path.
- Add the missing agent_extra_headers entry to the handle_streaming
  docstring for parity with handle_non_streaming.

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>

* fix(a2a/bedrock): forward agent_extra_headers to AgentCore HTTP request

Address greptile follow-up on PR #28277:

BedrockAgentCoreA2AConfig was absorbing agent_extra_headers via **kwargs
but never propagating to the underlying HTTP POST, so x-a2a-{agent}-*
rewrites and admin extra_headers were silently dropped on the
bedrock_agentcore path that bypasses the completion bridge.

Thread the parameter through the full Bedrock AgentCore stack:
- config.handle_non_streaming / handle_streaming pull
  agent_extra_headers from kwargs and pass to the handler.
- handler.handle_non_streaming / handle_streaming accept it and forward
  to the transformation layer.
- transformation.get_url_and_signed_request merges agent_extra_headers
  into the headers dict BEFORE signing, so SigV4 covers them in the
  signature. JWT/Bearer path: AgentCore signer always overwrites
  Authorization with api_key, so use api_key (not agent_extra_headers)
  to override the bearer token.

Also fix a pre-existing test assertion that was already broken by the
parent commit ab70ff6 (test_provider_config_receives_litellm_params
didn't include agent_extra_headers in the expected call).

Tests:
- TestTransformation::test_agent_extra_headers_merged_into_signed_headers_jwt
- TestTransformation::test_agent_extra_headers_signed_for_sigv4
- TestNonStreaming::test_agent_extra_headers_forwarded_on_outbound_post

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>

* fix(a2a/bedrock): drop reserved AWS headers from agent_extra_headers

Per veria-ai security review on PR #28277:

agent_extra_headers carries values rewritten from the client-controlled
x-a2a-{agent}-* convention, so the unconditional 'headers.update(agent_extra_headers)'
in BedrockAgentCoreA2ATransformation.get_url_and_signed_request let any
caller with access to an agent overwrite headers the proxy sets from
trusted server-side config -- most notably
X-Amzn-Bedrock-AgentCore-Runtime-User-Id, which AWS treats as the runtime
identity. Because the merge happened before SigV4 signing, the spoofed
value would also be bound into a valid signature.

Strip reserved AWS/AgentCore headers (authorization, host,
x-amzn-bedrock-agentcore-runtime-*, x-amz-*) from agent_extra_headers
before merging and log a warning when any are dropped. Legitimate
per-request headers (e.g. x-mcp-token, x-tenant) still pass through.

Adds two tests covering both the JWT path (verifies the spoof does not
land on the outbound headers) and the SigV4 path (verifies the signer
never sees the spoofed values).

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>

* fix(a2a): annotate completion_params dict for mypy

The dict literal initializing completion_params had heterogeneous value
types (str, list, bool), so mypy inferred the value type as a narrow
union that did not accept dict[str, str] when assigning extra_headers.

Annotate completion_params as Dict[str, Any] in both the non-streaming
and streaming bridge handlers so the agent_extra_headers merge
type-checks cleanly.

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>

* fix(a2a/pydantic_ai): forward agent_extra_headers to upstream HTTP request

* fix(a2a/bridge): admin litellm_params.extra_headers win over caller-rewritten headers

agent_extra_headers contains both admin static_headers and caller-derived
dynamic headers (from the x-a2a-{agent}-* rewrite). Merging it last would
let a caller replace headers that the proxy was configured to send upstream
via litellm_params.extra_headers. Flip the merge order so admin-configured
headers take precedence on conflict.

* fix(a2a/headers): merge_agent_headers compares case-insensitively

HTTP header names are case-insensitive, but the previous merge was a
case-sensitive dict update. That meant an admin-configured
static_headers['Authorization'] (capital A) did not strip a
caller-rewritten x-a2a-{agent}-authorization (lowercase, from the
inbound header normalization in a2a_endpoints) - both ended up on the
outbound request to pydantic_ai / langgraph / etc.

Restore the documented 'static wins on conflict' invariant by comparing
case-insensitively when overlaying static_headers. Static side's casing
is preserved on the output.

* fix(a2a/bridge): merge configured extra_headers case-insensitively over caller headers

A caller-rewritten lowercase header (e.g. authorization from the
x-a2a-{agent}-* convention) could ride alongside an admin-configured
case-variant key in litellm_params.extra_headers, sending duplicate
Authorization headers upstream. The bridge now reuses
merge_agent_headers so configured headers win case-insensitively, in
both the non-streaming and streaming paths. merge_agent_headers moved
to litellm.interactions.agents.utils (re-exported from the proxy utils)
so the SDK-level bridge does not import from litellm.proxy.

https://claude.ai/code/session_017cBvda8Y4CLo8wspB2kfSV

---------

Co-authored-by: Cursor Agent <cursoragent@cursor.com>
Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
Mateo Wang 2026-06-11 21:56:18 -07:00 • committed by GitHub
parent 2d576b5695
commit 250d8d2a96
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 692 additions and 35 deletions

View file

@ -19,6 +19,7 @@ from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
A2AStreamingContext,
)
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
from litellm.interactions.agents.utils import merge_agent_headers
# litellm_params key carrying the authenticated principal (hashed virtual key) so
# A2A provider configs can scope provider-side state (e.g. LangFlow session memory)
@ -48,6 +49,7 @@ class A2ACompletionBridgeHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> Dict[str, Any]:
@ -59,6 +61,8 @@ class A2ACompletionBridgeHandler:
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
api_base: API base URL from agent_card_params
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Returns:
A2A SendMessageResponse dict
@ -80,6 +84,7 @@ class A2ACompletionBridgeHandler:
params=params,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
)
# Extract message from params
@ -106,7 +111,7 @@ class A2ACompletionBridgeHandler:
)
# Build completion params dict
completion_params = {
completion_params: Dict[str, Any] = {
"model": full_model,
"messages": openai_messages,
"api_base": api_base,
@ -128,6 +133,12 @@ class A2ACompletionBridgeHandler:
params=params,
)
if agent_extra_headers:
completion_params["extra_headers"] = merge_agent_headers(
dynamic_headers=agent_extra_headers,
static_headers=completion_params.get("extra_headers"),
)
# Call litellm.acompletion
response = await litellm.acompletion(**completion_params)
@ -149,6 +160,7 @@ class A2ACompletionBridgeHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> AsyncIterator[Dict[str, Any]]:
@ -166,6 +178,8 @@ class A2ACompletionBridgeHandler:
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
api_base: API base URL from agent_card_params
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Yields:
A2A streaming response events
@ -187,6 +201,7 @@ class A2ACompletionBridgeHandler:
params=params,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
):
yield chunk
@ -222,7 +237,7 @@ class A2ACompletionBridgeHandler:
)
# Build completion params dict
completion_params = {
completion_params: Dict[str, Any] = {
"model": full_model,
"messages": openai_messages,
"api_base": api_base,
@ -244,6 +259,12 @@ class A2ACompletionBridgeHandler:
params=params,
)
if agent_extra_headers:
completion_params["extra_headers"] = merge_agent_headers(
dynamic_headers=agent_extra_headers,
static_headers=completion_params.get("extra_headers"),
)
# 1. Emit initial task event (kind: "task", status: "submitted")
task_event = A2ACompletionBridgeTransformation.create_task_event(ctx)
yield task_event
@ -305,6 +326,7 @@ async def handle_a2a_completion(
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""Convenience function for non-streaming A2A completion."""
return await A2ACompletionBridgeHandler.handle_non_streaming(
@ -312,6 +334,7 @@ async def handle_a2a_completion(
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
@ -320,6 +343,7 @@ async def handle_a2a_completion_streaming(
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncIterator[Dict[str, Any]]:
"""Convenience function for streaming A2A completion."""
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
@ -327,5 +351,6 @@ async def handle_a2a_completion_streaming(
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
):
yield chunk

View file

@ -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).
@ -157,6 +158,7 @@ async def _send_message_via_completion_bridge(
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
return LiteLLMSendMessageResponse.from_dict(
@ -283,6 +285,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 +512,7 @@ async def asend_message_streaming( # noqa: PLR0915
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
):
yield chunk
return

View file

@ -37,6 +37,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
litellm_params=litellm_params,
agent_extra_headers=kwargs.get("agent_extra_headers"),
)
async def handle_streaming(
@ -57,5 +58,6 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
litellm_params=litellm_params,
agent_extra_headers=kwargs.get("agent_extra_headers"),
):
yield chunk

View file

@ -6,7 +6,7 @@ completion bridge that would otherwise strip the envelope.
"""
import json
from typing import Any, AsyncIterator, Dict, cast
from typing import Any, AsyncIterator, Dict, Optional, cast
from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
@ -29,6 +29,7 @@ class BedrockAgentCoreA2AHandler:
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Handle non-streaming A2A request to AgentCore.
@ -37,6 +38,8 @@ class BedrockAgentCoreA2AHandler:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (model, api_key, etc.)
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Returns:
A2A JSON-RPC response dict from the AgentCore agent
@ -47,6 +50,7 @@ class BedrockAgentCoreA2AHandler:
params=params,
litellm_params=litellm_params,
method="message/send",
agent_extra_headers=agent_extra_headers,
)
)
@ -77,6 +81,7 @@ class BedrockAgentCoreA2AHandler:
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncIterator[Dict[str, Any]]:
"""
Handle streaming A2A request to AgentCore.
@ -85,6 +90,8 @@ class BedrockAgentCoreA2AHandler:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (model, api_key, etc.)
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Yields:
A2A streaming response events from the AgentCore agent
@ -96,6 +103,7 @@ class BedrockAgentCoreA2AHandler:
litellm_params=litellm_params,
method="message/send",
stream=True,
agent_extra_headers=agent_extra_headers,
)
)

View file

@ -6,11 +6,66 @@ and signs requests via AmazonAgentCoreConfig (SigV4 or JWT).
"""
import json
from typing import Any, AsyncIterator, Dict, Tuple
from typing import Any, AsyncIterator, Dict, Mapping, Optional, Tuple
from litellm._logging import verbose_logger
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
# Reserved outbound header names that must never be sourced from per-request
# ``agent_extra_headers`` for AgentCore requests. ``agent_extra_headers`` carries
# values rewritten from the client-controlled ``x-a2a-{agent}-*`` convention, so
# allowing these would let any caller with access to the agent spoof the AWS
# request identity / SigV4 metadata by overwriting headers the proxy sets from
# trusted server-side config.
#
# The runtime headers (session / user id) are derived server-side from
# ``runtimeSessionId`` / ``runtimeUserId`` in the agent's ``litellm_params``;
# ``authorization`` is set by the AgentCore signer (JWT or SigV4); ``host`` and
# the ``x-amz-*`` family are owned by SigV4 itself.
_RESERVED_EXACT_HEADERS = frozenset(
{
"authorization",
"host",
}
)
_RESERVED_PREFIX_HEADERS: Tuple[str, ...] = (
"x-amzn-bedrock-agentcore-runtime-",
"x-amz-",
)
def _filter_reserved_headers(
agent_extra_headers: Optional[Mapping[str, str]],
) -> Optional[Dict[str, str]]:
"""
Strip reserved AWS / AgentCore headers from caller-supplied
``agent_extra_headers`` before they are merged into the signed request.
Returns ``None`` if the result is empty.
"""
if not agent_extra_headers:
return None
filtered: Dict[str, str] = {}
dropped: list = []
for k, v in agent_extra_headers.items():
k_lower = k.lower()
if k_lower in _RESERVED_EXACT_HEADERS or any(
k_lower.startswith(prefix) for prefix in _RESERVED_PREFIX_HEADERS
):
dropped.append(k)
continue
filtered[k] = v
if dropped:
verbose_logger.warning(
"BedrockAgentCore A2A: dropping reserved header(s) from "
"agent_extra_headers (not forwarded to AgentCore): %s",
sorted(dropped),
)
return filtered or None
class BedrockAgentCoreA2ATransformation:
"""
@ -27,6 +82,7 @@ class BedrockAgentCoreA2ATransformation:
litellm_params: Dict[str, Any],
method: str = "message/send",
stream: bool = False,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[str, dict, bytes]:
"""
Build the AgentCore URL, construct a JSON-RPC envelope, and sign the request.
@ -37,6 +93,15 @@ class BedrockAgentCoreA2ATransformation:
litellm_params: Agent's litellm_params (model, api_key, etc.)
method: JSON-RPC method name (default: "message/send")
stream: Whether this is a streaming request
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call. Merged into
the headers dict before signing so SigV4 includes them in the signature.
Reserved AWS / AgentCore identity headers (``authorization``, ``host``,
``x-amzn-bedrock-agentcore-runtime-*``, ``x-amz-*``) are filtered out
here to prevent a caller-controlled ``x-a2a-{agent}-*`` header from
spoofing the AgentCore runtime user id or other SigV4 metadata. Use
``api_key`` / ``runtimeUserId`` / ``runtimeSessionId`` in litellm_params
(not ``agent_extra_headers``) to override those values.
Returns:
Tuple of (url, signed_headers, signed_body_bytes)
@ -85,6 +150,13 @@ class BedrockAgentCoreA2ATransformation:
if runtime_user_id:
headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id
# Merge per-request agent headers before signing so SigV4 covers them.
# Reserved headers are stripped first to prevent client-controlled values
# from spoofing the AgentCore runtime identity / SigV4 metadata.
safe_extra_headers = _filter_reserved_headers(agent_extra_headers)
if safe_extra_headers:
headers.update(safe_extra_headers)
# Sign the request (SigV4 or JWT depending on api_key presence)
signed_headers, signed_body = agentcore_config.sign_request(
headers=headers,

View file

@ -31,6 +31,7 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
params=params,
api_base=api_base,
timeout=kwargs.get("timeout", 60.0),
agent_extra_headers=kwargs.get("agent_extra_headers"),
)
async def handle_streaming(
@ -50,5 +51,6 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
timeout=kwargs.get("timeout", 60.0),
chunk_size=kwargs.get("chunk_size", 50),
delay_ms=kwargs.get("delay_ms", 10),
agent_extra_headers=kwargs.get("agent_extra_headers"),
):
yield chunk

View file

@ -28,6 +28,7 @@ class PydanticAIHandler:
params: Dict[str, Any],
api_base: Optional[str] = None,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Handle non-streaming request to Pydantic AI agent.
@ -37,6 +38,8 @@ class PydanticAIHandler:
params: A2A MessageSendParams containing the message
api_base: Base URL of the Pydantic AI agent
timeout: Request timeout in seconds
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Returns:
A2A SendMessageResponse dict
@ -51,6 +54,7 @@ class PydanticAIHandler:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
return response_data
@ -63,6 +67,7 @@ class PydanticAIHandler:
timeout: float = 60.0,
chunk_size: int = 50,
delay_ms: int = 10,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncIterator[Dict[str, Any]]:
"""
Handle streaming request to Pydantic AI agent with fake streaming.
@ -78,6 +83,8 @@ class PydanticAIHandler:
timeout: Request timeout in seconds
chunk_size: Number of characters per chunk
delay_ms: Delay between chunks in milliseconds
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Yields:
A2A streaming response events
@ -94,6 +101,7 @@ class PydanticAIHandler:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
# Convert raw task response to fake streaming chunks

View file

@ -6,7 +6,7 @@ This module provides fake streaming by converting non-streaming responses into s
"""
import asyncio
from typing import Any, AsyncIterator, Dict, cast
from typing import Any, AsyncIterator, Dict, Optional, cast
from uuid import uuid4
from litellm._logging import verbose_logger
@ -86,6 +86,7 @@ class PydanticAITransformation:
request_id: str,
max_attempts: int = 30,
poll_interval: float = 0.5,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Poll for task completion using tasks/get method.
@ -112,7 +113,10 @@ class PydanticAITransformation:
response = await client.post(
endpoint,
json=poll_request,
headers={"Content-Type": "application/json"},
headers={
**(agent_extra_headers or {}),
"Content-Type": "application/json",
},
)
response.raise_for_status()
poll_data = response.json()
@ -142,6 +146,7 @@ class PydanticAITransformation:
request_id: str,
params: Any,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Send a request to Pydantic AI agent and return the raw task response.
@ -189,7 +194,10 @@ class PydanticAITransformation:
response = await client.post(
endpoint,
json=a2a_request,
headers={"Content-Type": "application/json"},
headers={
**(agent_extra_headers or {}),
"Content-Type": "application/json",
},
)
response.raise_for_status()
response_data = response.json()
@ -211,6 +219,7 @@ class PydanticAITransformation:
endpoint=endpoint,
task_id=task_id,
request_id=request_id,
agent_extra_headers=agent_extra_headers,
)
verbose_logger.info(
@ -225,6 +234,7 @@ class PydanticAITransformation:
request_id: str,
params: Any,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Send a non-streaming A2A request to Pydantic AI agent and wait for completion.
@ -234,6 +244,7 @@ class PydanticAITransformation:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message (dict or Pydantic model)
timeout: Request timeout in seconds
agent_extra_headers: Per-request headers to forward on the upstream HTTP call.
Returns:
Standard A2A non-streaming response format with message
@ -244,6 +255,7 @@ class PydanticAITransformation:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
# Transform to standard A2A non-streaming format
@ -258,6 +270,7 @@ class PydanticAITransformation:
request_id: str,
params: Any,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Send a request to Pydantic AI agent and return the raw task response.
@ -269,6 +282,7 @@ class PydanticAITransformation:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
timeout: Request timeout in seconds
agent_extra_headers: Per-request headers to forward on the upstream HTTP call.
Returns:
Raw Pydantic AI task response (with history/artifacts)
@ -278,6 +292,7 @@ class PydanticAITransformation:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
@staticmethod

View file

@ -2,11 +2,40 @@
Utility functions for the Agents API SDK.
"""
from typing import Optional
from typing import Dict, Mapping, Optional
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
def merge_agent_headers(
*,
dynamic_headers: Optional[Mapping[str, str]] = None,
static_headers: Optional[Mapping[str, str]] = None,
) -> Optional[Dict[str, str]]:
"""Merge outbound HTTP headers for A2A agent calls.
Merge rules:
- Start with ``dynamic_headers`` (values extracted from the incoming client request).
- Overlay ``static_headers`` (admin-configured per agent).
- Comparison is case-insensitive (HTTP headers are case-insensitive), so a
static ``Authorization`` strips any dynamic ``authorization`` before the
static value is written. The static side's casing is preserved.
If both contain the same header (case-insensitively), ``static_headers`` wins.
"""
merged: Dict[str, str] = {}
if dynamic_headers:
merged.update({str(k): str(v) for k, v in dynamic_headers.items()})
if static_headers:
static_lower = {str(k).lower() for k in static_headers}
merged = {k: v for k, v in merged.items() if k.lower() not in static_lower}
merged.update({str(k): str(v) for k, v in static_headers.items()})
return merged or None
def get_provider_agents_api_config(
custom_llm_provider: Optional[str],
) -> Optional[BaseAgentsAPIConfig]:

View file

@ -1,33 +1,8 @@
"""Utility helpers for A2A agent endpoints."""
from typing import Dict, Mapping, Optional
# Re-export from the canonical SDK location so the proxy and SDK always
# share the same provider-config lookup logic.
# share the same provider-config lookup and header-merge logic.
from litellm.interactions.agents.utils import ( # noqa: F401
get_provider_agents_api_config,
merge_agent_headers,
)
def merge_agent_headers(
*,
dynamic_headers: Optional[Mapping[str, str]] = None,
static_headers: Optional[Mapping[str, str]] = None,
) -> Optional[Dict[str, str]]:
"""Merge outbound HTTP headers for A2A agent calls.
Merge rules:
- Start with ``dynamic_headers`` (values extracted from the incoming client request).
- Overlay ``static_headers`` (admin-configured per agent).
If both contain the same key, ``static_headers`` wins.
"""
merged: Dict[str, str] = {}
if dynamic_headers:
merged.update({str(k): str(v) for k, v in dynamic_headers.items()})
if static_headers:
merged.update({str(k): str(v) for k, v in static_headers.items()})
return merged or None

View file

@ -0,0 +1,200 @@
"""
Tests for Pydantic AI agents header forwarding via agent_extra_headers.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import (
PydanticAITransformation,
)
def _build_mock_client(response_payload):
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json = MagicMock(return_value=response_payload)
mock_client = MagicMock()
mock_client.post = AsyncMock(return_value=mock_response)
return mock_client
@pytest.mark.asyncio
async def test_send_non_streaming_request_forwards_agent_extra_headers():
"""agent_extra_headers should be merged into the outbound HTTP request headers."""
completed_payload = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {
"id": "task-1",
"kind": "task",
"status": {"state": "completed"},
"history": [
{
"role": "agent",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
],
"artifacts": [],
},
}
mock_client = _build_mock_client(completed_payload)
with patch(
"litellm.a2a_protocol.providers.pydantic_ai_agents.transformation.get_async_httpx_client",
return_value=mock_client,
):
await PydanticAITransformation.send_non_streaming_request(
api_base="http://example.test",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hello"}],
"messageId": "msg-user-1",
}
},
agent_extra_headers={
"x-tenant-id": "acme",
"authorization": "Bearer caller-supplied",
},
)
assert mock_client.post.await_count == 1
sent_headers = mock_client.post.await_args.kwargs["headers"]
assert sent_headers["x-tenant-id"] == "acme"
assert sent_headers["authorization"] == "Bearer caller-supplied"
assert sent_headers["Content-Type"] == "application/json"
@pytest.mark.asyncio
async def test_send_non_streaming_request_without_headers_preserves_content_type():
"""When no agent_extra_headers are passed, behavior is unchanged."""
completed_payload = {
"jsonrpc": "2.0",
"id": "req-2",
"result": {
"id": "task-2",
"kind": "task",
"status": {"state": "completed"},
"history": [],
"artifacts": [
{
"artifactId": "a-1",
"parts": [{"kind": "text", "text": "ok"}],
}
],
},
}
mock_client = _build_mock_client(completed_payload)
with patch(
"litellm.a2a_protocol.providers.pydantic_ai_agents.transformation.get_async_httpx_client",
return_value=mock_client,
):
await PydanticAITransformation.send_non_streaming_request(
api_base="http://example.test",
request_id="req-2",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hello"}],
"messageId": "msg-user-2",
}
},
)
sent_headers = mock_client.post.await_args.kwargs["headers"]
assert sent_headers == {"Content-Type": "application/json"}
@pytest.mark.asyncio
async def test_content_type_is_preserved_when_caller_tries_to_override():
"""A caller-supplied Content-Type must not displace application/json."""
completed_payload = {
"jsonrpc": "2.0",
"id": "req-3",
"result": {
"id": "task-3",
"kind": "task",
"status": {"state": "completed"},
"history": [],
"artifacts": [
{
"artifactId": "a-2",
"parts": [{"kind": "text", "text": "ok"}],
}
],
},
}
mock_client = _build_mock_client(completed_payload)
with patch(
"litellm.a2a_protocol.providers.pydantic_ai_agents.transformation.get_async_httpx_client",
return_value=mock_client,
):
await PydanticAITransformation.send_non_streaming_request(
api_base="http://example.test",
request_id="req-3",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hello"}],
"messageId": "msg-user-3",
}
},
agent_extra_headers={"Content-Type": "text/plain"},
)
sent_headers = mock_client.post.await_args.kwargs["headers"]
assert sent_headers["Content-Type"] == "application/json"
@pytest.mark.asyncio
async def test_provider_config_threads_agent_extra_headers():
"""End-to-end: PydanticAIProviderConfig forwards agent_extra_headers down the stack."""
from litellm.a2a_protocol.providers.pydantic_ai_agents.config import (
PydanticAIProviderConfig,
)
completed_payload = {
"jsonrpc": "2.0",
"id": "req-4",
"result": {
"id": "task-4",
"kind": "task",
"status": {"state": "completed"},
"history": [],
"artifacts": [
{
"artifactId": "a-3",
"parts": [{"kind": "text", "text": "ok"}],
}
],
},
}
mock_client = _build_mock_client(completed_payload)
with patch(
"litellm.a2a_protocol.providers.pydantic_ai_agents.transformation.get_async_httpx_client",
return_value=mock_client,
):
await PydanticAIProviderConfig().handle_non_streaming(
request_id="req-4",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hello"}],
"messageId": "msg-user-4",
}
},
api_base="http://example.test",
agent_extra_headers={"x-trace-id": "abc-123"},
)
sent_headers = mock_client.post.await_args.kwargs["headers"]
assert sent_headers["x-trace-id"] == "abc-123"
assert sent_headers["Content-Type"] == "application/json"

View file

@ -110,6 +110,153 @@ class TestTransformation:
)
assert headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] == "a" * 40
def test_agent_extra_headers_merged_into_signed_headers_jwt(self):
"""agent_extra_headers should appear on the outbound request (JWT path)."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
agent_extra_headers={"x-mcp-token": "mcp-abc", "x-tenant": "t1"},
)
assert headers["x-mcp-token"] == "mcp-abc"
assert headers["x-tenant"] == "t1"
def test_agent_extra_headers_signed_for_sigv4(self):
"""agent_extra_headers must be present in the dict passed to _sign_request."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
litellm_params_no_key = {
"model": SAMPLE_MODEL,
"custom_llm_provider": "bedrock",
"aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
"aws_region_name": "us-west-2",
}
captured: dict = {}
def fake_sign(self, headers, **kwargs):
captured.update(headers)
return headers, b'{"jsonrpc":"2.0"}'
with patch(
"litellm.llms.bedrock.chat.agentcore.transformation.AmazonAgentCoreConfig._sign_request",
new=fake_sign,
):
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=litellm_params_no_key,
agent_extra_headers={"x-mcp-token": "mcp-abc"},
)
assert captured.get("x-mcp-token") == "mcp-abc"
def test_reserved_headers_filtered_from_agent_extra_headers(self):
"""
Reserved AWS / AgentCore headers in agent_extra_headers must NOT overwrite
the values the proxy sets from trusted server-side config, otherwise a
caller could spoof the runtime user identity via the x-a2a-{agent}-*
header rewrite.
"""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
litellm_params_with_user = {
**SAMPLE_LITELLM_PARAMS,
"runtimeUserId": "legit-user",
}
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=litellm_params_with_user,
agent_extra_headers={
# Spoofing attempt — must be dropped.
"x-amzn-bedrock-agentcore-runtime-user-id": "victim-user",
"X-Amzn-Bedrock-AgentCore-Runtime-Session-Id": "spoofed-session",
"Authorization": "Bearer attacker-token",
"Host": "attacker.example.com",
"x-amz-content-sha256": "deadbeef",
# Legitimate per-request header — must pass through.
"x-mcp-token": "mcp-abc",
},
)
# Legitimate header is preserved.
assert headers["x-mcp-token"] == "mcp-abc"
# Reserved headers from agent_extra_headers must not appear at all
# (case-insensitive) — only the proxy/signer-controlled values may.
normalized = {k.lower(): v for k, v in headers.items()}
# Runtime user id is the value set from litellm_params, NOT the spoof.
assert normalized["x-amzn-bedrock-agentcore-runtime-user-id"] == "legit-user"
# Session id is the auto-generated one, not the spoofed value.
assert (
normalized["x-amzn-bedrock-agentcore-runtime-session-id"]
!= "spoofed-session"
)
# Authorization is the JWT bearer set by the signer, not the spoof.
assert normalized["authorization"] == "Bearer test-jwt-token"
# Host / x-amz-* must not have been carried over from the client.
assert normalized.get("host") != "attacker.example.com"
assert normalized.get("x-amz-content-sha256") != "deadbeef"
def test_reserved_headers_filtered_before_sigv4_signing(self):
"""
Reserved headers in agent_extra_headers must be stripped BEFORE the
SigV4 signer sees them, so the signature does not bind a spoofed
runtime user identity into a valid SigV4 request.
"""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
litellm_params_no_key = {
"model": SAMPLE_MODEL,
"custom_llm_provider": "bedrock",
"aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
"aws_region_name": "us-west-2",
"runtimeUserId": "legit-user",
}
captured: dict = {}
def fake_sign(self, headers, **kwargs):
captured.update(headers)
return headers, b'{"jsonrpc":"2.0"}'
with patch(
"litellm.llms.bedrock.chat.agentcore.transformation.AmazonAgentCoreConfig._sign_request",
new=fake_sign,
):
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=litellm_params_no_key,
agent_extra_headers={
"x-amzn-bedrock-agentcore-runtime-user-id": "victim-user",
"x-amz-date": "20990101T000000Z",
"authorization": "Bearer attacker",
"x-mcp-token": "mcp-abc",
},
)
normalized = {k.lower(): v for k, v in captured.items()}
assert normalized["x-amzn-bedrock-agentcore-runtime-user-id"] == "legit-user"
assert normalized.get("x-amz-date") != "20990101T000000Z"
assert normalized.get("authorization") != "Bearer attacker"
# Non-reserved header still makes it into the signed dict.
assert captured.get("x-mcp-token") == "mcp-abc"
def test_sigv4_auth_when_no_api_key(self):
"""When no api_key, falls through to SigV4 signing."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
@ -200,6 +347,39 @@ class TestNonStreaming:
# Verify response is passed through
assert result["result"]["message"]["parts"][0]["text"] == "2"
@pytest.mark.asyncio
async def test_agent_extra_headers_forwarded_on_outbound_post(self):
"""End-to-end: agent_extra_headers from the bridge land on the HTTP POST."""
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
mock_response = MagicMock()
mock_response.json.return_value = {
"jsonrpc": "2.0",
"id": "req-001",
"result": {},
}
mock_response.raise_for_status = MagicMock()
with patch(
"litellm.a2a_protocol.providers.bedrock_agentcore.handler.get_async_httpx_client"
) as mock_get_client:
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
config = BedrockAgentCoreA2AConfig()
await config.handle_non_streaming(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
agent_extra_headers={"x-mcp-token": "mcp-abc"},
)
sent_headers = mock_client.post.call_args.kwargs["headers"]
assert sent_headers.get("x-mcp-token") == "mcp-abc"
@pytest.mark.asyncio
async def test_a2a_error_response_passthrough(self):
"""JSON-RPC error responses from the agent are returned as-is."""
@ -301,6 +481,7 @@ class TestHandlerIntegration:
params=SAMPLE_PARAMS,
api_base=None,
litellm_params=SAMPLE_LITELLM_PARAMS,
agent_extra_headers=None,
)
@pytest.mark.asyncio

View file

@ -438,3 +438,139 @@ def test_merge_agent_headers_util_empty_dicts_returns_none():
result = merge_agent_headers(dynamic_headers={}, static_headers={})
assert result is None
def test_merge_agent_headers_util_case_insensitive_static_wins():
"""Static ``Authorization`` strips dynamic ``authorization`` (HTTP headers are case-insensitive)."""
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
result = merge_agent_headers(
dynamic_headers={"authorization": "Bearer caller-token", "x-extra": "d"},
static_headers={"Authorization": "Bearer admin-token"},
)
assert result == {"Authorization": "Bearer admin-token", "x-extra": "d"}
def test_merge_agent_headers_util_case_insensitive_no_dynamic_leak():
"""No case-variant of a static header can leak through from dynamic headers."""
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
result = merge_agent_headers(
dynamic_headers={"AUTHORIZATION": "Bearer caller", "authorization": "x"},
static_headers={"Authorization": "Bearer admin"},
)
assert result == {"Authorization": "Bearer admin"}
@pytest.mark.asyncio
async def test_convention_header_blocked_by_case_variant_static():
"""Static ``Authorization`` blocks caller-rewritten lowercase ``authorization``."""
mock_agent = _make_mock_agent(
static_headers={"Authorization": "Bearer admin-token"}
)
mock_agent.agent_name = "my-agent"
mock_request = _make_mock_request(
extra_headers={"x-a2a-my-agent-authorization": "Bearer caller-token"}
)
mock_asend = await _invoke(mock_agent, mock_request, None)
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
assert headers is not None
assert headers == {"Authorization": "Bearer admin-token"}
assert "authorization" not in headers
# ---------------------------------------------------------------------------
# Completion bridge: configured litellm_params.extra_headers win
# case-insensitively over caller-rewritten headers
# ---------------------------------------------------------------------------
_BRIDGE_MESSAGE_PARAMS = {
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "Hi"}],
"messageId": "msg-123",
}
}
@pytest.mark.asyncio
async def test_bridge_caller_header_cannot_shadow_configured_header():
"""A caller-rewritten lowercase ``authorization`` must not ride alongside the
admin-configured ``Authorization`` from ``litellm_params.extra_headers``."""
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
mock_response.choices[0].message = MagicMock()
mock_response.choices[0].message.content = "Hello!"
mock_response.id = "resp-123"
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
mock_acompletion.return_value = mock_response
await A2ACompletionBridgeHandler.handle_non_streaming(
request_id="req-456",
params=_BRIDGE_MESSAGE_PARAMS,
litellm_params={
"custom_llm_provider": "langgraph",
"model": "agent",
"extra_headers": {"Authorization": "Bearer admin-token"},
},
api_base="http://backend-agent:10001",
agent_extra_headers={
"authorization": "Bearer caller-token",
"x-mcp-token": "mcp-abc",
},
)
sent_headers = mock_acompletion.call_args.kwargs["extra_headers"]
assert sent_headers == {
"Authorization": "Bearer admin-token",
"x-mcp-token": "mcp-abc",
}
@pytest.mark.asyncio
async def test_bridge_streaming_caller_header_cannot_shadow_configured_header():
"""Streaming path applies the same case-insensitive precedence."""
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
mock_chunk = MagicMock()
mock_chunk.choices = [MagicMock()]
mock_chunk.choices[0].delta = MagicMock()
mock_chunk.choices[0].delta.content = "Hello"
async def mock_streaming_response():
yield mock_chunk
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
mock_acompletion.return_value = mock_streaming_response()
async for _ in A2ACompletionBridgeHandler.handle_streaming(
request_id="req-456",
params=_BRIDGE_MESSAGE_PARAMS,
litellm_params={
"custom_llm_provider": "langgraph",
"model": "agent",
"extra_headers": {"Authorization": "Bearer admin-token"},
},
api_base="http://backend-agent:10001",
agent_extra_headers={
"authorization": "Bearer caller-token",
"x-mcp-token": "mcp-abc",
},
):
pass
sent_headers = mock_acompletion.call_args.kwargs["extra_headers"]
assert sent_headers == {
"Authorization": "Bearer admin-token",
"x-mcp-token": "mcp-abc",
}