From f757c525703e1d3db5568781d4c94f46a3a80a0a Mon Sep 17 00:00:00 2001 From: Peter Boers Date: Thu, 3 Sep 2026 09:43:07 +0200 Subject: [PATCH] fix(a2a): stop echoing the prompt and repeating the reply when streaming `extract_text_from_a2a_response()` took text from every streaming event regardless of who authored it or whether it was a delta, so a reply of "OK" reached the client as "user: Reply with exactly: OKOKOKOK": - the opening `status-update` (`state: "submitted"`) carries the caller's own message, and A2A marks it `role: "user"`. It was emitted as assistant output. Now skipped: per spec `Message.role` is "user" or "agent", and only agent text belongs in a completion. - the terminal non-partial `status-update` and the `artifact-update` (`append: false`) each repeat the whole reply. Both were forwarded as deltas. The iterator now tracks what it has emitted and forwards only new text, which also lets servers that stream growing snapshots ("O", "OK") collapse to the same output as servers that stream true deltas ("O", "K"). Tests use an event sequence captured from a real kagent agent. Co-Authored-By: Claude Opus 5 (1M context) --- litellm/llms/a2a/chat/streaming_iterator.py | 37 +++++++- litellm/llms/a2a/common_utils.py | 17 ++++ .../chat/test_a2a_chat_streaming_iterator.py | 91 +++++++++++++++++++ 3 files changed, 143 insertions(+), 2 deletions(-) create mode 100644 tests/test_litellm/llms/a2a/chat/test_a2a_chat_streaming_iterator.py diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 1983c18a6b3..7be27b2fa0f 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -30,6 +30,9 @@ class A2AModelResponseIterator(BaseModelResponseIterator): json_mode=json_mode, ) self.model = model + # Text already emitted downstream, used to collapse cumulative snapshots. + self._emitted_text: str = "" + self._delta_count: int = 0 def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream: """ @@ -57,8 +60,8 @@ class A2AModelResponseIterator(BaseModelResponseIterator): } """ try: - # Extract text from A2A response - text: Final = extract_text_from_a2a_response(chunk) + # Extract text from A2A response, then reduce it to what is actually new. + text: Final = self._to_incremental_text(extract_text_from_a2a_response(chunk)) # Determine finish reason finish_reason: Final = self._get_finish_reason(chunk) @@ -83,6 +86,36 @@ class A2AModelResponseIterator(BaseModelResponseIterator): tool_use=None, ) + def _to_incremental_text(self, text: str) -> str: + """ + Reduce an A2A event's text to the portion not yet emitted. + + A2A servers interleave true deltas with cumulative snapshots of the whole reply: a + terminal non-partial ``status-update`` and an ``artifact-update`` carrying + ``append: false`` both repeat everything produced so far. Forwarding those verbatim + makes the client render the reply two or three times over, so emit only new text. + + Handles both streaming styles: servers that send deltas ("O", "K") and servers that + send growing snapshots ("O", "OK") collapse to the same output. + """ + 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 + 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. + if self._delta_count > 1: + return "" + + self._emitted_text = emitted + text + self._delta_count += 1 + return text + def _get_finish_reason(self, chunk: dict) -> str | None: """Extract finish reason from A2A chunk""" result: Final = chunk.get("result", {}) diff --git a/litellm/llms/a2a/common_utils.py b/litellm/llms/a2a/common_utils.py index 178b4c47a0f..d4a437027e2 100644 --- a/litellm/llms/a2a/common_utils.py +++ b/litellm/llms/a2a/common_utils.py @@ -61,6 +61,17 @@ def convert_messages_to_prompt(messages: list[AllMessageValues]) -> str: return "\n".join(conversation_parts) +def _is_user_authored(message: object) -> bool: + """ + Whether an A2A message was authored by the caller rather than the agent. + + A2A ``Message.role`` is either ``"user"`` or ``"agent"``. Agents commonly echo the + inbound message back on the first ``status-update`` (``state: "submitted"``), and that + text must never surface as assistant output in a chat completion. + """ + return isinstance(message, dict) and message.get("role") == "user" + + def extract_text_from_a2a_message(message: dict[str, Any], depth: int = 0, max_depth: int = 10) -> str: """ Extract text content from A2A message parts. @@ -115,11 +126,15 @@ def extract_text_from_a2a_response(response_dict: dict[str, Any], max_depth: int # Check if result itself has parts (direct message) if "parts" in result: + if _is_user_authored(result): + return "" return extract_text_from_a2a_message(result, depth=0, max_depth=max_depth) # Check for nested message message: Final = result.get("message") if message: + if _is_user_authored(message): + return "" return extract_text_from_a2a_message(message, depth=0, max_depth=max_depth) # Check for streaming artifact-update (singular artifact) @@ -132,6 +147,8 @@ def extract_text_from_a2a_response(response_dict: dict[str, Any], max_depth: int if isinstance(status, dict): status_message: Final = status.get("message") if status_message: + if _is_user_authored(status_message): + return "" return extract_text_from_a2a_message(status_message, depth=0, max_depth=max_depth) # Handle task result with artifacts (plural, array) 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 new file mode 100644 index 00000000000..c49f989d76d --- /dev/null +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_chat_streaming_iterator.py @@ -0,0 +1,91 @@ +"""Tests for litellm/llms/a2a/chat/streaming_iterator.py delta handling.""" + +import pytest + +from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator + + +def _iterator() -> A2AModelResponseIterator: + return A2AModelResponseIterator(streaming_response=iter(()), sync_stream=True) + + +def _status_update( + *, + text: str | None = None, + role: str = "agent", + state: str = "working", + final: bool = False, +) -> dict: + status: dict = {"state": state} + if text is not None: + status["message"] = { + "kind": "message", + "role": role, + "parts": [{"kind": "text", "text": text}], + } + return { + "jsonrpc": "2.0", + "id": "1", + "result": {"kind": "status-update", "final": final, "status": status}, + } + + +def _artifact_update(*, text: str, append: bool = False) -> dict: + return { + "jsonrpc": "2.0", + "id": "1", + "result": { + "kind": "artifact-update", + "append": append, + "lastChunk": True, + "artifact": {"parts": [{"kind": "text", "text": text}]}, + }, + } + + +# Event sequence captured from a real kagent A2A agent replying "OK": the submitted +# status-update echoes the caller's own message, then two true deltas are followed by two +# cumulative snapshots of the whole reply. +KAGENT_OK_STREAM = [ + _status_update(text="Reply with exactly: OK", role="user", state="submitted"), + _status_update(), + _status_update(text="O"), + _status_update(text="K"), + _status_update(text="OK"), + _artifact_update(text="OK"), + _status_update(state="completed", final=True), +] + + +def test_kagent_stream_yields_reply_exactly_once(): + """ + Regression: the caller's echoed message must not be emitted as assistant output, and + cumulative snapshots must not repeat the reply. + + Previously this stream rendered as "user: Reply with exactly: OKOKOKOK". + """ + iterator = _iterator() + assert "".join(iterator.chunk_parser(e)["text"] for e in KAGENT_OK_STREAM) == "OK" + + +def test_kagent_stream_finishes_on_completed_state(): + iterator = _iterator() + chunks = [iterator.chunk_parser(e) for e in KAGENT_OK_STREAM] + assert [c["finish_reason"] for c in chunks if c["is_finished"]] == ["stop"] + + +@pytest.mark.parametrize( + "texts, expected", + [ + pytest.param(["Hello", " world"], "Hello world", id="incremental_deltas"), + pytest.param(["O", "OK", "OKAY"], "OKAY", id="cumulative_snapshots"), + 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(["", "OK", ""], "OK", id="empty_events_ignored"), + ], +) +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