mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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) <noreply@anthropic.com>
This commit is contained in:
parent
d88bdaa17c
commit
f757c52570
3 changed files with 143 additions and 2 deletions
|
|
@ -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", {})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue