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)