From 7ce56c0f80131c0fe93b6e2884a55455946b425d Mon Sep 17 00:00:00 2001 From: Peter Boers Date: Thu, 3 Sep 2026 10:58:53 +0200 Subject: [PATCH] fix(a2a): detect cumulative snapshots across part-join whitespace Snapshot detection compared raw strings, but A2A joins a multi-part message's text parts with spaces. A server that streams deltas ("Hello", "world") and then sends the whole reply as one two-part artifact produces "Helloworld" accumulated against a "Hello world" snapshot, so the prefix test missed and the reply was emitted twice: "HelloworldHello world". Compare with whitespace removed, and map the match back to a raw offset so the emitted tail keeps the server's own spacing and newlines. Reported by Greptile on #39513. Co-Authored-By: Claude Opus 5 (1M context) --- litellm/llms/a2a/chat/streaming_iterator.py | 43 +++++++++++++++---- .../chat/test_a2a_chat_streaming_iterator.py | 30 +++++++++++-- .../proxy/test_route_a2a_models.py | 6 +-- 3 files changed, 63 insertions(+), 16 deletions(-) diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 7be27b2fa0f..f2b073efd73 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -2,6 +2,7 @@ A2A Streaming Response Iterator """ +from itertools import accumulate from typing import Final from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator @@ -10,6 +11,29 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream from ..common_utils import extract_text_from_a2a_response +def _ignoring_whitespace(text: str) -> str: + """ + Comparison key for snapshot detection. + + A2A joins a multi-part message's text parts with spaces, so the same content arrives + with different whitespace depending on how the server chunked it: two delta events + ("Hello", "world") accumulate to "Helloworld" while a single two-part snapshot of the + same content renders "Hello world". Comparing without whitespace makes them equal. + """ + return "".join(text.split()) + + +def _index_after(text: str, non_space_count: int) -> int: + """Index in `text` just past its first `non_space_count` non-whitespace characters.""" + if non_space_count <= 0: + return 0 + running: Final = accumulate(0 if char.isspace() else 1 for char in text) + return next( + (index + 1 for index, total in enumerate(running) if total >= non_space_count), + len(text), + ) + + class A2AModelResponseIterator(BaseModelResponseIterator): """ Iterator for parsing A2A streaming responses. @@ -101,18 +125,21 @@ class A2AModelResponseIterator(BaseModelResponseIterator): if not text: return "" - emitted: str = self._emitted_text - if emitted and text.startswith(emitted): - suffix: str = text[len(emitted) :] - if suffix: - self._emitted_text = text + emitted_key: Final = _ignoring_whitespace(self._emitted_text) + text_key: Final = _ignoring_whitespace(text) + + if emitted_key and text_key.startswith(emitted_key): + suffix: Final = text[_index_after(text, len(emitted_key)) :] + if suffix.strip(): + self._emitted_text += suffix return suffix - # text == emitted. Treat as a snapshot repeat, except while only a single delta - # has been emitted, where a genuinely repeated delta is still indistinguishable. + # Same content as everything emitted so far: a snapshot repeat, except while + # only a single delta has been emitted, where a genuinely repeated delta is + # still indistinguishable from one. if self._delta_count > 1: return "" - self._emitted_text = emitted + text + self._emitted_text += text self._delta_count += 1 return text diff --git a/tests/test_litellm/llms/a2a/chat/test_a2a_chat_streaming_iterator.py b/tests/test_litellm/llms/a2a/chat/test_a2a_chat_streaming_iterator.py index c49f989d76d..4becf808cd9 100644 --- a/tests/test_litellm/llms/a2a/chat/test_a2a_chat_streaming_iterator.py +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_chat_streaming_iterator.py @@ -30,7 +30,7 @@ def _status_update( } -def _artifact_update(*, text: str, append: bool = False) -> dict: +def _artifact_update(*texts: str, append: bool = False) -> dict: return { "jsonrpc": "2.0", "id": "1", @@ -38,7 +38,7 @@ def _artifact_update(*, text: str, append: bool = False) -> dict: "kind": "artifact-update", "append": append, "lastChunk": True, - "artifact": {"parts": [{"kind": "text", "text": text}]}, + "artifact": {"parts": [{"kind": "text", "text": text} for text in texts]}, }, } @@ -52,7 +52,7 @@ KAGENT_OK_STREAM = [ _status_update(text="O"), _status_update(text="K"), _status_update(text="OK"), - _artifact_update(text="OK"), + _artifact_update("OK"), _status_update(state="completed", final=True), ] @@ -82,6 +82,12 @@ def test_kagent_stream_finishes_on_completed_state(): pytest.param(["O", "K", "OK"], "OK", id="deltas_then_final_snapshot"), pytest.param(["O", "K", "OK", "OK"], "OK", id="deltas_then_repeated_snapshots"), pytest.param(["a", "a", "a"], "aaa", id="genuinely_repeated_deltas"), + pytest.param(["Hello", "world", "Hello world"], "Helloworld", id="multipart_snapshot_respaced"), + pytest.param( + ["Hello", "world", "Hello world again"], + "Helloworld again", + id="multipart_snapshot_extends", + ), pytest.param(["", "OK", ""], "OK", id="empty_events_ignored"), ], ) @@ -89,3 +95,21 @@ def test_incremental_text_reduction(texts, expected): """Delta-style and snapshot-style servers must collapse to the same output.""" iterator = _iterator() assert "".join(iterator._to_incremental_text(t) for t in texts) == expected + + +# A server that chunks its reply into separate delta events but sends the final artifact +# as one multi-part message: A2A joins those parts with a space, so the snapshot reads +# "Hello world" while the deltas accumulated to "Helloworld". +MULTIPART_SNAPSHOT_STREAM = [ + _status_update(text="Hello"), + _status_update(text="world"), + _artifact_update("Hello", "world"), + _status_update(state="completed", final=True), +] + + +def test_multipart_snapshot_is_not_re_emitted(): + """Regression: whitespace introduced by part-joining must not defeat snapshot detection.""" + iterator = _iterator() + rendered = "".join(iterator.chunk_parser(e)["text"] for e in MULTIPART_SNAPSHOT_STREAM) + assert rendered == "Helloworld" diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 6dc69e566d8..d1e85d56817 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -4,8 +4,6 @@ Test A2A model routing in proxy. Maps to: litellm/proxy/agent_endpoints/a2a_routing.py """ - - from unittest.mock import AsyncMock, Mock, patch import pytest @@ -214,9 +212,7 @@ def _router_without_a2a_deployment( id="team_scoped_key", ), pytest.param({"patterns": ("openrouter/*",)}, {}, id="wildcard_model_group"), - pytest.param( - {"default_deployment": {"model_name": "*"}}, {}, id="default_deployment" - ), + pytest.param({"default_deployment": {"model_name": "*"}}, {}, id="default_deployment"), ], ) async def test_a2a_model_resolves_before_router_branches(router_kwargs, extra_data):