mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
Fix: use registered agent_card_params for A2A message/send instead of re-discovering via well-known paths (#40586)
This commit is contained in:
parent
fbed17d567
commit
ce8b116396
3 changed files with 121 additions and 16 deletions
|
|
@ -387,9 +387,7 @@ async def asend_message(
|
|||
api_base: str | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
agent_id: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
**kwargs: object,
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
agent_extra_headers: dict[str, str] | None = None, agent_card_params: dict[str, object] | None = None, **kwargs: object, ) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
Async: Send a message to an A2A agent.
|
||||
|
||||
|
|
@ -474,7 +472,7 @@ async def asend_message(
|
|||
# Overlay agent-level headers (agent headers take precedence over LiteLLM internal ones)
|
||||
if agent_extra_headers:
|
||||
extra_headers.update(agent_extra_headers)
|
||||
a2a_client = await create_a2a_client(base_url=api_base, extra_headers=extra_headers)
|
||||
a2a_client = await create_a2a_client( base_url=api_base, extra_headers=extra_headers, agent_card_params=agent_card_params, )
|
||||
|
||||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
assert a2a_client is not None
|
||||
|
|
@ -609,9 +607,9 @@ async def asend_message_streaming(
|
|||
agent_id: str | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
proxy_server_request: dict[str, object] | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
**kwargs: object,
|
||||
) -> AsyncIterator[Any]:
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
agent_card_params: dict[str, object] | None = None,
|
||||
**kwargs: object, ) -> AsyncIterator[Any]:
|
||||
"""
|
||||
Async: Send a streaming message to an A2A agent.
|
||||
|
||||
|
|
@ -695,11 +693,7 @@ async def asend_message_streaming(
|
|||
extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
if agent_extra_headers:
|
||||
extra_headers.update(agent_extra_headers)
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url=api_base,
|
||||
extra_headers=extra_headers,
|
||||
streaming=True,
|
||||
)
|
||||
a2a_client = await create_a2a_client( base_url=api_base, extra_headers=extra_headers, streaming=True, agent_card_params=agent_card_params, )
|
||||
|
||||
assert a2a_client is not None
|
||||
|
||||
|
|
@ -745,6 +739,7 @@ async def create_a2a_client(
|
|||
timeout: float = DEFAULT_A2A_AGENT_TIMEOUT,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
streaming: bool = False,
|
||||
agent_card_params: dict[str, object] | None = None,
|
||||
) -> "A2AClientType":
|
||||
"""
|
||||
Create an A2A client for the given agent URL.
|
||||
|
|
@ -787,10 +782,34 @@ async def create_a2a_client(
|
|||
if extra_headers:
|
||||
verbose_proxy_logger.debug("A2A client created with extra_headers=%s", list(extra_headers.keys()))
|
||||
|
||||
resolver: Final = A2ACardResolver(httpx_client=httpx_client, base_url=base_url)
|
||||
agent_card: Final = normalize_agent_card_interfaces(
|
||||
await resolver.get_agent_card(http_kwargs={"headers": extra_headers} if extra_headers else None)
|
||||
)
|
||||
agent_card: "AgentCard | None" = None
|
||||
|
||||
if agent_card_params:
|
||||
from pydantic import ValidationError as _ValidationError
|
||||
|
||||
from a2a.compat.v0_3.types import AgentCard as _AgentCard
|
||||
|
||||
try:
|
||||
agent_card = normalize_agent_card_interfaces(
|
||||
_AgentCard.model_validate(agent_card_params)
|
||||
)
|
||||
verbose_logger.info(
|
||||
"Using pre-registered agent card for %s (skipping well-known discovery)", base_url
|
||||
)
|
||||
except _ValidationError as e:
|
||||
verbose_logger.warning(
|
||||
"Stored agent_card_params for %s failed AgentCard validation (%s); "
|
||||
"falling back to well-known discovery.",
|
||||
base_url,
|
||||
e,
|
||||
)
|
||||
agent_card = None
|
||||
|
||||
if agent_card is None:
|
||||
resolver: Final = A2ACardResolver(httpx_client=httpx_client, base_url=base_url)
|
||||
agent_card = normalize_agent_card_interfaces(
|
||||
await resolver.get_agent_card(http_kwargs={"headers": extra_headers} if extra_headers else None)
|
||||
)
|
||||
|
||||
a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
agent_card,
|
||||
|
|
|
|||
|
|
@ -404,6 +404,7 @@ async def _handle_stream_message(
|
|||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
request_data: dict[str, object] | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
agent_card_params: dict[str, object] | None = None,
|
||||
served_version: A2AVersion = "0.3",
|
||||
) -> StreamingResponse:
|
||||
"""Handle message/stream method via SDK functions.
|
||||
|
|
@ -474,6 +475,7 @@ async def _handle_stream_message(
|
|||
metadata=metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
agent_card_params=agent_card_params,
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
@ -863,6 +865,7 @@ async def invoke_agent_a2a(
|
|||
proxy_server_request=data.get("proxy_server_request"),
|
||||
litellm_logging_obj=logging_obj,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
agent_card_params=agent_card_params or None,
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
@ -904,6 +907,7 @@ async def invoke_agent_a2a(
|
|||
agent_extra_headers=agent_extra_headers,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=data,
|
||||
agent_card_params=agent_card_params or None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
served_version=served_version,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,82 @@
|
|||
"""
|
||||
Regression tests for GH #40586: message/send re-discovered the agent card
|
||||
from well-known paths instead of using the registered agent_card_params.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.a2a_protocol.main import create_a2a_client
|
||||
|
||||
MINIMAL_AGENT_CARD_PARAMS = {
|
||||
"protocolVersion": "1.0",
|
||||
"name": "agent",
|
||||
"description": "A test assistant.",
|
||||
"url": "https://example.com/agents/a2a",
|
||||
"version": "1.0",
|
||||
"capabilities": {"streaming": False},
|
||||
"defaultInputModes": ["text"],
|
||||
"defaultOutputModes": ["text"],
|
||||
"skills": [],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_a2a_client_skips_discovery_when_agent_card_params_given():
|
||||
"""
|
||||
When a caller already has a resolved agent card (e.g. one stored via
|
||||
POST /v1/agents), create_a2a_client must build the AgentCard directly
|
||||
from it instead of probing the upstream server's well-known paths.
|
||||
"""
|
||||
with (
|
||||
patch("litellm.a2a_protocol.main.get_async_httpx_client") as mock_get_httpx,
|
||||
patch("litellm.a2a_protocol.main.A2ACardResolver") as mock_resolver_cls,
|
||||
patch("litellm.a2a_protocol.main.create_client", new_callable=AsyncMock) as mock_create_client,
|
||||
patch(
|
||||
"litellm.a2a_protocol.main.normalize_agent_card_interfaces",
|
||||
side_effect=lambda card: card,
|
||||
),
|
||||
):
|
||||
mock_get_httpx.return_value.client = MagicMock()
|
||||
mock_create_client.return_value = MagicMock()
|
||||
|
||||
await create_a2a_client(
|
||||
base_url="https://example.com/agents/a2a",
|
||||
agent_card_params=MINIMAL_AGENT_CARD_PARAMS,
|
||||
)
|
||||
|
||||
# The whole point of the fix: the resolver must never be touched
|
||||
# when a card was already supplied at registration time.
|
||||
mock_resolver_cls.assert_not_called()
|
||||
|
||||
mock_create_client.assert_awaited_once()
|
||||
called_card = mock_create_client.await_args.args[0]
|
||||
assert called_card.name == "agent"
|
||||
assert str(called_card.url) == "https://example.com/agents/a2a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_a2a_client_falls_back_to_discovery_without_agent_card_params():
|
||||
"""
|
||||
Unchanged behavior: when no agent_card_params is supplied, resolve the
|
||||
card via the well-known-path resolver, same as before the fix.
|
||||
"""
|
||||
with (
|
||||
patch("litellm.a2a_protocol.main.get_async_httpx_client") as mock_get_httpx,
|
||||
patch("litellm.a2a_protocol.main.A2ACardResolver") as mock_resolver_cls,
|
||||
patch("litellm.a2a_protocol.main.create_client", new_callable=AsyncMock) as mock_create_client,
|
||||
patch(
|
||||
"litellm.a2a_protocol.main.normalize_agent_card_interfaces",
|
||||
side_effect=lambda card: card,
|
||||
),
|
||||
):
|
||||
mock_get_httpx.return_value.client = MagicMock()
|
||||
mock_resolver_instance = mock_resolver_cls.return_value
|
||||
mock_resolver_instance.get_agent_card = AsyncMock(return_value=MagicMock())
|
||||
mock_create_client.return_value = MagicMock()
|
||||
|
||||
await create_a2a_client(base_url="https://example.com/agents/a2a")
|
||||
|
||||
mock_resolver_cls.assert_called_once()
|
||||
mock_resolver_instance.get_agent_card.assert_awaited_once()
|
||||
Loading…
Add table
Reference in a new issue