From 0c3c1d4c709dc8aa918b67b6705c9919ac1a193d Mon Sep 17 00:00:00 2001 From: shin-bot-litellm Date: Tue, 3 Feb 2026 22:45:05 +0000 Subject: [PATCH] fix(a2a): use text/event-stream SSE format for message/stream endpoint The A2A gateway's streaming response was using application/x-ndjson Content-Type and raw NDJSON body format. The A2A protocol spec requires text/event-stream with SSE framing (data: ...\n\n). The official a2a-sdk client validates the Content-Type header and raises SSEError when it doesn't contain text/event-stream. Changes: - Changed media_type from application/x-ndjson to text/event-stream - Updated response body to use SSE framing (data: prefix + \n\n suffix) - Added tests validating Content-Type and SSE body format Fixes #20278 --- .../proxy/agent_endpoints/a2a_endpoints.py | 18 +-- .../agent_endpoints/test_a2a_endpoints.py | 126 ++++++++++++++++++ 2 files changed, 135 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 24727aacd75..338c054c177 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -62,7 +62,7 @@ async def _handle_stream_message( if not A2A_SDK_AVAILABLE: # Return a streaming response that yields an error async def _error_stream(): - yield json.dumps( + yield "data: " + json.dumps( { "jsonrpc": "2.0", "id": request_id, @@ -71,9 +71,9 @@ async def _handle_stream_message( "message": "Server error: 'a2a' package not installed", }, } - ) + "\n" + ) + "\n\n" - return StreamingResponse(_error_stream(), media_type="application/x-ndjson") + return StreamingResponse(_error_stream(), media_type="text/event-stream") from a2a.types import ( MessageSendParams, @@ -96,22 +96,22 @@ async def _handle_stream_message( ): # Chunk may be dict or object depending on bridge vs standard path if hasattr(chunk, "model_dump"): - yield json.dumps( + yield "data: " + json.dumps( chunk.model_dump(mode="json", exclude_none=True) - ) + "\n" + ) + "\n\n" else: - yield json.dumps(chunk) + "\n" + yield "data: " + json.dumps(chunk) + "\n\n" except Exception as e: verbose_proxy_logger.exception(f"Error streaming A2A response: {e}") - yield json.dumps( + yield "data: " + json.dumps( { "jsonrpc": "2.0", "id": request_id, "error": {"code": -32603, "message": f"Streaming error: {str(e)}"}, } - ) + "\n" + ) + "\n\n" - return StreamingResponse(stream_response(), media_type="application/x-ndjson") + return StreamingResponse(stream_response(), media_type="text/event-stream") @router.get( diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 9588c3b55c3..135e24953a9 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -4,6 +4,7 @@ Mock tests for A2A endpoints. Tests that invoke_agent_a2a properly integrates with add_litellm_data_to_request. """ +import json import sys from unittest.mock import AsyncMock, MagicMock, patch @@ -168,3 +169,128 @@ async def test_invoke_agent_a2a_adds_litellm_data(): # Verify proxy_server_request was added assert "proxy_server_request" in captured_data assert captured_data["proxy_server_request"]["method"] == "POST" + + +@pytest.mark.asyncio +async def test_handle_stream_message_returns_sse_content_type(): + """ + Test that _handle_stream_message returns Content-Type: text/event-stream + with SSE-framed body (data: ...\\n\\n), not application/x-ndjson. + + The A2A protocol spec requires text/event-stream for streaming responses. + The official a2a-sdk client validates this header. + + Ref: https://github.com/BerriAI/litellm/issues/20278 + """ + # Mock chunk with model_dump + mock_chunk = MagicMock() + mock_chunk.model_dump.return_value = { + "jsonrpc": "2.0", + "id": "test-id", + "result": {"kind": "status-update"}, + } + + async def mock_streaming(*args, **kwargs): + yield mock_chunk + + # Try to use real a2a.types if available + try: + from a2a.types import ( + MessageSendParams, + SendStreamingMessageRequest, + ) + except ImportError: + + class MessageSendParams: + def __init__(self, **kwargs): + self.__dict__.update(kwargs) + + class SendStreamingMessageRequest: + def __init__(self, **kwargs): + self.__dict__.update(kwargs) + + mock_a2a_types = MagicMock() + mock_a2a_types.MessageSendParams = MessageSendParams + mock_a2a_types.SendStreamingMessageRequest = SendStreamingMessageRequest + + with patch.dict( + sys.modules, + {"a2a": MagicMock(), "a2a.types": mock_a2a_types}, + ), patch( + "litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", + True, + ), patch( + "litellm.a2a_protocol.asend_message_streaming", + side_effect=mock_streaming, + ): + from litellm.proxy.agent_endpoints.a2a_endpoints import ( + _handle_stream_message, + ) + + response = await _handle_stream_message( + api_base="http://backend:10001", + request_id="test-id", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "Hello"}], + "messageId": "msg-123", + } + }, + ) + + # Verify Content-Type is text/event-stream (required by A2A spec) + assert response.media_type == "text/event-stream" + + # Collect streamed body and verify SSE framing + body_parts = [] + async for chunk in response.body_iterator: + body_parts.append(chunk) + + assert len(body_parts) > 0 + for part in body_parts: + # Each SSE event must start with "data: " and end with "\n\n" + assert part.startswith("data: "), ( + f"SSE event must start with 'data: ', got: {part!r}" + ) + assert part.endswith("\n\n"), ( + f"SSE event must end with '\\n\\n', got: {part!r}" + ) + # The payload between "data: " and "\n\n" must be valid JSON + payload = part[len("data: "):-2] + parsed = json.loads(payload) + assert isinstance(parsed, dict) + + +@pytest.mark.asyncio +async def test_handle_stream_message_error_uses_sse_format(): + """ + Test that when A2A SDK is not available, the error stream also uses + text/event-stream with SSE framing. + """ + with patch( + "litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", + False, + ): + from litellm.proxy.agent_endpoints.a2a_endpoints import ( + _handle_stream_message, + ) + + response = await _handle_stream_message( + api_base=None, + request_id="err-id", + params={}, + ) + + assert response.media_type == "text/event-stream" + + body_parts = [] + async for chunk in response.body_iterator: + body_parts.append(chunk) + + assert len(body_parts) == 1 + part = body_parts[0] + assert part.startswith("data: ") + assert part.endswith("\n\n") + payload = json.loads(part[len("data: "):-2]) + assert payload["error"]["code"] == -32603