mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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
This commit is contained in:
parent
59cab4d2aa
commit
0c3c1d4c70
2 changed files with 135 additions and 9 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue