From 3207014906a18ba25c30beba30febf8e14fcd252 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 20 Aug 2026 16:19:58 -0700 Subject: [PATCH] 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 --- .../proxy/agent_endpoints/a2a_endpoints.py | 181 +++---- .../unified_guardrail/unified_guardrail.py | 85 ++-- .../agent_endpoints/test_a2a_endpoints.py | 440 +++++++++++++++++- .../agent_endpoints/test_a2a_version_e2e.py | 3 +- .../test_unified_guardrail.py | 10 +- 5 files changed, 566 insertions(+), 153 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index cea30ffad52..bd02cfdf907 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -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: \\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: \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( diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 5bbb01c6c8e..e95e97bfe74 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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 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 e54358c1f00..2ff38af80b1 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -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: \\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: \\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" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py index 069c72af53a..3dc4d3427cd 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py @@ -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 = [ diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index 8a551f749d0..8b9ecfbbeee 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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"