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) <noreply@anthropic.com>
This commit is contained in:
Peter Boers 2026-09-03 10:58:53 +02:00
parent a1dbd1f28c
commit 7ce56c0f80
No known key found for this signature in database
3 changed files with 63 additions and 16 deletions

View file

@ -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

View file

@ -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"

View file

@ -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):