diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 71993135037..8d84c65ab4e 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -280,6 +280,9 @@ class A2ACompletionBridgeHandler: # 3. Forward content as artifact updates 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_finish_reasons: dict[int, str] = {} stream_usage: object | None = None stream_finish_reason: str | None = None chunk_count = 0 @@ -304,24 +307,34 @@ class A2ACompletionBridgeHandler: stream_usage = dumped_usage # Extract delta content - content = "" - if chunk is not None and hasattr(chunk, "choices") and chunk.choices: - choice = chunk.choices[0] - raw_finish_reason = getattr(choice, "finish_reason", None) - if isinstance(raw_finish_reason, str) and raw_finish_reason: - stream_finish_reason = raw_finish_reason - if hasattr(choice, "delta") and choice.delta: - content = choice.delta.content or "" - tool_calls = getattr(choice.delta, "tool_calls", None) - if isinstance(tool_calls, (list, tuple)): - accumulated_tool_calls.extend(tool_calls) + choices = getattr(chunk, "choices", None) if chunk is not None else None + if isinstance(choices, (list, tuple)): + for choice_position, choice in enumerate(choices): + raw_index = getattr(choice, "index", choice_position) + choice_index = raw_index if isinstance(raw_index, int) else choice_position + choice_texts.setdefault(choice_index, "") + raw_finish_reason = getattr(choice, "finish_reason", None) + if isinstance(raw_finish_reason, str) and raw_finish_reason: + choice_finish_reasons[choice_index] = raw_finish_reason + if choice_index == 0 or stream_finish_reason is None: + stream_finish_reason = raw_finish_reason + content = "" + delta = getattr(choice, "delta", None) + if delta: + raw_content = getattr(delta, "content", None) + content = raw_content if isinstance(raw_content, str) else "" + choice_texts[choice_index] += content + tool_calls = getattr(delta, "tool_calls", None) + if isinstance(tool_calls, (list, tuple)): + accumulated_tool_calls.extend(tool_calls) + choice_tool_calls.setdefault(choice_index, []).extend(tool_calls) - if content: - artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( - ctx=ctx, - text=content, - ) - yield artifact_event + if content: + artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( + ctx=ctx, + text=content, + ) + yield artifact_event finally: close_response = getattr(response, "aclose", None) if close_response is not None: @@ -339,6 +352,28 @@ class A2ACompletionBridgeHandler: completed_event["result"]["finish_reason"] = stream_finish_reason if stream_usage is not None: completed_event["usage"] = stream_usage + if len(choice_texts) > 1: + completed_event["result"]["choices"] = [ + { + "index": choice_index, + "message": { + "kind": "message", + "role": "agent", + "parts": [{"kind": "text", "text": choice_texts[choice_index]}], + **( + {"tool_calls": choice_tool_calls[choice_index]} + if choice_tool_calls.get(choice_index) + else {} + ), + }, + **( + {"finish_reason": choice_finish_reasons[choice_index]} + if choice_index in choice_finish_reasons + else {} + ), + } + for choice_index in sorted(choice_texts) + ] yield completed_event verbose_logger.info( diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index f6340426c1b..1e416a21e84 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1233,7 +1233,8 @@ class CustomStreamWrapper: ) if "tool_use" in anthropic_response_obj and anthropic_response_obj["tool_use"] is not None: - completion_obj["tool_calls"] = [anthropic_response_obj["tool_use"]] + tool_use = anthropic_response_obj["tool_use"] + completion_obj["tool_calls"] = tool_use if isinstance(tool_use, list) else [tool_use] if ( "provider_specific_fields" in anthropic_response_obj @@ -2559,6 +2560,8 @@ def convert_generic_chunk_to_model_response_stream( ) -> ModelResponseStream: from litellm.types.utils import Delta + tool_use = chunk.get("tool_use", None) + tool_calls = tool_use if isinstance(tool_use, list) else [tool_use] if tool_use is not None else None model_response_stream: Final = ModelResponseStream( id=str(uuid.uuid4()), model="", @@ -2567,7 +2570,7 @@ def convert_generic_chunk_to_model_response_stream( index=chunk.get("index", 0), delta=Delta( content=chunk["text"], - tool_calls=chunk.get("tool_use", None), + tool_calls=tool_calls, ), ) ], diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 636aef0e724..32f2d0babc2 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -149,16 +149,18 @@ class A2AModelResponseIterator(BaseModelResponseIterator): return raw_usage return raw_usage - def _get_tool_calls(self, chunk: dict) -> 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 tool_calls = result.get("tool_calls") if isinstance(tool_calls, list) and tool_calls: - return self._serialize_tool_call(tool_calls[0]) + return self._serialize_tool_calls(tool_calls) message = result.get("message") if isinstance(message, dict) and isinstance(message.get("tool_calls"), list) and message["tool_calls"]: - return self._serialize_tool_call(message["tool_calls"][0]) + return self._serialize_tool_calls(message["tool_calls"]) return None @staticmethod @@ -171,6 +173,19 @@ class A2AModelResponseIterator(BaseModelResponseIterator): return tool_call.dict(exclude_none=True) return None + @classmethod + def _serialize_tool_calls( + cls, tool_calls: list[object] + ) -> ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None: + serialized: Final = [ + tool_call_value + for tool_call in tool_calls + if (tool_call_value := cls._serialize_tool_call(tool_call)) is not None + ] + if len(serialized) == 1: + return serialized[0] + return serialized or None + async def aclose(self) -> None: streaming_response = self.streaming_response self.streaming_response = None diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 34e945f1415..d6dcd58c02b 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -163,6 +163,13 @@ async def _route_registered_provider( logging_obj: Final = data.get("litellm_logging_obj") if isinstance(logging_obj, Logging): provider_params["no-log"] = True + provider_model: Final = litellm_params.get("model") + if isinstance(provider_model, str): + logging_obj.model_call_details["model"] = provider_model + logging_obj.model_call_details.setdefault("litellm_params", {})["model"] = provider_model + provider_name: Final = litellm_params.get("custom_llm_provider") + if isinstance(provider_name, str): + logging_obj.model_call_details["custom_llm_provider"] = provider_name pricing_params = { key: litellm_params[key] for key in _A2A_PRICING_PARAMS diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 16a10a16e70..5927214a202 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1841,11 +1841,6 @@ class ProxyBaseLLMRequestProcessing: merge_a2a_agent_guardrails_before_hooks, ) - self.data = await authorize_a2a_agent_before_hooks( - data=self.data, - user_api_key_dict=user_api_key_dict, - ) - logging_obj, self.data = litellm.utils.function_setup( original_function=route_type, rules_obj=litellm.utils.Rules(), @@ -1855,6 +1850,11 @@ class ProxyBaseLLMRequestProcessing: self.data["litellm_logging_obj"] = logging_obj + self.data = await authorize_a2a_agent_before_hooks( + data=self.data, + user_api_key_dict=user_api_key_dict, + ) + self.data = await merge_a2a_agent_guardrails_before_hooks(self.data) # Merge model-level guardrails before pre_call_hook so DB/UI-configured diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 67eae2b4f21..61d7aca0430 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -317,7 +317,7 @@ class ModelInfo(ModelInfoBase, total=False): class GenericStreamingChunk(TypedDict, total=False): text: Required[str] - tool_use: ChatCompletionToolCallChunk | None + tool_use: ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None is_finished: Required[bool] finish_reason: Required[str] usage: Required[ChatCompletionUsageBlock | None] 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 d2b74d416a0..6279bc78dab 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -249,6 +249,44 @@ async def test_handle_streaming_emits_proper_events(): assert events[4]["usage"]["total_tokens"] == 5 +@pytest.mark.asyncio +async def test_handle_streaming_preserves_multiple_choices(): + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + mock_chunk = MagicMock() + first_choice = MagicMock() + first_choice.index = 0 + first_choice.finish_reason = None + first_choice.delta.content = "first" + second_choice = MagicMock() + second_choice.index = 1 + second_choice.finish_reason = "length" + second_choice.delta.content = "second" + mock_chunk.choices = [first_choice, second_choice] + + async def mock_streaming_response(): + yield mock_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-choices", + params={"message": {"role": "user", "parts": []}}, + litellm_params={"custom_llm_provider": "langgraph", "model": "agent", "n": 2}, + ) + ] + + choices = events[-1]["result"]["choices"] + assert [choice["index"] for choice in choices] == [0, 1] + assert choices[0]["message"]["parts"][0]["text"] == "first" + assert choices[1]["message"]["parts"][0]["text"] == "second" + assert choices[1]["finish_reason"] == "length" + + @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 db303719864..443254e69fd 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 @@ -49,6 +49,31 @@ async def test_async_iterator_preserves_tool_calls(): assert chunk["finish_reason"] == "tool_calls" +@pytest.mark.asyncio +async def test_async_iterator_preserves_parallel_tool_calls(): + tool_calls = [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + }, + { + "id": "call-2", + "type": "function", + "function": {"name": "write", "arguments": "{}"}, + }, + ] + + async def _events(): + yield {"jsonrpc": "2.0", "result": {"tool_calls": tool_calls}} + + iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False) + + chunk = await iterator.__aiter__().__anext__() + + assert chunk["tool_use"] == tool_calls + + @pytest.mark.asyncio async def test_async_iterator_serializes_delta_tool_calls_and_usage(): delta = Delta( diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index ca8a7ff5b2a..2a34f4404bc 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -11,6 +11,7 @@ from fastapi import HTTPException from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.agent_endpoints.a2a_routing import ( + _route_registered_provider, merge_a2a_agent_guardrails_before_hooks, route_a2a_agent_request, ) @@ -642,6 +643,39 @@ async def test_route_a2a_stream_uses_registered_provider(): assert response is wrapper +@pytest.mark.asyncio +async def test_registered_provider_logging_uses_provider_model_for_builtin_pricing(): + class FakeLogging: + def __init__(self) -> None: + self.model_call_details = {"litellm_params": {}} + self.litellm_params = self.model_call_details["litellm_params"] + self.custom_pricing = False + + logging_obj = FakeLogging() + response = {"result": {"message": {"parts": [{"kind": "text", "text": "hello"}]}}} + with ( + patch("litellm.litellm_core_utils.litellm_logging.Logging", FakeLogging), + patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming", + AsyncMock(return_value=response), + ), + ): + await _route_registered_provider( + data={ + "messages": [{"role": "user", "content": "hello"}], + "litellm_logging_obj": logging_obj, + }, + model_name="a2a/agent", + api_base="https://provider.example", + litellm_params={"model": "gpt-4o", "custom_llm_provider": "openai"}, + static_headers=None, + ) + + assert logging_obj.model_call_details["model"] == "gpt-4o" + assert logging_obj.model_call_details["custom_llm_provider"] == "openai" + assert logging_obj.model_call_details["litellm_params"]["model"] == "gpt-4o" + + @pytest.mark.asyncio async def test_route_non_a2a_model_raises_error_if_not_in_router(): """Test that non-a2a models that aren't in router raise an error"""