fix(a2a): return SSE (text/event-stream) for message/stream instead of NDJSON (#35037)

* fix(a2a): return SSE (text/event-stream) for message/stream instead of NDJSON

* test(a2a): cover message/stream SSE framing on proxy-hook and sdk-unavailable paths

* test(a2a): cover SSE error framing paths for streaming

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(a2a): send sse keepalive pings on message/stream

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: yassin <yassin@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-08-20 16:19:58 -07:00 • committed by GitHub
parent d40aea865d
commit 3207014906
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 566 additions and 153 deletions

View file

@ -77,6 +77,27 @@ _PASCAL_TO_WIRE: Final[Mapping[str, str]] = {
}
def _sse_event(payload: object) -> str:
"""Frame a JSON-RPC object as a single A2A SSE event (``data: <json>\\n\\n``)."""
return f"data: {json.dumps(payload)}\n\n"
def _to_jsonrpc_object(chunk: object) -> object:
"""Coerce a streamed chunk to the JSON-RPC object it carries.
Chunks arrive as SDK models, plain dicts, or, when a guardrail terminates a
stream, as an already serialized JSON-RPC object.
"""
if isinstance(chunk, (str, bytes, bytearray)):
try:
return json.loads(chunk)
except (json.JSONDecodeError, UnicodeDecodeError):
return chunk
if hasattr(chunk, "model_dump"):
return chunk.model_dump(mode="json", exclude_none=True)
return chunk
def _build_message_send_params(params: dict[str, Any]) -> "MessageSendParams":
"""Build MessageSendParams from wire (0.3) or A2A 1.0 JSON-RPC params."""
from a2a.compat.v0_3.types import MessageSendParams
@ -280,6 +301,22 @@ async def _a2a_sse_event_source(
await resp.aclose()
def _sse_streaming_response(generator: AsyncGenerator[str, None]) -> StreamingResponse:
# The upstream agent is only contacted once this generator is first pulled, so
# a slow first event leaves the response body idle for its whole
# time-to-first-token and an intermediary with an idle read timeout drops a
# healthy connection. Off until an operator sets an interval, and the
# buffering hint only goes out when there are keepalives to protect.
keepalive_interval: Final = coerce_keepalive_interval(litellm.sse_keepalive_ping_interval_seconds)
if keepalive_interval is None:
return StreamingResponse(generator, media_type="text/event-stream")
return StreamingResponse(
wrap_sse_stream_with_keepalive_pings(generator, keepalive_interval, ping_chunk=SSE_COMMENT_PING),
media_type="text/event-stream",
headers=_SSE_KEEPALIVE_HEADERS,
)
async def _forward_jsonrpc_sse(
agent_url: str,
body: Mapping[str, object],
@ -341,19 +378,7 @@ async def _forward_jsonrpc_sse(
generator = _passthrough()
# The upstream agent is only contacted once this generator is first pulled, so
# a slow first event leaves the response body idle for its whole
# time-to-first-token and an intermediary with an idle read timeout drops a
# healthy connection. Off until an operator sets an interval, and the
# buffering hint only goes out when there are keepalives to protect.
keepalive_interval: Final = coerce_keepalive_interval(litellm.sse_keepalive_ping_interval_seconds)
if keepalive_interval is None:
return StreamingResponse(generator, media_type="text/event-stream")
return StreamingResponse(
wrap_sse_stream_with_keepalive_pings(generator, keepalive_interval, ping_chunk=SSE_COMMENT_PING),
media_type="text/event-stream",
headers=_SSE_KEEPALIVE_HEADERS,
)
return _sse_streaming_response(generator)
async def _handle_stream_message(
@ -373,9 +398,12 @@ async def _handle_stream_message(
) -> StreamingResponse:
"""Handle message/stream method via SDK functions.
When user_api_key_dict, request_data, and proxy_logging_obj are provided,
uses common_request_processing.async_streaming_data_generator with NDJSON
serializers so proxy hooks and cost injection apply.
The A2A JSON-RPC binding streams responses as SSE (text/event-stream) with
each JSON-RPC object framed as ``data: <json>\n\n``, matching the official
a2a-sdk client which rejects any other Content-Type. When user_api_key_dict,
request_data, and proxy_logging_obj are provided, events are routed through
common_request_processing.async_streaming_data_generator so proxy hooks and
cost injection apply.
"""
from litellm.a2a_protocol import asend_message_streaming
from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE
@ -383,21 +411,18 @@ async def _handle_stream_message(
if not A2A_SDK_AVAILABLE:
async def _error_stream():
yield (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": "Server error: 'a2a' package not installed",
},
}
)
+ "\n"
yield _sse_event(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": "Server error: 'a2a' package not installed",
},
}
)
return StreamingResponse(_error_stream(), media_type="application/x-ndjson")
return StreamingResponse(_error_stream(), media_type="text/event-stream")
from a2a.compat.v0_3.types import SendStreamingMessageRequest
@ -409,18 +434,21 @@ async def _handle_stream_message(
invalid_params_message: Final = f"Invalid params: {e}"
async def _invalid_params_stream():
yield (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {"code": -32602, "message": invalid_params_message},
}
)
+ "\n"
yield _sse_event(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {"code": -32602, "message": invalid_params_message},
}
)
return StreamingResponse(_invalid_params_stream(), media_type="application/x-ndjson")
return StreamingResponse(_invalid_params_stream(), media_type="text/event-stream")
def _sse_chunk(chunk: object) -> str:
obj = _to_jsonrpc_object(chunk)
if isinstance(obj, dict):
obj = normalize_stream_event(obj, served_version, request_id=request_id)
return _sse_event(obj)
async def stream_response():
try:
@ -448,32 +476,20 @@ async def _handle_stream_message(
ProxyBaseLLMRequestProcessing,
)
def _ndjson_chunk(chunk: Any) -> str:
if hasattr(chunk, "model_dump"):
obj = chunk.model_dump(mode="json", exclude_none=True)
else:
obj = chunk
if isinstance(obj, dict):
obj = normalize_stream_event(obj, served_version, request_id=request_id)
return json.dumps(obj) + "\n"
def _ndjson_error(proxy_exc: object) -> str:
return (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": getattr(
proxy_exc,
"message",
f"Streaming error: {proxy_exc}",
),
},
}
)
+ "\n"
def _sse_error(proxy_exc: object) -> str:
return _sse_event(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": getattr(
proxy_exc,
"message",
f"Streaming error: {proxy_exc}",
),
},
}
)
async for line in ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
@ -481,19 +497,13 @@ async def _handle_stream_message(
user_api_key_dict=user_api_key_dict,
request_data=request_data,
proxy_logging_obj=proxy_logging_obj,
serialize_chunk=_ndjson_chunk,
serialize_error=_ndjson_error,
serialize_chunk=_sse_chunk,
serialize_error=_sse_error,
):
yield line
else:
async for chunk in a2a_stream:
if hasattr(chunk, "model_dump"):
obj = chunk.model_dump(mode="json", exclude_none=True)
else:
obj = chunk
if isinstance(obj, dict):
obj = normalize_stream_event(obj, served_version, request_id=request_id)
yield json.dumps(obj) + "\n"
yield _sse_chunk(chunk)
except Exception as e:
verbose_proxy_logger.exception("Error streaming A2A response: %s", e)
if (
@ -511,21 +521,18 @@ async def _handle_stream_message(
e = transformed_exception
if isinstance(e, HTTPException):
raise
yield (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": f"Streaming error: {e}",
},
}
)
+ "\n"
yield _sse_event(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": f"Streaming error: {e}",
},
}
)
return StreamingResponse(stream_response(), media_type="application/x-ndjson")
return _sse_streaming_response(stream_response())
@router.get(

View file

@ -38,7 +38,7 @@ if TYPE_CHECKING:
BaseTranslation,
)
# Call types that use NDJSON streaming (A2A); guardrail HTTPException is emitted as in-stream error
# Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error
A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message)
GUARDRAIL_NAME: Final = "unified_llm_guardrails"
@ -90,6 +90,24 @@ def _get_a2a_request_id(responses_so_far: Sequence[object], request_data: dict)
return None
def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapping[str, object]:
"""Build the in-stream JSON-RPC error object for a mid-stream A2A failure.
Returned as an object, not a serialized string: the A2A endpoint owns wire
framing and serializes whatever the stream yields.
"""
detail: Final = exc.detail if isinstance(exc.detail, dict) else {"message": str(exc.detail)}
return {
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get("error", detail.get("message", str(exc.detail))),
"data": {k: v for k, v in detail.items() if k not in ("error", "message")},
},
}
endpoint_guardrail_translation_mappings = None
@ -391,28 +409,12 @@ class UnifiedLLMGuardrails(CustomLogger):
responses_so_far: Sequence[object],
request_data: dict,
) -> AsyncGenerator[object, None]:
"""Surface a mid-stream HTTPException. For A2A (NDJSON) call types the
response has already started, so emit an in-stream JSON-RPC error chunk;
otherwise re-raise so the proxy can report it.
"""Surface a mid-stream HTTPException. For A2A call types the response has
already started, so emit an in-stream JSON-RPC error chunk; otherwise
re-raise so the proxy can report it.
"""
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
request_id: Final = _get_a2a_request_id(responses_so_far, request_data)
detail: Final = exc.detail if isinstance(exc.detail, dict) else {"message": str(exc.detail)}
error_chunk: Final = (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get("error", detail.get("message", str(exc.detail))),
"data": {k: v for k, v in detail.items() if k not in ("error", "message")},
},
}
)
+ "\n"
)
yield error_chunk
yield _a2a_jsonrpc_error_chunk(exc, _get_a2a_request_id(responses_so_far, request_data))
return
raise exc
@ -1068,28 +1070,9 @@ class UnifiedLLMGuardrails(CustomLogger):
return
except HTTPException as e:
# Response already started (we already yielded chunks); cannot send 400.
# For A2A (NDJSON), yield an in-stream JSON-RPC error so the client sees it.
# For A2A, yield an in-stream JSON-RPC error so the client sees it.
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
request_id = _get_a2a_request_id(responses_so_far, request_data)
detail = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)}
error_chunk = (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get(
"error",
detail.get("message", str(e.detail)),
),
"data": {k: v for k, v in detail.items() if k not in ("error", "message")},
},
}
)
+ "\n"
)
yield error_chunk
yield _a2a_jsonrpc_error_chunk(e, _get_a2a_request_id(responses_so_far, request_data))
return
raise
chunks_yielded = True
@ -1151,22 +1134,6 @@ class UnifiedLLMGuardrails(CustomLogger):
return
except HTTPException as e:
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
request_id = _get_a2a_request_id(responses_so_far, request_data)
detail = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)}
error_chunk = (
json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": detail.get("error", detail.get("message", str(e.detail))),
"data": {k: v for k, v in detail.items() if k not in ("error", "message")},
},
}
)
+ "\n"
)
yield error_chunk
yield _a2a_jsonrpc_error_chunk(e, _get_a2a_request_id(responses_so_far, request_data))
else:
raise

View file

@ -1317,15 +1317,373 @@ async def test_handle_stream_message_rejects_invalid_params_with_32602():
request_id="req-1",
params={"message": 12345},
)
assert response.media_type == "text/event-stream"
chunks = [chunk async for chunk in response.body_iterator]
body = "".join(
chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks
)
payload = json.loads(body.strip())
assert body.startswith("data: ")
assert body.endswith("\n\n")
payload = json.loads(body.removeprefix("data: ").strip())
assert payload["error"]["code"] == -32602
assert payload["id"] == "req-1"
@pytest.mark.asyncio
async def test_handle_stream_message_frames_events_as_sse():
"""message/stream must return text/event-stream with each JSON-RPC object
framed as ``data: <json>\\n\\n``. Regression for #35027: NDJSON framing
breaks the official a2a-sdk client, which requires SSE."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
events = [
{
"jsonrpc": "2.0",
"id": "req-1",
"result": {"kind": "task", "id": "t-1", "status": {"state": "working"}},
},
{
"jsonrpc": "2.0",
"id": "req-1",
"result": {"kind": "message", "parts": [{"kind": "text", "text": "pong"}]},
},
]
async def fake_stream(**kwargs):
for event in events:
yield event
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch(
"litellm.a2a_protocol.asend_message_streaming",
new=fake_stream,
)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
assert response.media_type == "text/event-stream"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == len(events)
for chunk, event in zip(chunks, events):
assert chunk.startswith("data: ")
assert chunk.endswith("\n\n")
assert json.loads(chunk.removeprefix("data: ").strip()) == event
@pytest.mark.asyncio
async def test_handle_stream_message_sdk_unavailable_frames_error_as_sse():
"""When the a2a package is unavailable the -32603 error must still be
emitted as a single SSE event so the a2a-sdk client can parse it."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
with patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", False):
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={"message": {"role": "user", "parts": []}},
)
assert response.media_type == "text/event-stream"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == 1
assert chunks[0].startswith("data: ")
assert chunks[0].endswith("\n\n")
payload = json.loads(chunks[0].removeprefix("data: ").strip())
assert payload["error"]["code"] == -32603
assert payload["id"] == "req-1"
@pytest.mark.asyncio
async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse():
"""When proxy hooks are wired the events are routed through
async_streaming_data_generator; that path must also frame each JSON-RPC
object as ``data: <json>\\n\\n`` (regression for #35027)."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
from litellm.proxy.utils import ProxyLogging
events = [
{"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}},
{
"jsonrpc": "2.0",
"id": "req-1",
"result": {"kind": "message", "parts": [{"kind": "text", "text": "pong"}]},
},
]
async def fake_stream(**kwargs):
for event in events:
yield event
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "a2a/test"},
proxy_logging_obj=proxy_logging_obj,
)
assert response.media_type == "text/event-stream"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == len(events)
for chunk, event in zip(chunks, events):
assert chunk.startswith("data: ")
assert chunk.endswith("\n\n")
assert json.loads(chunk.removeprefix("data: ").strip()) == event
@pytest.mark.asyncio
async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once():
"""A stream chunk that is already a serialized JSON-RPC object (what a
guardrail may yield when it terminates an A2A stream mid-flight) must be
framed as one SSE event carrying that object, not JSON-encoded a second time
into a bare string."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
error_event = {
"jsonrpc": "2.0",
"id": "req-1",
"error": {"code": -32603, "message": "blocked by guardrail", "data": {}},
}
async def fake_stream(**kwargs):
yield json.dumps(error_event) + "\n"
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == 1
payload = json.loads(chunks[0].removeprefix("data: ").strip())
assert payload == error_event
@pytest.mark.asyncio
async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse():
"""A failure while the hooked generator is streaming must reach the client as
a ``data:``-framed JSON-RPC error, not as a bare NDJSON line."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
from litellm.proxy.utils import ProxyLogging
async def fake_stream(**kwargs):
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
raise ValueError("upstream died")
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "a2a/test"},
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == 2
assert chunks[-1].startswith("data: ")
error_payload = json.loads(chunks[-1].removeprefix("data: ").strip())
assert error_payload["id"] == "req-1"
assert error_payload["error"]["code"] == -32603
assert "upstream died" in error_payload["error"]["message"]
@pytest.mark.asyncio
async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error():
"""A failure raised before any event is streamed (with proxy hooks wired) is
still delivered as a ``data:``-framed JSON-RPC error."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
from litellm.proxy.utils import ProxyLogging
def fake_stream(**kwargs):
raise ValueError("could not reach agent")
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "a2a/test"},
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == 1
error_payload = json.loads(chunks[0].removeprefix("data: ").strip())
assert error_payload["id"] == "req-1"
assert error_payload["error"]["code"] == -32603
assert "could not reach agent" in error_payload["error"]["message"]
@pytest.mark.asyncio
async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event():
"""A chunk that is not JSON at all still leaves as one well-formed SSE event
instead of raising and killing the stream."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
async def fake_stream(**kwargs):
yield "not json at all"
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert chunks == ['data: "not json at all"\n\n']
@pytest.mark.asyncio
async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error():
"""An upstream failure after the response started is reported as a
``data:``-framed JSON-RPC error object, so an SSE client sees the failure."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
async def fake_stream(**kwargs):
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
raise RuntimeError("upstream died")
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert len(chunks) == 2
error_payload = json.loads(chunks[-1].removeprefix("data: ").strip())
assert error_payload["id"] == "req-1"
assert error_payload["error"]["code"] == -32603
assert "upstream died" in error_payload["error"]["message"]
@pytest.mark.asyncio
async def test_send_message_pascal_case_routes_to_asend_message():
from litellm.proxy._types import UserAPIKeyAuth
@ -2029,3 +2387,83 @@ async def test_forward_jsonrpc_sse_is_untouched_while_keepalives_are_unconfigure
assert not any(chunk.startswith(":") for chunk in chunks)
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
async def _stream_message_response():
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
return await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
@pytest.mark.asyncio
async def test_handle_stream_message_pings_while_the_upstream_agent_is_still_silent(
monkeypatch,
):
"""message/stream is SSE like tasks/resubscribe, so a slow first event must be
held open by the same keepalives rather than sitting idle for the whole
time-to-first-token."""
import asyncio
import litellm
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", 0.05)
async def fake_stream(**kwargs):
await asyncio.sleep(0.3)
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _stream_message_response()
assert response.headers["x-accel-buffering"] == "no"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert chunks[0] == ": ping\n\n"
assert chunks.count(": ping\n\n") >= 3
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
@pytest.mark.asyncio
async def test_handle_stream_message_is_untouched_while_keepalives_are_unconfigured(
monkeypatch,
):
"""Off until an operator sets an interval, so the default stream is unchanged."""
import litellm
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", None)
async def fake_stream(**kwargs):
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
response = await _stream_message_response()
assert "x-accel-buffering" not in response.headers
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
assert not any(chunk.startswith(":") for chunk in chunks)
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"

View file

@ -304,6 +304,7 @@ async def test_proxy_streaming_serves_1_0_envelopes():
user_api_key_dict=user_api_key_dict,
)
assert response.media_type == "text/event-stream"
lines: List[Dict[str, Any]] = []
async for raw_line in response.body_iterator:
line = (
@ -312,7 +313,7 @@ async def test_proxy_streaming_serves_1_0_envelopes():
else str(raw_line).strip()
)
if line:
lines.append(json.loads(line))
lines.append(json.loads(line.removeprefix("data:").strip()))
assert lines, "expected at least one streamed JSON-RPC event"
message_events = [

View file

@ -1016,10 +1016,9 @@ class TestStreamingTransform:
@pytest.mark.asyncio
async def test_emit_streaming_http_error_a2a_yields_jsonrpc_chunk(self):
"""The shared streaming error helper emits an in-stream JSON-RPC error for
A2A call types instead of raising."""
import json
"""The shared streaming error helper emits an in-stream JSON-RPC error
object (not a pre-serialized string, which the A2A endpoint would frame as
a JSON string instead of an error object) for A2A call types."""
handler = UnifiedLLMGuardrails()
exc = unified_module.HTTPException(
status_code=400,
@ -1036,7 +1035,8 @@ class TestStreamingTransform:
emitted.append(item)
assert len(emitted) == 1
payload = json.loads(emitted[0])
payload = emitted[0]
assert isinstance(payload, dict)
assert payload["error"]["message"] == "stream_transform_underflow"
assert payload["id"] == "req-1"