diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 8d84c65ab4e..f99a91729d4 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -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) ] diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 32f2d0babc2..1a598bb7a50 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index d6dcd58c02b..38cd5e9f4f2 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -55,6 +55,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset( "safety_identifier", "stop", "store", + "stream_options", "temperature", "thinking", "timeout", diff --git a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py index 6279bc78dab..11f2e75107d 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -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 ( diff --git a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py index 443254e69fd..9c4212ee135 100644 --- a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py @@ -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 = [