fix: harden A2A provider routing

This commit is contained in:
aiedwardyi 2026-08-26 12:32:12 +09:00
parent d2bfb5ed1b
commit 9af6eacf84
No known key found for this signature in database
3 changed files with 53 additions and 8 deletions

View file

@ -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)

View file

@ -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)

View file

@ -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,