diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index ffc94a15890..3546ae0891d 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -174,6 +174,12 @@ class A2ACompletionBridgeHandler: verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider) provider_params: Final = dict(params) + if custom_llm_provider == "pydantic_ai_agents" or ( + custom_llm_provider == "bedrock" + and isinstance(litellm_params.get("model"), str) + and "agentcore" in litellm_params["model"] + ): + provider_params.pop("messages", None) provider_kwargs: Final[dict[str, Any]] = { "request_id": request_id, "params": provider_params, @@ -250,6 +256,12 @@ class A2ACompletionBridgeHandler: verbose_logger.info("A2A: Using provider config for %s (streaming)", custom_llm_provider) provider_params: Final = dict(params) + if custom_llm_provider == "pydantic_ai_agents" or ( + custom_llm_provider == "bedrock" + and isinstance(litellm_params.get("model"), str) + and "agentcore" in litellm_params["model"] + ): + provider_params.pop("messages", None) provider_kwargs: Final[dict[str, Any]] = { "request_id": request_id, "params": provider_params, @@ -429,7 +441,7 @@ class A2ACompletionBridgeHandler: "message": { "kind": "message", "role": "agent", - "parts": [{"kind": "text", "text": choice_texts[choice_index]}], + "parts": [{"kind": "text", "text": ""}], **( {"tool_calls": choice_tool_calls[choice_index]} if choice_tool_calls.get(choice_index) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index b1af6121839..5f5b5fc7705 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -147,7 +147,16 @@ async def _route_registered_provider( **_OBJECT_DICT_ADAPTER.validate_python(litellm_params), **{key: data[key] for key in _FORWARDED_REQUEST_PARAMS if key in data and data[key] is not None}, } - bridge_params: Final = _OBJECT_DICT_ADAPTER.validate_python(params) + registered_provider: Final = litellm_params.get("custom_llm_provider") + registered_model: Final = litellm_params.get("model") + native_provider: Final = registered_provider == "pydantic_ai_agents" or ( + registered_provider == "bedrock" + and isinstance(registered_model, str) + and "agentcore" in registered_model + ) + bridge_params: Final = _OBJECT_DICT_ADAPTER.validate_python( + {"message": params["message"]} if native_provider else params + ) configured_headers: Final = litellm_params.get("extra_headers") or litellm_params.get("headers") configured_headers_dict: Final = ( _HEADERS_ADAPTER.validate_python(configured_headers) if isinstance(configured_headers, dict) else None @@ -550,10 +559,7 @@ async def route_a2a_agent_request( ) configured_api_base: Final = registered_params_value.get("api_base") if registered_params_value else None api_base: Final = configured_api_base if isinstance(configured_api_base, str) and configured_api_base else agent_url - registered_model: Final = registered_params_value.get("model") if registered_params_value else None - cardless_provider: Final = registered_provider == "watsonx_orchestrate" or ( - registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model - ) + cardless_provider: Final = registered_provider is not None and registered_provider != "a2a" 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) 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 acd584b36db..7fe0a6ced75 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -352,8 +352,7 @@ async def test_handle_streaming_preserves_multiple_choices(): 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 [choice["message"]["parts"][0]["text"] for choice in choices] == ["", ""] assert choices[1]["finish_reason"] == "length" @@ -433,6 +432,34 @@ async def test_provider_config_receives_full_message_history(): assert provider_config.handle_non_streaming.await_args.kwargs["params"]["messages"] == messages +@pytest.mark.asyncio +async def test_native_provider_config_drops_internal_message_history(): + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + provider_config = MagicMock() + provider_config.handle_non_streaming = AsyncMock(return_value={"result": {}}) + params = { + "message": {"role": "user", "parts": []}, + "messages": [{"role": "user", "content": "Hello"}], + } + + with patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config", + return_value=provider_config, + ): + await A2ACompletionBridgeHandler.handle_non_streaming( + request_id="req-native", + params=params, + litellm_params={"custom_llm_provider": "pydantic_ai_agents", "model": "agent"}, + ) + + assert provider_config.handle_non_streaming.await_args.kwargs["params"] == { + "message": params["message"] + } + + def test_response_transform_preserves_audio_and_logprobs(): from litellm.a2a_protocol.litellm_completion_bridge.transformation import ( A2ACompletionBridgeTransformation,