Merge pull request #25514 from milan-berri/fix/a2a-create-a2a-client-default-timeout

fix: a2a create a2a client default 60 second timeout
This commit is contained in:
ishaan-berri 2026-04-10 16:51:34 -07:00 • committed by GitHub
commit b049aadc96
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 87 additions and 4 deletions

View file

@ -615,7 +615,7 @@ async def asend_message_streaming( # noqa: PLR0915
async def create_a2a_client(
base_url: str,
timeout: float = 60.0,
timeout: float = DEFAULT_A2A_AGENT_TIMEOUT,
extra_headers: Optional[Dict[str, str]] = None,
) -> "A2AClientType":
"""
@ -626,7 +626,7 @@ async def create_a2a_client(
Args:
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
timeout: Request timeout in seconds (default: 60.0)
timeout: Request timeout in seconds (default: ``DEFAULT_A2A_AGENT_TIMEOUT`` / env ``DEFAULT_A2A_AGENT_TIMEOUT``)
extra_headers: Optional additional headers to include in requests
Returns:
@ -711,7 +711,7 @@ async def aget_agent_card(
Args:
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
timeout: Request timeout in seconds (default: 60.0)
timeout: Request timeout in seconds (default: ``DEFAULT_A2A_AGENT_TIMEOUT`` / env ``DEFAULT_A2A_AGENT_TIMEOUT``)
extra_headers: Optional additional headers to include in requests
Returns:

View file

@ -4,6 +4,9 @@ Tests that prove header isolation between agents.
Before the fix these tests FAIL — agent A's headers bleed into agent B
because create_a2a_client mutates a globally cached httpx client.
After the fix they pass.
Also includes direct unit tests for create_a2a_client (fresh httpx client
per call; default timeout uses DEFAULT_A2A_AGENT_TIMEOUT).
"""
import sys
@ -11,6 +14,8 @@ from unittest.mock import AsyncMock, MagicMock, call, patch
import pytest
from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT
# ---------------------------------------------------------------------------
# Helpers
@ -199,7 +204,7 @@ async def test_each_agent_gets_only_its_own_static_headers():
# ---------------------------------------------------------------------------
# Unit test: create_a2a_client uses a fresh httpx client per call
# Unit tests: create_a2a_client (httpx client per call + timeout defaults)
# ---------------------------------------------------------------------------
@ -246,3 +251,81 @@ async def test_create_a2a_client_uses_fresh_httpx_client():
assert created_clients[0] is not created_clients[1], (
"create_a2a_client reused a cached httpx client — headers will bleed between agents"
)
@pytest.mark.asyncio
async def test_create_a2a_client_default_timeout_matches_constant():
"""When timeout is omitted, httpx client params must use DEFAULT_A2A_AGENT_TIMEOUT."""
from litellm.a2a_protocol.main import create_a2a_client
captured: dict = {}
def _capture_get_async_httpx_client(llm_provider, params, **kwargs):
captured["params"] = params
handler = MagicMock()
handler.client = MagicMock()
handler.client.headers = MagicMock()
return handler
fake_agent_card = MagicMock()
fake_agent_card.name = "test-agent"
class _FakeResolver:
def __init__(self, **kw):
pass
async def get_agent_card(self):
return fake_agent_card
class _FakeA2AClient:
def __init__(self, httpx_client, agent_card):
pass
with patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), patch(
"litellm.a2a_protocol.main.get_async_httpx_client",
side_effect=_capture_get_async_httpx_client,
), patch("litellm.a2a_protocol.main.A2ACardResolver", _FakeResolver), patch(
"litellm.a2a_protocol.main._A2AClient", _FakeA2AClient
):
await create_a2a_client(base_url="http://127.0.0.1:9")
assert captured["params"]["timeout"] == DEFAULT_A2A_AGENT_TIMEOUT
@pytest.mark.asyncio
async def test_create_a2a_client_explicit_timeout_overrides_default():
"""Explicit timeout= must be passed through to the httpx client params."""
from litellm.a2a_protocol.main import create_a2a_client
captured: dict = {}
def _capture_get_async_httpx_client(llm_provider, params, **kwargs):
captured["params"] = params
handler = MagicMock()
handler.client = MagicMock()
handler.client.headers = MagicMock()
return handler
fake_agent_card = MagicMock()
fake_agent_card.name = "test-agent"
class _FakeResolver:
def __init__(self, **kw):
pass
async def get_agent_card(self):
return fake_agent_card
class _FakeA2AClient:
def __init__(self, httpx_client, agent_card):
pass
with patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), patch(
"litellm.a2a_protocol.main.get_async_httpx_client",
side_effect=_capture_get_async_httpx_client,
), patch("litellm.a2a_protocol.main.A2ACardResolver", _FakeResolver), patch(
"litellm.a2a_protocol.main._A2AClient", _FakeA2AClient
):
await create_a2a_client(base_url="http://127.0.0.1:9", timeout=42.5)
assert captured["params"]["timeout"] == 42.5