mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: honor registered A2A providers
This commit is contained in:
parent
f005afa146
commit
c3d599dd27
4 changed files with 294 additions and 13 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue