diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 8a795214750..0373033ee10 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -5,20 +5,134 @@ Handles routing for A2A agents (models with "a2a/" 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 diff --git a/litellm/proxy/agent_endpoints/model_list_helpers.py b/litellm/proxy/agent_endpoints/model_list_helpers.py index 4a88644bae5..624136c020b 100644 --- a/litellm/proxy/agent_endpoints/model_list_helpers.py +++ b/litellm/proxy/agent_endpoints/model_list_helpers.py @@ -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, diff --git a/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py b/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py index 939ab1cab40..b82cd1dd231 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py @@ -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" diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 35308474949..6273fdc0df0 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -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"""