fix: honor registered A2A providers

This commit is contained in:
aiedwardyi 2026-08-24 00:23:31 +09:00
parent f005afa146
commit c3d599dd27
No known key found for this signature in database
4 changed files with 294 additions and 13 deletions

View file

@ -5,20 +5,134 @@ Handles routing for A2A agents (models with "a2a/<agent-name>" prefix).
Looks up agents in the registry and injects their API base URL.
"""
from typing import Any, Final
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from uuid import uuid4
from fastapi import HTTPException
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.llms.a2a.common_utils import A2AError, convert_messages_to_prompt, extract_text_from_a2a_response
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, Message, ModelResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
_HEADERS_ADAPTER: Final = TypeAdapter(dict[str, str])
_MESSAGES_ADAPTER: Final = TypeAdapter(list[AllMessageValues])
class _A2ATextPart(TypedDict):
kind: ReadOnly[str]
text: ReadOnly[str]
class _A2AMessage(TypedDict):
role: ReadOnly[str]
parts: ReadOnly[tuple[_A2ATextPart, ...]]
messageId: ReadOnly[str]
class _A2AParams(TypedDict):
message: ReadOnly[_A2AMessage]
async def _route_registered_provider(
data: Mapping[str, object],
model_name: str,
api_base: str,
litellm_params: Mapping[str, object],
) -> ModelResponse | CustomStreamWrapper:
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator
raw_messages: Final = data.get("messages")
messages: Final = _MESSAGES_ADAPTER.validate_python(raw_messages)
stream: Final = data.get("stream") is True
request_id: Final = str(uuid4())
params: Final[_A2AParams] = {
"message": {
"role": "user",
"parts": ({"kind": "text", "text": convert_messages_to_prompt(messages)},),
"messageId": str(uuid4()),
}
}
provider_params: Final = _OBJECT_DICT_ADAPTER.validate_python(litellm_params)
bridge_params: Final = _OBJECT_DICT_ADAPTER.validate_python(params)
configured_headers: Final = litellm_params.get("extra_headers") or litellm_params.get("headers")
agent_extra_headers: Final = (
_HEADERS_ADAPTER.validate_python(configured_headers) if isinstance(configured_headers, dict) else None
)
if stream:
streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming(
request_id=request_id,
params=bridge_params,
litellm_params=provider_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
completion_stream: Final = A2AModelResponseIterator(
streaming_response=streaming_response,
sync_stream=False,
model=model_name,
)
logging_obj: Final = data.get("litellm_logging_obj")
if not isinstance(logging_obj, Logging):
raise TypeError("litellm_logging_obj is required for streaming A2A requests")
return CustomStreamWrapper(
completion_stream=completion_stream,
model=model_name,
custom_llm_provider="a2a",
logging_obj=logging_obj,
stream_options=data.get("stream_options"),
)
response: Final = await A2ACompletionBridgeHandler.handle_non_streaming(
request_id=request_id,
params=bridge_params,
litellm_params=provider_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
error_value: Final = response.get("error")
if isinstance(error_value, dict):
error: Final = _OBJECT_DICT_ADAPTER.validate_python(error_value)
error_message: Final = error.get("message")
raise A2AError(
status_code=500,
message=f"A2A error: {error_message if isinstance(error_message, str) else 'Unknown error'}",
)
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
Choices(finish_reason="stop", index=0, message=Message(content=text, role="assistant"))
],
)
return model_response
async def route_a2a_agent_request(
data: dict,
data: Mapping[str, object],
route_type: str,
user_api_key_dict: UserAPIKeyAuth | None = None,
) -> Any | None:
) -> Awaitable[object] | None:
"""
Route A2A agent requests directly to litellm with injected API base.
@ -69,13 +183,34 @@ async def route_a2a_agent_request(
)
# Get API base URL from agent config
if not agent.agent_card_params or "url" not in agent.agent_card_params:
agent_card_params: Final = agent.agent_card_params
agent_url: Final = agent_card_params.get("url") if agent_card_params else None
if not isinstance(agent_url, str) or not agent_url:
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)
# Inject API base and route to litellm
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
registered_params_value: Final = agent.litellm_params
registered_provider_value: Final = (
registered_params_value.get("custom_llm_provider") if registered_params_value else None
)
registered_provider: Final = registered_provider_value if isinstance(registered_provider_value, str) else None
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) else agent_url
if (
registered_provider
and registered_provider != "a2a"
and route_type == "acompletion"
and registered_params_value is not None
):
verbose_proxy_logger.debug("[A2A] Routing %s through %s", model_name, registered_provider)
return _route_registered_provider(
data=data,
model_name=model_name,
api_base=api_base,
litellm_params=registered_params_value,
)
return getattr(litellm, f"{route_type}")(**data)
completion_data: Final = MappingProxyType({**data, "api_base": api_base})
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, api_base)
return getattr(litellm, f"{route_type}")(**completion_data) # pyright: ignore[reportAny] # dynamic SDK route

View file

