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