mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
d40aea865d
commit
3207014906
5 changed files with 566 additions and 153 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue