mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
a1dbd1f28c
commit
7ce56c0f80
3 changed files with 63 additions and 16 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue