diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 4a3f7d608ce..a5c4463da8d 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -155,14 +155,16 @@ class A2ACompletionBridgeHandler: verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider) provider_params: Final = {key: value for key, value in params.items() if key != "messages"} - return await a2a_provider_config.handle_non_streaming( - request_id=request_id, - params=provider_params, - api_base=api_base, - timeout=litellm_params.get("timeout") or 60.0, - litellm_params=litellm_params, - agent_extra_headers=agent_extra_headers, - ) + provider_kwargs: Final[dict[str, Any]] = { + "request_id": request_id, + "params": provider_params, + "api_base": api_base, + "litellm_params": litellm_params, + "agent_extra_headers": agent_extra_headers, + } + if litellm_params.get("timeout") is not None: + provider_kwargs["timeout"] = litellm_params["timeout"] + return await a2a_provider_config.handle_non_streaming(**provider_kwargs) completion_params: Final = A2ACompletionBridgeHandler._build_completion_params( params=params, @@ -226,14 +228,16 @@ class A2ACompletionBridgeHandler: verbose_logger.info("A2A: Using provider config for %s (streaming)", custom_llm_provider) provider_params: Final = {key: value for key, value in params.items() if key != "messages"} - async for chunk in a2a_provider_config.handle_streaming( - request_id=request_id, - params=provider_params, - api_base=api_base, - timeout=litellm_params.get("timeout") or 60.0, - litellm_params=litellm_params, - agent_extra_headers=agent_extra_headers, - ): + provider_kwargs: Final[dict[str, Any]] = { + "request_id": request_id, + "params": provider_params, + "api_base": api_base, + "litellm_params": litellm_params, + "agent_extra_headers": agent_extra_headers, + } + if litellm_params.get("timeout") is not None: + provider_kwargs["timeout"] = litellm_params["timeout"] + async for chunk in a2a_provider_config.handle_streaming(**provider_kwargs): yield chunk return @@ -268,8 +272,7 @@ class A2ACompletionBridgeHandler: # Call litellm.acompletion with streaming response: Final = await A2ACompletionBridgeHandler._acompletion(completion_params) - # 3. Accumulate content and emit artifact update - accumulated_text = "" + # 3. Forward content as artifact updates accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas chunk_count = 0 async for chunk in response: @@ -286,15 +289,11 @@ class A2ACompletionBridgeHandler: accumulated_tool_calls.extend(tool_calls) if content: - accumulated_text += content - - # Emit artifact update with accumulated content - if accumulated_text: - artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( - ctx=ctx, - text=accumulated_text, - ) - yield artifact_event + artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( + ctx=ctx, + text=content, + ) + yield artifact_event # 4. Emit final status update (kind: "status-update", status: "completed", final: true) completed_event: Final = A2ACompletionBridgeTransformation.create_status_update_event( diff --git a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py index c87b1367377..c24af243c83 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -166,41 +166,44 @@ class A2ACompletionBridgeTransformation: Returns: A2A SendMessageResponse dict """ - # Extract content from response - content = "" - if hasattr(response, "choices") and response.choices: - choice: Final = response.choices[0] - if hasattr(choice, "message") and choice.message: - content = choice.message.content or "" + serialized_choices: list[dict[str, Any]] = [] + raw_choices: Final = getattr(response, "choices", None) + if raw_choices: + for choice in raw_choices: + content: Final = ( + getattr(getattr(choice, "message", None), "content", None) or "" + ) + message: Final = { + "kind": "message", + "role": "agent", + "parts": [{"kind": "text", "text": content}], + "messageId": uuid4().hex, + } + raw_tool_calls = getattr(getattr(choice, "message", None), "tool_calls", None) + if raw_tool_calls: + message["tool_calls"] = [ + call.model_dump(exclude_none=True) + if hasattr(call, "model_dump") + else call.dict(exclude_none=True) + if hasattr(call, "dict") + else call + for call in raw_tool_calls + ] + finish_reason: Final = getattr(choice, "finish_reason", None) + if finish_reason: + message["finish_reason"] = finish_reason + serialized_choices.append({"index": len(serialized_choices), "message": message}) - tool_calls: list[Any] | None = None - finish_reason: str | None = None - if hasattr(response, "choices") and response.choices: - choice = response.choices[0] - finish_reason = getattr(choice, "finish_reason", None) - message = getattr(choice, "message", None) - raw_tool_calls = getattr(message, "tool_calls", None) - if raw_tool_calls: - tool_calls = [ - call.model_dump(exclude_none=True) - if hasattr(call, "model_dump") - else call.dict(exclude_none=True) - if hasattr(call, "dict") - else call - for call in raw_tool_calls - ] - - # Build A2A message - a2a_message: Final = { - "kind": "message", - "role": "agent", - "parts": [{"kind": "text", "text": content}], - "messageId": uuid4().hex, - } - if tool_calls: - a2a_message["tool_calls"] = tool_calls - if finish_reason: - a2a_message["finish_reason"] = finish_reason + a2a_message: Final = ( + serialized_choices[0]["message"] + if serialized_choices + else { + "kind": "message", + "role": "agent", + "parts": [{"kind": "text", "text": ""}], + "messageId": uuid4().hex, + } + ) usage: Final = getattr(response, "usage", None) @@ -212,8 +215,10 @@ class A2ACompletionBridgeTransformation: } if usage is not None: a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage + if len(serialized_choices) > 1: + a2a_response["choices"] = serialized_choices - verbose_logger.debug("OpenAI -> A2A transform: content_length=%s", len(content)) + verbose_logger.debug("OpenAI -> A2A transform: content_length=%s", len(a2a_message["parts"][0]["text"])) return a2a_response diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 9da95139c4e..58b3b396a0e 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -5,6 +5,7 @@ A2A Streaming Response Iterator from typing import Final from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator +from litellm.types.llms.openai import ChatCompletionToolCallChunk from litellm.types.utils import GenericStreamingChunk, ModelResponseStream from ..common_utils import A2AError, extract_text_from_a2a_response @@ -118,14 +119,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator): return None - def _get_tool_calls(self, chunk: dict) -> list[dict] | None: + def _get_tool_calls(self, chunk: dict) -> 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): - return tool_calls + if isinstance(tool_calls, list) and tool_calls: + first_tool_call: Final = tool_calls[0] + return first_tool_call if isinstance(first_tool_call, dict) else None message = result.get("message") - if isinstance(message, dict) and isinstance(message.get("tool_calls"), list): - return message["tool_calls"] + if isinstance(message, dict) and isinstance(message.get("tool_calls"), list) and message["tool_calls"]: + first_tool_call = message["tool_calls"][0] + return first_tool_call if isinstance(first_tool_call, dict) else None return None diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 69c3aebf4f6..5826442865d 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -213,14 +213,53 @@ async def _route_registered_provider( result_dict: Final = result if isinstance(result, Mapping) else {} nested_message: Final = result_dict.get("message") response_message: Final = nested_message if isinstance(nested_message, Mapping) else result_dict - tool_calls: Final = response_message.get("tool_calls") - normalized_tool_calls: Final = tool_calls if isinstance(tool_calls, list) else None - finish_reason: Final = response_message.get("finish_reason") - text: Final = extract_text_from_a2a_response(response) - model_response: Final = ModelResponse( - id=str(response.get("id") or request_id), - model=model_name, - choices=[ # mutable-ok: ModelResponse requires a choices list + response_choices: Final = response.get("choices") + choice_payloads: Final = ( + response_choices + if isinstance(response_choices, list) + else result_dict.get("choices") + ) + if isinstance(choice_payloads, list) and choice_payloads: + model_choices = [ + Choices( + finish_reason=( + choice.get("finish_reason") + if isinstance(choice, Mapping) and isinstance(choice.get("finish_reason"), str) + else choice.get("message", {}).get("finish_reason") + if isinstance(choice, Mapping) + and isinstance(choice.get("message"), Mapping) + and isinstance(choice.get("message", {}).get("finish_reason"), str) + else "stop" + ), + index=choice.get("index", choice_index) + if isinstance(choice, Mapping) and isinstance(choice.get("index", choice_index), int) + else choice_index, + message=Message( + content=extract_text_from_a2a_response( + {"result": choice.get("message", choice)} + if isinstance(choice, Mapping) + else {"result": {}} + ), + role="assistant", + tool_calls=( + choice.get("message", {}).get("tool_calls") + if isinstance(choice, Mapping) + and isinstance(choice.get("message"), Mapping) + and isinstance(choice.get("message", {}).get("tool_calls"), list) + else choice.get("tool_calls") + if isinstance(choice, Mapping) and isinstance(choice.get("tool_calls"), list) + else None + ), + ), + ) + for choice_index, choice in enumerate(choice_payloads) + ] + else: + tool_calls: Final = response_message.get("tool_calls") + normalized_tool_calls: Final = tool_calls if isinstance(tool_calls, list) else None + finish_reason: Final = response_message.get("finish_reason") + text: Final = extract_text_from_a2a_response(response) + model_choices = [ Choices( finish_reason=( finish_reason @@ -232,12 +271,16 @@ async def _route_registered_provider( index=0, message=Message(content=text, role="assistant", tool_calls=normalized_tool_calls), ) - ], + ] + model_response: Final = ModelResponse( + id=str(response.get("id") or request_id), + model=model_name, + choices=model_choices, ) raw_usage: Final = response.get("usage") usage: Final = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage if usage is not None: - setattr(model_response, "usage", usage) + model_response.usage = usage if isinstance(logging_obj, Logging): logging_obj.model_call_details["usage"] = usage @@ -477,7 +520,8 @@ async def route_a2a_agent_request( cardless_provider: Final = registered_provider == "watsonx_orchestrate" or ( registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model ) - if (not isinstance(agent_url, str) or not agent_url) and not cardless_provider: + has_configured_api_base: Final = isinstance(configured_api_base, str) and bool(configured_api_base) + if (not isinstance(agent_url, str) or not agent_url) and not has_configured_api_base and not cardless_provider: verbose_proxy_logger.error("[A2A] Agent '%s' has no URL configured", agent_name) route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False) 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 1b3e5f86020..23495d06629 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -222,8 +222,8 @@ async def test_handle_streaming_emits_proper_events(): ): events.append(event) - # Should have 4 events: task, working, artifact, completed - assert len(events) == 4 + # Should have 5 events: task, working, two artifacts, completed + assert len(events) == 5 # Event 1: task submitted assert events[0]["result"]["kind"] == "task" @@ -234,14 +234,18 @@ async def test_handle_streaming_emits_proper_events(): assert events[1]["result"]["status"]["state"] == "working" assert events[1]["result"]["final"] is False - # Event 3: artifact update with accumulated content + # Event 3: first artifact update assert events[2]["result"]["kind"] == "artifact-update" - assert events[2]["result"]["artifact"]["parts"][0]["text"] == "Hello world" + assert events[2]["result"]["artifact"]["parts"][0]["text"] == "Hello" - # Event 4: status completed - assert events[3]["result"]["kind"] == "status-update" - assert events[3]["result"]["status"]["state"] == "completed" - assert events[3]["result"]["final"] is True + # Event 4: second artifact update + assert events[3]["result"]["kind"] == "artifact-update" + assert events[3]["result"]["artifact"]["parts"][0]["text"] == " world" + + # Event 5: status completed + assert events[4]["result"]["kind"] == "status-update" + assert events[4]["result"]["status"]["state"] == "completed" + assert events[4]["result"]["final"] is True @pytest.mark.asyncio 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 f436fe27a57..dafa2839ace 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 @@ -44,7 +44,7 @@ async def test_async_iterator_preserves_tool_calls(): chunk = await iterator.__aiter__().__anext__() - assert chunk["tool_use"] == tool_calls + assert chunk["tool_use"] == tool_calls[0] assert chunk["finish_reason"] == "tool_calls" diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 8cbeaf6f6d1..766745425c3 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -194,6 +194,92 @@ async def test_route_a2a_cardless_bedrock_agentcore_uses_registered_model(): assert bridge.await_args.kwargs["api_base"] is None +@pytest.mark.asyncio +async def test_route_a2a_registered_provider_uses_configured_api_base_without_card_url(): + from litellm.types.agents import AgentResponse + + agent = AgentResponse( + agent_id="test-agent-id", + agent_name="test-agent", + agent_card_params={}, + litellm_params={ + "custom_llm_provider": "langflow", + "model": "flow", + "api_base": "https://flow.example.com", + }, + ) + bridge_response = { + "jsonrpc": "2.0", + "id": "request-id", + "result": {"kind": "message", "parts": [{"kind": "text", "text": "Hello back"}]}, + } + + with ( + patch( + "litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through", + AsyncMock(return_value=agent), + ), + patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + AsyncMock(return_value=True), + ), + patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming", + AsyncMock(return_value=bridge_response), + ) as bridge, + ): + call = await route_a2a_agent_request( + {"model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}]}, + "acompletion", + ) + await call + + assert bridge.await_args.kwargs["api_base"] == "https://flow.example.com" + + +@pytest.mark.asyncio +async def test_registered_provider_response_preserves_multiple_choices(): + from litellm.types.agents import AgentResponse + + agent = AgentResponse( + agent_id="test-agent-id", + agent_name="test-agent", + agent_card_params={"url": "http://agent.example.com"}, + litellm_params={"custom_llm_provider": "pydantic_ai_agents"}, + ) + bridge_response = { + "jsonrpc": "2.0", + "id": "request-id", + "choices": [ + {"index": 0, "message": {"parts": [{"kind": "text", "text": "first"}]}}, + {"index": 1, "message": {"parts": [{"kind": "text", "text": "second"}]}}, + ], + "result": {}, + } + + with ( + patch( + "litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through", + AsyncMock(return_value=agent), + ), + patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + AsyncMock(return_value=True), + ), + patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming", + AsyncMock(return_value=bridge_response), + ), + ): + call = await route_a2a_agent_request( + {"model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}]}, + "acompletion", + ) + response = await call + + assert [choice.message.content for choice in response.choices] == ["first", "second"] + + @pytest.mark.asyncio async def test_route_a2a_cardless_watsonx_orchestrate_uses_registered_model(): from litellm.types.agents import AgentResponse