fix: forward A2A headers and errors

This commit is contained in:
aiedwardyi 2026-08-24 16:13:57 +09:00
parent 02c5c74b37
commit 41bec1c2d0
No known key found for this signature in database
4 changed files with 80 additions and 3 deletions

View file

@ -7,7 +7,7 @@ from typing import Final
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
from ..common_utils import extract_text_from_a2a_response
from ..common_utils import A2AError, extract_text_from_a2a_response
class A2AModelResponseIterator(BaseModelResponseIterator):
@ -56,6 +56,15 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
}
}
"""
if "error" in chunk:
error_value: Final = chunk["error"]
error_message: Final = (
error_value.get("message")
if isinstance(error_value, dict) and isinstance(error_value.get("message"), str)
else str(error_value)
)
raise A2AError(status_code=500, message=f"A2A error: {error_message}")
try:
# Extract text from A2A response
text: Final = extract_text_from_a2a_response(chunk)

View file

@ -89,6 +89,7 @@ async def _route_registered_provider(
api_base: str | None,
litellm_params: Mapping[str, object],
static_headers: Mapping[str, str] | None,
dynamic_headers: Mapping[str, str] | None = None,
) -> ModelResponse | CustomStreamWrapper:
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
@ -122,7 +123,10 @@ async def _route_registered_provider(
_HEADERS_ADAPTER.validate_python(configured_headers) if isinstance(configured_headers, dict) else None
)
agent_extra_headers: Final = merge_agent_headers(
dynamic_headers=configured_headers_dict,
dynamic_headers=merge_agent_headers(
dynamic_headers=dynamic_headers,
static_headers=configured_headers_dict,
),
static_headers=static_headers,
)
if agent_extra_headers:
@ -233,6 +237,40 @@ def _merge_agent_guardrails(
return merged_data
def _get_agent_dynamic_headers(
data: Mapping[str, object],
agent_id: str,
agent_name: str,
extra_headers: list[str] | None,
) -> dict[str, str]:
proxy_request: Final = data.get("proxy_server_request")
raw_headers: object = proxy_request.get("headers") if isinstance(proxy_request, Mapping) else None
if not isinstance(raw_headers, Mapping):
metadata: Final = data.get("metadata")
raw_headers = metadata.get("headers") if isinstance(metadata, Mapping) else None
normalized_headers: Final = (
{str(key).lower(): str(value) for key, value in raw_headers.items()}
if isinstance(raw_headers, Mapping)
else {}
)
dynamic_headers: dict[str, str] = {}
for header_name in extra_headers or []:
header_name_str: Final = str(header_name)
value: Final = normalized_headers.get(header_name_str.lower())
if value is not None:
dynamic_headers[header_name_str] = value
for alias in (agent_id.lower(), agent_name.lower()):
prefix: Final = f"x-a2a-{alias}-"
for key, value in normalized_headers.items():
if key.startswith(prefix):
header_name: Final = key[len(prefix) :]
if header_name:
dynamic_headers[header_name] = value
return dynamic_headers
async def route_a2a_agent_request(
data: Mapping[str, object],
route_type: str,
@ -310,6 +348,12 @@ async def route_a2a_agent_request(
data=data,
agent_guardrails=registered_params_value.get("guardrails") if registered_params_value else None,
)
registered_dynamic_headers: Final = _get_agent_dynamic_headers(
data=routed_data,
agent_id=agent.agent_id,
agent_name=agent.agent_name,
extra_headers=agent.extra_headers,
)
if (
registered_provider
and registered_provider != "a2a"
@ -323,6 +367,7 @@ async def route_a2a_agent_request(
api_base=api_base,
litellm_params=registered_params_value,
static_headers=agent.static_headers,
dynamic_headers=registered_dynamic_headers,
)
completion_data: Final = MappingProxyType({**routed_data, "api_base": api_base})

View file

@ -3,6 +3,7 @@
import pytest
from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator
from litellm.llms.a2a.common_utils import A2AError
@pytest.mark.asyncio
@ -24,3 +25,14 @@ async def test_async_iterator_accepts_decoded_a2a_events():
chunk = await iterator.__aiter__().__anext__()
assert chunk["text"] == "Hello"
@pytest.mark.asyncio
async def test_async_iterator_propagates_jsonrpc_errors():
async def _events():
yield {"jsonrpc": "2.0", "error": {"code": -32000, "message": "agent failed"}}
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
with pytest.raises(A2AError, match="agent failed"):
await iterator.__aiter__().__anext__()

View file

@ -83,6 +83,7 @@ async def test_route_a2a_model_uses_registered_provider():
"guardrails": ["agent-guardrail"],
},
static_headers={"Authorization": "Bearer static"},
extra_headers=["X-Tenant"],
)
data = {
"model": "a2a/test-agent",
@ -92,6 +93,12 @@ async def test_route_a2a_model_uses_registered_provider():
"temperature": 0.2,
"timeout": 12.0,
"tools": [{"type": "function", "function": {"name": "lookup"}}],
"proxy_server_request": {
"headers": {
"x-tenant": "tenant-1",
"x-a2a-test-agent-x-run": "run-1",
}
},
}
bridge_response = {
"jsonrpc": "2.0",
@ -133,7 +140,11 @@ async def test_route_a2a_model_uses_registered_provider():
assert bridge_kwargs["litellm_params"]["timeout"] == 12.0
assert bridge_kwargs["litellm_params"]["tools"] == data["tools"]
assert bridge_kwargs["litellm_params"]["guardrails"] == ["request-guardrail", "agent-guardrail"]
assert bridge_kwargs["litellm_params"]["extra_headers"] == {"Authorization": "Bearer static"}
assert bridge_kwargs["litellm_params"]["extra_headers"] == {
"X-Tenant": "tenant-1",
"x-run": "run-1",
"Authorization": "Bearer static",
}
@pytest.mark.asyncio