Fix: use registered agent_card_params for A2A message/send instead of re-discovering via well-known paths (#40586)

This commit is contained in:
AlphaRex-pixel 2026-09-11 00:30:38 +05:30
parent fbed17d567
commit ce8b116396
3 changed files with 121 additions and 16 deletions

View file

@ -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,

View file

@ -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,
)

View file

@ -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()