mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: preserve A2A stream metadata
This commit is contained in:
parent
c1add95b62
commit
a52bf355c6
5 changed files with 170 additions and 4 deletions
|
|
@ -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)
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset(
|
|||
"safety_identifier",
|
||||
"stop",
|
||||
"store",
|
||||
"stream_options",
|
||||
"temperature",
|
||||
"thinking",
|
||||
"timeout",
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue