From 633813e6eecc878f523c563a05c4e38dcbb22b0d Mon Sep 17 00:00:00 2001 From: shivam Date: Sat, 3 Oct 2026 23:08:29 +0000 Subject: [PATCH 1/2] fix(a2a): add per-agent a2a_protocol_version override for langgraph-api 0.15 cards Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/a2a_protocol/card_resolver.py | 42 +++++++-- .../a2a_protocol/exception_mapping_utils.py | 8 +- .../litellm_completion_bridge/handler.py | 3 +- litellm/a2a_protocol/main.py | 13 ++- .../test_a2a_exception_mapping_utils.py | 16 +++- tests/unit/a2a_protocol/test_card_resolver.py | 72 +++++++++++++++ .../test_completion_bridge_streaming.py | 2 + tests/unit/a2a_protocol/test_main.py | 87 +++++++++++++++++++ 8 files changed, 231 insertions(+), 12 deletions(-) diff --git a/litellm/a2a_protocol/card_resolver.py b/litellm/a2a_protocol/card_resolver.py index 8614c794ac4..08e95f99c87 100644 --- a/litellm/a2a_protocol/card_resolver.py +++ b/litellm/a2a_protocol/card_resolver.py @@ -21,6 +21,7 @@ AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent-card.json" PREV_AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent.json" FOUNDRY_AGENT_CARD_PATH: Final = "/agentCard/v1.0" AGENT_CARD_PATH_PARAM: Final = "agent_card_path" +A2A_PROTOCOL_VERSION_PARAM: Final = "a2a_protocol_version" try: from a2a.client import A2ACardResolver as _A2ACardResolver @@ -76,22 +77,43 @@ _CANONICAL_PROTOCOL_BINDINGS: Final = MappingProxyType( ) _LEGACY_PROTOCOL_VERSION: Final = "0.3" +_A2A_PROTOCOL_VERSION_ALIASES: Final = MappingProxyType( + { + "0.3": "0.3", + "0.3.0": "0.3", + "1.0": "1.0", + "1.0.0": "1.0", + } +) -def normalize_agent_card_interfaces(agent_card: "AgentCard") -> "AgentCard": +def resolve_a2a_protocol_version(value: object) -> str | None: + if value is None: + return None + if isinstance(value, bool): + verbose_logger.warning("Ignoring invalid %s value: %r", A2A_PROTOCOL_VERSION_PARAM, value) + return None + if not isinstance(value, (str, int, float)): + verbose_logger.warning("Ignoring invalid %s value: %r", A2A_PROTOCOL_VERSION_PARAM, value) + return None + + protocol_version: Final = _A2A_PROTOCOL_VERSION_ALIASES.get(str(value).strip()) + if protocol_version is None: + verbose_logger.warning("Ignoring invalid %s value: %r", A2A_PROTOCOL_VERSION_PARAM, value) + return protocol_version + + +def normalize_agent_card_interfaces(agent_card: "AgentCard", protocol_version: str | None = None) -> "AgentCard": """ Canonicalize the supported interfaces of spec-adjacent agent cards. Some A2A servers (e.g. LangGraph Platform) serve agent cards with lowercase bindings like "jsonrpc", but a2a-sdk's ClientFactory matches bindings case-sensitively against its uppercase TransportProtocol constants and fails - with "no compatible transports found." for spec-adjacent casings. - - The same servers also speak the A2A 0.3 JSON dialect ("kind"-discriminated - payloads) while declaring protocolVersion "1.0", which a2a-sdk's strict v1 - proto parsing rejects. A mis-cased binding fingerprints such a server, so its - declared version is downgraded to 0.3 to route the SDK's ClientFactory onto - its v0.3 compat transport, which speaks that dialect. + with "no compatible transports found." for spec-adjacent casings. The casing + fingerprint stopped working with langgraph-api 0.15, so operators can pin + the version per agent via ``litellm_params.a2a_protocol_version``. An explicit + "1.0" opts out of the mis-cased downgrade. """ normalized: Final = type(agent_card)() normalized.CopyFrom(agent_card) @@ -101,6 +123,10 @@ def normalize_agent_card_interfaces(agent_card: "AgentCard") -> "AgentCard": continue interface.protocol_binding = canonical interface.protocol_version = _LEGACY_PROTOCOL_VERSION + if protocol_version is not None: + for interface in normalized.supported_interfaces: + if interface.protocol_binding in _CANONICAL_PROTOCOL_BINDINGS.values(): + interface.protocol_version = protocol_version return normalized diff --git a/litellm/a2a_protocol/exception_mapping_utils.py b/litellm/a2a_protocol/exception_mapping_utils.py index 62bef6e02ae..1abf32915f5 100644 --- a/litellm/a2a_protocol/exception_mapping_utils.py +++ b/litellm/a2a_protocol/exception_mapping_utils.py @@ -23,6 +23,12 @@ if TYPE_CHECKING: from a2a.client import Client as A2AClientType +_A2A_KIND_DIALECT_HINT: Final = ( + " (the agent replied in the A2A 0.3 dialect while its card declares 1.0; " + 'set litellm_params.a2a_protocol_version: "0.3" on this agent)' +) + + try: from a2a.client import Client, ClientConfig, create_client @@ -152,7 +158,7 @@ def map_a2a_exception( # Default: wrap in generic A2AError raise A2AError( - message=error_str, + message=error_str + _A2A_KIND_DIALECT_HINT if 'has no field named "kind"' in error_str else error_str, model=model, ) diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index bad17f05923..d664acb5d58 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -15,7 +15,7 @@ from typing import Any, Final import litellm from litellm._logging import verbose_logger -from litellm.a2a_protocol.card_resolver import AGENT_CARD_PATH_PARAM +from litellm.a2a_protocol.card_resolver import A2A_PROTOCOL_VERSION_PARAM, AGENT_CARD_PATH_PARAM from litellm.a2a_protocol.litellm_completion_bridge.transformation import ( A2ACompletionBridgeTransformation, A2AStreamingContext, @@ -38,6 +38,7 @@ _AGENT_ONLY_PARAMS: Final = frozenset( "agent_id", "agent_card_params", AGENT_CARD_PATH_PARAM, + A2A_PROTOCOL_VERSION_PARAM, A2A_USER_API_KEY_HASH_PARAM, } ) diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 3a1d2c70b12..dafe041ba3f 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -72,10 +72,12 @@ except ImportError: # Import our custom card resolver that supports multiple well-known paths from litellm.a2a_protocol.card_resolver import ( + A2A_PROTOCOL_VERSION_PARAM, AGENT_CARD_PATH_PARAM, LiteLLMA2ACardResolver, get_agent_card_url, normalize_agent_card_interfaces, + resolve_a2a_protocol_version, ) from litellm.a2a_protocol.exception_mapping_utils import ( handle_a2a_localhost_retry, @@ -153,6 +155,10 @@ def _agent_card_path(litellm_params: Mapping[str, object]) -> str | None: return configured_path if isinstance(configured_path, str) and configured_path else None +def _agent_protocol_version(litellm_params: Mapping[str, object]) -> str | None: + return resolve_a2a_protocol_version(litellm_params.get(A2A_PROTOCOL_VERSION_PARAM)) + + def _set_litellm_params_on_logging_obj( kwargs: Mapping[str, object], litellm_params: Mapping[str, object], @@ -498,6 +504,7 @@ async def asend_message( base_url=api_base, extra_headers=extra_headers, relative_card_path=_agent_card_path(litellm_params), + protocol_version=_agent_protocol_version(litellm_params), ) # Type assertion: a2a_client is guaranteed to be non-None here @@ -723,6 +730,7 @@ async def asend_message_streaming( extra_headers=extra_headers, streaming=True, relative_card_path=_agent_card_path(litellm_params), + protocol_version=_agent_protocol_version(litellm_params), ) assert a2a_client is not None @@ -770,6 +778,7 @@ async def create_a2a_client( extra_headers: dict[str, str] | None = None, streaming: bool = False, relative_card_path: str | None = None, + protocol_version: str | None = None, ) -> "A2AClientType": """ Create an A2A client for the given agent URL. @@ -783,6 +792,7 @@ async def create_a2a_client( extra_headers: Optional additional headers to include in requests relative_card_path: Optional card path relative to ``base_url`` (e.g. ``agentCard/v1.0`` for a Microsoft Foundry agent); when None the well-known paths are probed in order + protocol_version: Optional per-agent protocol version override for supported interfaces Returns: An initialized a2a.client.A2AClient instance @@ -819,7 +829,8 @@ async def create_a2a_client( await resolver.get_agent_card( relative_card_path=relative_card_path, http_kwargs=_card_http_kwargs(extra_headers), - ) + ), + protocol_version=protocol_version, ) a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall] diff --git a/tests/unit/a2a_protocol/test_a2a_exception_mapping_utils.py b/tests/unit/a2a_protocol/test_a2a_exception_mapping_utils.py index 5f097570bc2..4043a95c794 100644 --- a/tests/unit/a2a_protocol/test_a2a_exception_mapping_utils.py +++ b/tests/unit/a2a_protocol/test_a2a_exception_mapping_utils.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from litellm.a2a_protocol import exception_mapping_utils as emu -from litellm.a2a_protocol.exceptions import A2ALocalhostURLError +from litellm.a2a_protocol.exceptions import A2AError, A2ALocalhostURLError def _localhost_error() -> A2ALocalhostURLError: @@ -17,6 +17,20 @@ def _localhost_error() -> A2ALocalhostURLError: ) +def test_map_a2a_exception_adds_protocol_version_hint_for_kind_field_error(): + with pytest.raises(A2AError) as error: + emu.map_a2a_exception(RuntimeError('Message type "lf.a2a.v1.Task" has no field named "kind"')) + + assert "a2a_protocol_version: " in str(error.value) + + +def test_map_a2a_exception_does_not_add_protocol_version_hint_for_other_errors(): + with pytest.raises(A2AError) as error: + emu.map_a2a_exception(RuntimeError("unrelated agent error")) + + assert "a2a_protocol_version" not in str(error.value) + + @pytest.mark.asyncio async def test_localhost_retry_reuses_stashed_httpx_client(): """The retry must reuse the httpx client LiteLLM attached at creation (it carries diff --git a/tests/unit/a2a_protocol/test_card_resolver.py b/tests/unit/a2a_protocol/test_card_resolver.py index fdfb51987a3..4cfa8dab45b 100644 --- a/tests/unit/a2a_protocol/test_card_resolver.py +++ b/tests/unit/a2a_protocol/test_card_resolver.py @@ -16,6 +16,7 @@ from litellm.a2a_protocol.card_resolver import ( fix_agent_card_url, is_localhost_or_internal_url, normalize_agent_card_interfaces, + resolve_a2a_protocol_version, set_agent_card_url, ) from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError @@ -139,6 +140,77 @@ def test_normalize_agent_card_interfaces_downgrades_miscased_interfaces_to_the_0 assert card.supported_interfaces[0].protocol_version == "1.0" +def test_normalize_agent_card_interfaces_applies_override_to_canonical_bindings_without_mutating_input(): + pb2 = pytest.importorskip("a2a.types.a2a_pb2") + card = pb2.AgentCard( + name="langgraph", + supported_interfaces=[ + pb2.AgentInterface(url="http://a/", protocol_binding="jsonrpc", protocol_version="1.0"), + pb2.AgentInterface(url="http://b/", protocol_binding="JSONRPC", protocol_version="1.0"), + pb2.AgentInterface(url="http://c/", protocol_binding="HTTP+JSON", protocol_version="1.0"), + pb2.AgentInterface(url="http://d/", protocol_binding="GRPC", protocol_version="1.0"), + pb2.AgentInterface(url="http://e/", protocol_binding="websocket", protocol_version="1.0"), + ], + ) + + normalized = normalize_agent_card_interfaces(card, protocol_version="0.3") + + assert [(item.protocol_binding, item.protocol_version) for item in normalized.supported_interfaces] == [ + ("JSONRPC", "0.3"), + ("JSONRPC", "0.3"), + ("HTTP+JSON", "0.3"), + ("GRPC", "0.3"), + ("websocket", "1.0"), + ] + assert card.supported_interfaces[0].protocol_binding == "jsonrpc" + assert [(item.protocol_binding, item.protocol_version) for item in card.supported_interfaces] == [ + ("jsonrpc", "1.0"), + ("JSONRPC", "1.0"), + ("HTTP+JSON", "1.0"), + ("GRPC", "1.0"), + ("websocket", "1.0"), + ] + + +def test_normalize_agent_card_interfaces_override_opts_out_of_mis_cased_downgrade(): + pb2 = pytest.importorskip("a2a.types.a2a_pb2") + card = pb2.AgentCard( + name="langgraph", + supported_interfaces=[ + pb2.AgentInterface(url="http://a/", protocol_binding="jsonrpc", protocol_version="1.0"), + ], + ) + + normalized = normalize_agent_card_interfaces(card, protocol_version="1.0") + + assert normalized.supported_interfaces[0].protocol_binding == "JSONRPC" + assert normalized.supported_interfaces[0].protocol_version == "1.0" + assert card.supported_interfaces[0].protocol_binding == "jsonrpc" + assert card.supported_interfaces[0].protocol_version == "1.0" + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("0.3", "0.3"), + (" 0.3 ", "0.3"), + ("0.3.0", "0.3"), + (0.3, "0.3"), + (1.0, "1.0"), + ("1.0.0", "1.0"), + (None, None), + (True, None), + (1, None), + ("2.0", None), + ("abc", None), + ("", None), + (["0.3"], None), + ], +) +def test_resolve_a2a_protocol_version(value, expected): + assert resolve_a2a_protocol_version(value) == expected + + _FOUNDRY_BASE_URL: Final = "https://foundry.example.com/a2a" _FOUNDRY_CARD_JSON: Final = { diff --git a/tests/unit/a2a_protocol/test_completion_bridge_streaming.py b/tests/unit/a2a_protocol/test_completion_bridge_streaming.py index 913c917bd2d..2e7cac5e7be 100644 --- a/tests/unit/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/unit/a2a_protocol/test_completion_bridge_streaming.py @@ -356,6 +356,7 @@ async def test_handle_streaming_keeps_agent_card_path_out_of_the_completion_call "custom_llm_provider": "langgraph", "model": "agent", "agent_card_path": "agentCard/v1.0", + "a2a_protocol_version": "0.3", }, api_base="http://localhost:2024", ) @@ -363,3 +364,4 @@ async def test_handle_streaming_keeps_agent_card_path_out_of_the_completion_call assert len(events) == 4 assert "agent_card_path" not in mock_acompletion.call_args.kwargs + assert "a2a_protocol_version" not in mock_acompletion.call_args.kwargs diff --git a/tests/unit/a2a_protocol/test_main.py b/tests/unit/a2a_protocol/test_main.py index 4ba0ef8fa04..eead12e3cf7 100644 --- a/tests/unit/a2a_protocol/test_main.py +++ b/tests/unit/a2a_protocol/test_main.py @@ -1,6 +1,7 @@ """Tests for litellm/a2a_protocol/main.py non-streaming send behavior.""" import asyncio +import json import httpx import pytest @@ -23,6 +24,7 @@ from litellm.a2a_protocol.main import ( asend_message, create_a2a_client, ) +from litellm.a2a_protocol.exceptions import A2AError from litellm.caching.llm_caching_handler import LLMClientCache from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT from litellm.llms.custom_httpx.http_handler import ( @@ -230,6 +232,16 @@ _LOWERCASE_BINDING_CARD = { "supportedInterfaces": [{"url": "http://127.0.0.1:9/", "protocolBinding": "jsonrpc", "protocolVersion": "1.0"}], } +_UPPERCASE_BINDING_CARD = { + "name": "langgraph-agent", + "version": "1.0.0", + "capabilities": {"streaming": True}, + "defaultInputModes": ["text/plain"], + "defaultOutputModes": ["text/plain"], + "skills": [], + "supportedInterfaces": [{"url": "http://127.0.0.1:9/", "protocolBinding": "JSONRPC", "protocolVersion": "1.0"}], +} + class _RequestRecorder: """Records the headers httpx put on the wire, per outbound request.""" @@ -240,6 +252,7 @@ class _RequestRecorder: self.card_requests = [] self.card_urls = [] self.rpc_requests = [] + self.rpc_bodies = [] self.client = None def __call__(self, request: httpx.Request) -> httpx.Response: @@ -249,6 +262,7 @@ class _RequestRecorder: self.card_urls.append(str(request.url)) return httpx.Response(200, json=self.card) self.rpc_requests.append(headers) + self.rpc_bodies.append(json.loads(request.content)) return httpx.Response(200, json=self.rpc_reply) @@ -392,6 +406,79 @@ async def test_lowercase_protocol_binding_card_round_trips_the_langgraph_dialect assert interface.protocol_version == "0.3" +@pytest.mark.asyncio +async def test_uppercase_binding_card_without_override_fails_with_actionable_hint(isolated_client_cache): + await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_reply=_LANGGRAPH_TASK_REPLY) + + with pytest.raises(A2AError) as error: + await asend_message(request=_send_request("uppercase-default"), api_base="http://127.0.0.1:9") + + assert 'has no field named "kind"' in str(error.value) + assert "a2a_protocol_version" in str(error.value) + + +@pytest.mark.parametrize("card", [_LOWERCASE_BINDING_CARD, _UPPERCASE_BINDING_CARD]) +@pytest.mark.asyncio +async def test_protocol_version_override_round_trips_the_langgraph_dialect(card, isolated_client_cache): + recorder = await _seed_shared_a2a_client(card=card, rpc_reply=_LANGGRAPH_TASK_REPLY) + + a2a_client = await create_a2a_client(base_url="http://127.0.0.1:9", protocol_version="0.3") + response = await _send_message(a2a_client, _send_request("override-round-trip")) + + assert type(response.root.result).__name__ == "Task" + assert response.root.result.artifacts[0].parts[0].root.text == "langgraph echo: hi" + assert a2a_client._litellm_agent_card.supported_interfaces[0].protocol_version == "0.3" + assert recorder.rpc_bodies[-1]["method"] == "message/send" + + +@pytest.mark.asyncio +async def test_asend_message_resolves_float_protocol_version_from_agent_params(isolated_client_cache): + recorder = await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_reply=_LANGGRAPH_TASK_REPLY) + + response = await asend_message( + request=_send_request("float-protocol-override"), + api_base="http://127.0.0.1:9", + litellm_params={"a2a_protocol_version": 0.3}, + ) + + assert response.result["artifacts"][0]["parts"][0]["text"] == "langgraph echo: hi" + assert recorder.rpc_bodies[-1]["method"] == "message/send" + + +@pytest.mark.parametrize( + ("litellm_params", "expected_protocol_version"), + [({"a2a_protocol_version": "0.3"}, "0.3"), (None, None)], +) +@pytest.mark.asyncio +async def test_streaming_passes_agent_protocol_version_to_client(litellm_params, expected_protocol_version): + from unittest.mock import AsyncMock, patch + + from litellm.a2a_protocol import main as a2a_main + + request = SendStreamingMessageRequest( + id="stream-protocol-version", + params=MessageSendParams( + message={"messageId": "m1", "role": "user", "parts": [{"kind": "text", "text": "hi"}]} + ), + ) + captured = {} + + async def _capture(*, base_url, extra_headers=None, streaming=False, protocol_version=None, **_): + captured["protocol_version"] = protocol_version + raise RuntimeError("stop") + + with patch.object(a2a_main, "create_a2a_client", new=AsyncMock(side_effect=_capture)): + with pytest.raises(RuntimeError, match="stop"): + async for _ in a2a_main.asend_message_streaming( + request=request, + api_base="http://upstream.local", + litellm_params=litellm_params, + ): + pass + + assert captured["protocol_version"] == expected_protocol_version + + @pytest.mark.asyncio async def test_agent_card_fetch_carries_the_callers_headers(isolated_client_cache): """Agent cards can sit behind the same auth as the agent, so the card fetch must stay From ba709d3de5fd60b4180566a13283c3919ce924f4 Mon Sep 17 00:00:00 2001 From: shivam Date: Sat, 3 Oct 2026 23:55:34 +0000 Subject: [PATCH 2/2] test(a2a): cover streamed 0.3 replies through the real client Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/unit/a2a_protocol/test_main.py | 128 ++++++++++++++++++++++++++- 1 file changed, 125 insertions(+), 3 deletions(-) diff --git a/tests/unit/a2a_protocol/test_main.py b/tests/unit/a2a_protocol/test_main.py index eead12e3cf7..4327a1d0bc1 100644 --- a/tests/unit/a2a_protocol/test_main.py +++ b/tests/unit/a2a_protocol/test_main.py @@ -222,6 +222,77 @@ _LANGGRAPH_TASK_REPLY = { } +_LANGGRAPH_STREAM_EVENTS = [ + { + "jsonrpc": "2.0", + "id": "reply", + "result": { + "kind": "task", + "id": "run-1:task-1", + "contextId": "thread-1", + "history": [ + { + "kind": "message", + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "m-user", + "taskId": "run-1:task-1", + "contextId": "thread-1", + } + ], + "status": {"state": "submitted"}, + }, + }, + { + "jsonrpc": "2.0", + "id": "reply", + "result": { + "kind": "status-update", + "taskId": "run-1:task-1", + "contextId": "thread-1", + "status": { + "state": "working", + "message": { + "kind": "message", + "role": "agent", + "parts": [{"kind": "text", "text": "langgraph echo: hi"}], + "messageId": "m-agent", + "taskId": "run-1:task-1", + "contextId": "thread-1", + }, + }, + "final": False, + }, + }, + { + "jsonrpc": "2.0", + "id": "reply", + "result": { + "kind": "artifact-update", + "taskId": "run-1:task-1", + "contextId": "thread-1", + "artifact": { + "artifactId": "art-1", + "name": "Assistant Response", + "parts": [{"kind": "text", "text": "langgraph echo: hi"}], + }, + "lastChunk": True, + }, + }, + { + "jsonrpc": "2.0", + "id": "reply", + "result": { + "kind": "status-update", + "taskId": "run-1:task-1", + "contextId": "thread-1", + "status": {"state": "completed"}, + "final": True, + }, + }, +] + + _LOWERCASE_BINDING_CARD = { "name": "langgraph-agent", "version": "1.0.0", @@ -246,9 +317,10 @@ _UPPERCASE_BINDING_CARD = { class _RequestRecorder: """Records the headers httpx put on the wire, per outbound request.""" - def __init__(self, card=_AGENT_CARD, rpc_reply=_RPC_REPLY): + def __init__(self, card=_AGENT_CARD, rpc_reply=_RPC_REPLY, rpc_stream_events=None): self.card = card self.rpc_reply = rpc_reply + self.rpc_stream_events = rpc_stream_events self.card_requests = [] self.card_urls = [] self.rpc_requests = [] @@ -263,6 +335,9 @@ class _RequestRecorder: return httpx.Response(200, json=self.card) self.rpc_requests.append(headers) self.rpc_bodies.append(json.loads(request.content)) + if self.rpc_stream_events is not None: + body = "".join(f"data: {json.dumps(event)}\n\n" for event in self.rpc_stream_events) + return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=body) return httpx.Response(200, json=self.rpc_reply) @@ -271,7 +346,10 @@ def _a2a_client_cache_key(timeout: float, provider: str = httpxSpecialProvider.A async def _seed_shared_a2a_client( - card=_AGENT_CARD, rpc_reply=_RPC_REPLY, provider: str = httpxSpecialProvider.A2AProvider + card=_AGENT_CARD, + rpc_reply=_RPC_REPLY, + rpc_stream_events=None, + provider: str = httpxSpecialProvider.A2AProvider, ) -> _RequestRecorder: """Put the one A2A client the cache will hand out behind a mock transport. @@ -279,7 +357,7 @@ async def _seed_shared_a2a_client( it. The injected client is a real httpx.AsyncClient, so the merge of per-request headers over client defaults, which is what these tests are about, stays real. """ - recorder = _RequestRecorder(card=card, rpc_reply=rpc_reply) + recorder = _RequestRecorder(card=card, rpc_reply=rpc_reply, rpc_stream_events=rpc_stream_events) handler = AsyncHTTPHandler(timeout=DEFAULT_A2A_AGENT_TIMEOUT) owned_client = handler.client handler.client = httpx.AsyncClient(transport=httpx.MockTransport(recorder)) @@ -431,6 +509,50 @@ async def test_protocol_version_override_round_trips_the_langgraph_dialect(card, assert recorder.rpc_bodies[-1]["method"] == "message/send" +@pytest.mark.parametrize("card", [_LOWERCASE_BINDING_CARD, _UPPERCASE_BINDING_CARD]) +@pytest.mark.asyncio +async def test_protocol_version_override_streams_the_langgraph_dialect(card, isolated_client_cache): + recorder = await _seed_shared_a2a_client(card=card, rpc_stream_events=_LANGGRAPH_STREAM_EVENTS) + + a2a_client = await create_a2a_client( + base_url="http://127.0.0.1:9", + streaming=True, + protocol_version="0.3", + ) + request = SendStreamingMessageRequest( + id="reply", + params=MessageSendParams( + message={"messageId": "m-user", "role": "user", "parts": [{"kind": "text", "text": "hi"}]} + ), + ) + streamed = [chunk async for chunk in _stream_messages(a2a_client, request)] + + artifact_event = next(chunk.root.result for chunk in streamed if chunk.root.result.kind == "artifact-update") + completed_status = next( + chunk.root.result for chunk in streamed if chunk.root.result.kind == "status-update" and chunk.root.result.final + ) + assert artifact_event.artifact.parts[0].root.text == "langgraph echo: hi" + assert completed_status.status.state.value == "completed" + assert recorder.rpc_bodies[-1]["method"] == "message/stream" + + +@pytest.mark.asyncio +async def test_uppercase_binding_card_stream_without_override_fails(isolated_client_cache): + await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_stream_events=_LANGGRAPH_STREAM_EVENTS) + + a2a_client = await create_a2a_client(base_url="http://127.0.0.1:9", streaming=True) + request = SendStreamingMessageRequest( + id="reply", + params=MessageSendParams( + message={"messageId": "m-user", "role": "user", "parts": [{"kind": "text", "text": "hi"}]} + ), + ) + + with pytest.raises(Exception, match='has no field named "kind"'): + async for _ in _stream_messages(a2a_client, request): + pass + + @pytest.mark.asyncio async def test_asend_message_resolves_float_protocol_version_from_agent_params(isolated_client_cache): recorder = await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_reply=_LANGGRAPH_TASK_REPLY)