diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 0cbff95b9d3..1c3fa0c6c92 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -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) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 2963240f41a..bd8b222ce4f 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -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}) diff --git a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py index 579d6f8efef..b301b2f3c6e 100644 --- a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py @@ -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__() diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 1dd2edc8bf7..d96ee37dbdf 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -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