fix: preserve A2A stream metadata

This commit is contained in:
aiedwardyi 2026-08-25 21:05:09 +09:00
parent c1add95b62
commit a52bf355c6
No known key found for this signature in database
5 changed files with 170 additions and 4 deletions

View file

@ -282,6 +282,8 @@ class A2ACompletionBridgeHandler:
accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas
choice_texts: dict[int, str] = {}
choice_tool_calls: dict[int, list[object]] = {}
choice_delta_fields: dict[int, dict[str, object]] = {}
choice_logprobs: dict[int, object] = {}
choice_finish_reasons: dict[int, str] = {}
stream_usage: object | None = None
stream_finish_reason: str | None = None
@ -328,6 +330,26 @@ class A2ACompletionBridgeHandler:
if isinstance(tool_calls, (list, tuple)):
accumulated_tool_calls.extend(tool_calls)
choice_tool_calls.setdefault(choice_index, []).extend(tool_calls)
delta_fields = A2ACompletionBridgeTransformation._model_dump(delta)
if delta_fields:
choice_fields = choice_delta_fields.setdefault(choice_index, {})
for field, value in delta_fields.items():
if field in {"content", "role", "tool_calls"} or value is None:
continue
previous = choice_fields.get(field)
if (isinstance(previous, str) and isinstance(value, str)) or (
isinstance(previous, list) and isinstance(value, list)
):
choice_fields[field] = previous + value
elif isinstance(previous, Mapping) and isinstance(value, Mapping):
choice_fields[field] = {**previous, **value}
else:
choice_fields[field] = value
raw_logprobs = getattr(choice, "logprobs", None)
serialized_logprobs = A2ACompletionBridgeTransformation._model_dump(raw_logprobs)
if serialized_logprobs:
choice_logprobs[choice_index] = serialized_logprobs
if content:
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
@ -352,6 +374,10 @@ class A2ACompletionBridgeHandler:
completed_event["result"]["finish_reason"] = stream_finish_reason
if stream_usage is not None:
completed_event["usage"] = stream_usage
if choice_delta_fields.get(0):
completed_event["result"].update(choice_delta_fields[0])
if 0 in choice_logprobs:
completed_event["result"]["logprobs"] = choice_logprobs[0]
if len(choice_texts) > 1:
completed_event["result"]["choices"] = [
{
@ -365,12 +391,14 @@ class A2ACompletionBridgeHandler:
if choice_tool_calls.get(choice_index)
else {}
),
**choice_delta_fields.get(choice_index, {}),
},
**(
{"finish_reason": choice_finish_reasons[choice_index]}
if choice_index in choice_finish_reasons
else {}
),
**({"logprobs": choice_logprobs[choice_index]} if choice_index in choice_logprobs else {}),
}
for choice_index in sorted(choice_texts)
]

View file

@ -70,7 +70,56 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
try:
# Extract text from A2A response
text: Final = extract_text_from_a2a_response(chunk)
result: Final = chunk.get("result", {})
status: Final = result.get("status", {}) if isinstance(result, Mapping) else {}
is_working_status: Final = (
isinstance(result, Mapping)
and result.get("kind") == "status-update"
and isinstance(status, Mapping)
and status.get("state") == "working"
)
text: Final = "" if is_working_status else extract_text_from_a2a_response(chunk)
provider_fields: dict[str, object] = {}
if isinstance(result, Mapping) and not is_working_status:
control_fields = {
"artifacts",
"choices",
"contextId",
"final",
"finish_reason",
"history",
"id",
"kind",
"message",
"parts",
"status",
"taskId",
"tool_calls",
"usage",
}
provider_fields.update(
{key: value for key, value in result.items() if key not in control_fields and value is not None}
)
choices = result.get("choices")
if isinstance(choices, list) and choices:
first_choice = choices[0]
if isinstance(first_choice, Mapping):
provider_fields.update(
{
key: value
for key, value in first_choice.items()
if key not in {"index", "message", "finish_reason"} and value is not None
}
)
first_message = first_choice.get("message")
if isinstance(first_message, Mapping):
provider_fields.update(
{
key: value
for key, value in first_message.items()
if key not in {"kind", "role", "parts", "tool_calls"} and value is not None
}
)
# Determine finish reason
finish_reason: Final = self._get_finish_reason(chunk)
@ -85,6 +134,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
usage=usage,
index=0,
tool_use=tool_calls,
provider_specific_fields=provider_fields or None,
)
except Exception:
# Return empty chunk on parse error
@ -149,9 +199,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
return raw_usage
return raw_usage
def _get_tool_calls(
self, chunk: dict
) -> ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None:
def _get_tool_calls(self, chunk: dict) -> ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None:
result: Final = chunk.get("result", {})
if not isinstance(result, dict):
return None

View file

@ -55,6 +55,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset(
"safety_identifier",
"stop",
"store",
"stream_options",
"temperature",
"thinking",
"timeout",

View file

@ -287,6 +287,49 @@ async def test_handle_streaming_preserves_multiple_choices():
assert choices[1]["finish_reason"] == "length"
@pytest.mark.asyncio
async def test_handle_streaming_preserves_non_text_delta_fields():
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
delta = MagicMock()
delta.content = ""
delta.tool_calls = None
delta.model_dump.return_value = {
"audio": {"data": "abc"},
"reasoning_content": "thinking",
"provider_specific_fields": {"trace_id": "trace-1"},
}
choice = MagicMock()
choice.index = 0
choice.finish_reason = "stop"
choice.delta = delta
choice.logprobs = {"content": []}
chunk = MagicMock()
chunk.choices = [choice]
async def mock_streaming_response():
yield chunk
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
mock_acompletion.return_value = mock_streaming_response()
events = [
event
async for event in A2ACompletionBridgeHandler.handle_streaming(
request_id="req-fields",
params={"message": {"role": "user", "parts": []}},
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
)
]
result = events[-1]["result"]
assert result["audio"] == {"data": "abc"}
assert result["reasoning_content"] == "thinking"
assert result["provider_specific_fields"] == {"trace_id": "trace-1"}
assert result["logprobs"] == {"content": []}
@pytest.mark.asyncio
async def test_provider_config_receives_full_message_history():
from litellm.a2a_protocol.litellm_completion_bridge.handler import (

View file

@ -28,6 +28,52 @@ async def test_async_iterator_accepts_decoded_a2a_events():
assert chunk["text"] == "Hello"
@pytest.mark.asyncio
async def test_async_iterator_ignores_status_message_text():
async def _events():
yield {
"jsonrpc": "2.0",
"result": {
"kind": "status-update",
"status": {
"state": "working",
"message": {"parts": [{"kind": "text", "text": "Processing request..."}]},
},
},
}
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
chunk = await iterator.__aiter__().__anext__()
assert chunk["text"] == ""
@pytest.mark.asyncio
async def test_async_iterator_preserves_non_text_fields():
async def _events():
yield {
"jsonrpc": "2.0",
"result": {
"kind": "status-update",
"status": {"state": "completed"},
"audio": {"data": "abc"},
"reasoning_content": "thinking",
"logprobs": {"content": []},
},
}
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
chunk = await iterator.__aiter__().__anext__()
assert chunk["provider_specific_fields"] == {
"audio": {"data": "abc"},
"reasoning_content": "thinking",
"logprobs": {"content": []},
}
@pytest.mark.asyncio
async def test_async_iterator_preserves_tool_calls():
tool_calls = [