@ -4,12 +4,18 @@ Helper functions for appending A2A agents to model lists.
Used by proxy model endpoints to make agents appear in UI alongside models.
"""
from typing import Final
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
ModelGroupInfoProxy,
)
_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
async def append_agents_to_model_group(
model_groups: list[ModelGroupInfoProxy],
@ -70,12 +76,15 @@ async def append_agents_to_model_info(
for agent_id in allowed_agent_ids:
agent = global_agent_registry.get_agent_by_id(agent_id)
if agent is not None:
agent_params = agent.litellm_params
provider_value = agent_params.get("custom_llm_provider") if agent_params else None
custom_llm_provider = provider_value if isinstance(provider_value, str) else "a2a"
models.append(
{
"model_name": f"a2a/{agent.agent_name}",
"litellm_params": {
"model": f"a2a/{agent.agent_name}",
"custom_llm_provider": "a2a",
"custom_llm_provider": custom_llm_provider,
},
"model_info": {
"id": agent.agent_id,

View file

@ -4,8 +4,6 @@ Test appending A2A agents to model lists.
Maps to: litellm/proxy/agent_endpoints/model_list_helpers.py
"""
from unittest.mock import AsyncMock, Mock, patch
import pytest
@ -109,3 +107,32 @@ async def test_append_agents_to_model_info():
assert result[0]["litellm_params"]["custom_llm_provider"] == "a2a"
assert result[0]["model_info"]["id"] == "agent-123"
assert result[0]["model_info"]["mode"] == "chat"
@pytest.mark.asyncio
async def test_append_agents_to_model_info_preserves_registered_provider():
agent = AgentResponse(
agent_id="agent-123",
agent_name="test-agent",
agent_card_params={"url": "http://example.com"},
litellm_params={"custom_llm_provider": "pydantic_ai_agents"},
)
registry = Mock()
registry.get_agent_by_id = Mock(return_value=agent)
with (
patch( # test-quality-ok: access resolution is outside model-list assembly
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.resolve_agent_access",
AsyncMock(return_value=RestrictedAgentAccess(frozenset({"agent-123"}))),
),
patch( # test-quality-ok: registry output drives model-list assembly
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
registry,
),
):
result = await append_agents_to_model_info(
models=[],
user_api_key_dict=Mock(spec=UserAPIKeyAuth),
)
assert result[0]["litellm_params"]["custom_llm_provider"] == "pydantic_ai_agents"

View file

@ -4,8 +4,6 @@ Test A2A model routing in proxy.
Maps to: litellm/proxy/agent_endpoints/a2a_routing.py
"""
from unittest.mock import AsyncMock, Mock, patch
import pytest
@ -72,6 +70,118 @@ async def test_route_a2a_model_bypasses_router():
assert call_kwargs["api_base"] == "http://agent.example.com"
@pytest.mark.asyncio
async def test_route_a2a_model_uses_registered_provider():
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"},
)
data = {
"model": "a2a/test-agent",
"messages": [{"role": "user", "content": "Hello"}],
}
bridge_response = {
"jsonrpc": "2.0",
"id": "request-id",
"result": {
"kind": "message",
"role": "agent",
"parts": [{"kind": "text", "text": "Hello back"}],
"messageId": "message-id",
},
}
with (
patch( # test-quality-ok: registry lookup is the routing seam
"litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through",
AsyncMock(return_value=agent),
),
patch( # test-quality-ok: access control is outside this routing test
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
AsyncMock(return_value=True),
),
patch( # test-quality-ok: provider dispatch is the tested seam
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming",
AsyncMock(return_value=bridge_response),
) as bridge,
patch( # test-quality-ok: generic dispatch must stay unused
"litellm.acompletion", AsyncMock()
) as generic_completion,
):
call = await route_a2a_agent_request(data, "acompletion")
response = await call
bridge.assert_awaited_once()
generic_completion.assert_not_called()
assert response.choices[0].message.content == "Hello back"
@pytest.mark.asyncio
async def test_route_a2a_stream_uses_registered_provider():
from litellm.litellm_core_utils.litellm_logging import Logging
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"},
)
logging_obj = Mock(spec=Logging)
data = {
"model": "a2a/test-agent",
"messages": [{"role": "user", "content": "Hello"}],
"stream": True,
"litellm_logging_obj": logging_obj,
}
provider_stream = object()
completion_stream = object()
wrapper = object()
with (
patch( # test-quality-ok: registry lookup is the routing seam
"litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through",
AsyncMock(return_value=agent),
),
patch( # test-quality-ok: access control is outside this routing test
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
AsyncMock(return_value=True),
),
patch( # test-quality-ok: provider dispatch is the tested seam
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_streaming",
Mock(return_value=provider_stream),
) as bridge,
patch( # test-quality-ok: iterator wiring is the tested seam
"litellm.llms.a2a.chat.streaming_iterator.A2AModelResponseIterator",
Mock(return_value=completion_stream),
),
patch( # test-quality-ok: wrapper wiring is the tested seam
"litellm.litellm_core_utils.streaming_handler.CustomStreamWrapper",
Mock(return_value=wrapper),
) as stream_wrapper,
patch( # test-quality-ok: generic dispatch must stay unused
"litellm.acompletion", AsyncMock()
) as generic_completion,
):
call = await route_a2a_agent_request(data, "acompletion")
response = await call
bridge.assert_called_once()
stream_wrapper.assert_called_once_with(
completion_stream=completion_stream,
model="a2a/test-agent",
custom_llm_provider="a2a",
logging_obj=logging_obj,
stream_options=None,
)
generic_completion.assert_not_called()
assert response is wrapper
@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"""