mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: harden A2A provider routing
This commit is contained in:
parent
d2bfb5ed1b
commit
9af6eacf84
3 changed files with 53 additions and 8 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue