mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(a2a): add per-agent a2a_protocol_version override for langgraph-api 0.15 cards
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5f969982fc
commit
633813e6ee
8 changed files with 231 additions and 12 deletions
|
|
@ -21,6 +21,7 @@ AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent-card.json"
|
|||
PREV_AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent.json"
|
||||
FOUNDRY_AGENT_CARD_PATH: Final = "/agentCard/v1.0"
|
||||
AGENT_CARD_PATH_PARAM: Final = "agent_card_path"
|
||||
A2A_PROTOCOL_VERSION_PARAM: Final = "a2a_protocol_version"
|
||||
|
||||
try:
|
||||
from a2a.client import A2ACardResolver as _A2ACardResolver
|
||||
|
|
@ -76,22 +77,43 @@ _CANONICAL_PROTOCOL_BINDINGS: Final = MappingProxyType(
|
|||
)
|
||||
|
||||
_LEGACY_PROTOCOL_VERSION: Final = "0.3"
|
||||
_A2A_PROTOCOL_VERSION_ALIASES: Final = MappingProxyType(
|
||||
{
|
||||
"0.3": "0.3",
|
||||
"0.3.0": "0.3",
|
||||
"1.0": "1.0",
|
||||
"1.0.0": "1.0",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def normalize_agent_card_interfaces(agent_card: "AgentCard") -> "AgentCard":
|
||||
def resolve_a2a_protocol_version(value: object) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool):
|
||||
verbose_logger.warning("Ignoring invalid %s value: %r", A2A_PROTOCOL_VERSION_PARAM, value)
|
||||
return None
|
||||
if not isinstance(value, (str, int, float)):
|
||||
verbose_logger.warning("Ignoring invalid %s value: %r", A2A_PROTOCOL_VERSION_PARAM, value)
|
||||
return None
|
||||
|
||||
protocol_version: Final = _A2A_PROTOCOL_VERSION_ALIASES.get(str(value).strip())
|
||||
if protocol_version is None:
|
||||
verbose_logger.warning("Ignoring invalid %s value: %r", A2A_PROTOCOL_VERSION_PARAM, value)
|
||||
return protocol_version
|
||||
|
||||
|
||||
def normalize_agent_card_interfaces(agent_card: "AgentCard", protocol_version: str | None = None) -> "AgentCard":
|
||||
"""
|
||||
Canonicalize the supported interfaces of spec-adjacent agent cards.
|
||||
|
||||
Some A2A servers (e.g. LangGraph Platform) serve agent cards with lowercase
|
||||
bindings like "jsonrpc", but a2a-sdk's ClientFactory matches bindings
|
||||
case-sensitively against its uppercase TransportProtocol constants and fails
|
||||
with "no compatible transports found." for spec-adjacent casings.
|
||||
|
||||
The same servers also speak the A2A 0.3 JSON dialect ("kind"-discriminated
|
||||
payloads) while declaring protocolVersion "1.0", which a2a-sdk's strict v1
|
||||
proto parsing rejects. A mis-cased binding fingerprints such a server, so its
|
||||
declared version is downgraded to 0.3 to route the SDK's ClientFactory onto
|
||||
its v0.3 compat transport, which speaks that dialect.
|
||||
with "no compatible transports found." for spec-adjacent casings. The casing
|
||||
fingerprint stopped working with langgraph-api 0.15, so operators can pin
|
||||
the version per agent via ``litellm_params.a2a_protocol_version``. An explicit
|
||||
"1.0" opts out of the mis-cased downgrade.
|
||||
"""
|
||||
normalized: Final = type(agent_card)()
|
||||
normalized.CopyFrom(agent_card)
|
||||
|
|
@ -101,6 +123,10 @@ def normalize_agent_card_interfaces(agent_card: "AgentCard") -> "AgentCard":
|
|||
continue
|
||||
interface.protocol_binding = canonical
|
||||
interface.protocol_version = _LEGACY_PROTOCOL_VERSION
|
||||
if protocol_version is not None:
|
||||
for interface in normalized.supported_interfaces:
|
||||
if interface.protocol_binding in _CANONICAL_PROTOCOL_BINDINGS.values():
|
||||
interface.protocol_version = protocol_version
|
||||
return normalized
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -23,6 +23,12 @@ if TYPE_CHECKING:
|
|||
from a2a.client import Client as A2AClientType
|
||||
|
||||
|
||||
_A2A_KIND_DIALECT_HINT: Final = (
|
||||
" (the agent replied in the A2A 0.3 dialect while its card declares 1.0; "
|
||||
'set litellm_params.a2a_protocol_version: "0.3" on this agent)'
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
from a2a.client import Client, ClientConfig, create_client
|
||||
|
||||
|
|
@ -152,7 +158,7 @@ def map_a2a_exception(
|
|||
|
||||
# Default: wrap in generic A2AError
|
||||
raise A2AError(
|
||||
message=error_str,
|
||||
message=error_str + _A2A_KIND_DIALECT_HINT if 'has no field named "kind"' in error_str else error_str,
|
||||
model=model,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from typing import Any, Final
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.card_resolver import AGENT_CARD_PATH_PARAM
|
||||
from litellm.a2a_protocol.card_resolver import A2A_PROTOCOL_VERSION_PARAM, AGENT_CARD_PATH_PARAM
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
||||
A2ACompletionBridgeTransformation,
|
||||
A2AStreamingContext,
|
||||
|
|
@ -38,6 +38,7 @@ _AGENT_ONLY_PARAMS: Final = frozenset(
|
|||
"agent_id",
|
||||
"agent_card_params",
|
||||
AGENT_CARD_PATH_PARAM,
|
||||
A2A_PROTOCOL_VERSION_PARAM,
|
||||
A2A_USER_API_KEY_HASH_PARAM,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -72,10 +72,12 @@ except ImportError:
|
|||
|
||||
# Import our custom card resolver that supports multiple well-known paths
|
||||
from litellm.a2a_protocol.card_resolver import (
|
||||
A2A_PROTOCOL_VERSION_PARAM,
|
||||
AGENT_CARD_PATH_PARAM,
|
||||
LiteLLMA2ACardResolver,
|
||||
get_agent_card_url,
|
||||
normalize_agent_card_interfaces,
|
||||
resolve_a2a_protocol_version,
|
||||
)
|
||||
from litellm.a2a_protocol.exception_mapping_utils import (
|
||||
handle_a2a_localhost_retry,
|
||||
|
|
@ -153,6 +155,10 @@ def _agent_card_path(litellm_params: Mapping[str, object]) -> str | None:
|
|||
return configured_path if isinstance(configured_path, str) and configured_path else None
|
||||
|
||||
|
||||
def _agent_protocol_version(litellm_params: Mapping[str, object]) -> str | None:
|
||||
return resolve_a2a_protocol_version(litellm_params.get(A2A_PROTOCOL_VERSION_PARAM))
|
||||
|
||||
|
||||
def _set_litellm_params_on_logging_obj(
|
||||
kwargs: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
|
|
@ -498,6 +504,7 @@ async def asend_message(
|
|||
base_url=api_base,
|
||||
extra_headers=extra_headers,
|
||||
relative_card_path=_agent_card_path(litellm_params),
|
||||
protocol_version=_agent_protocol_version(litellm_params),
|
||||
)
|
||||
|
||||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
|
|
@ -723,6 +730,7 @@ async def asend_message_streaming(
|
|||
extra_headers=extra_headers,
|
||||
streaming=True,
|
||||
relative_card_path=_agent_card_path(litellm_params),
|
||||
protocol_version=_agent_protocol_version(litellm_params),
|
||||
)
|
||||
|
||||
assert a2a_client is not None
|
||||
|
|
@ -770,6 +778,7 @@ async def create_a2a_client(
|
|||
extra_headers: dict[str, str] | None = None,
|
||||
streaming: bool = False,
|
||||
relative_card_path: str | None = None,
|
||||
protocol_version: str | None = None,
|
||||
) -> "A2AClientType":
|
||||
"""
|
||||
Create an A2A client for the given agent URL.
|
||||
|
|
@ -783,6 +792,7 @@ async def create_a2a_client(
|
|||
extra_headers: Optional additional headers to include in requests
|
||||
relative_card_path: Optional card path relative to ``base_url`` (e.g. ``agentCard/v1.0`` for a
|
||||
Microsoft Foundry agent); when None the well-known paths are probed in order
|
||||
protocol_version: Optional per-agent protocol version override for supported interfaces
|
||||
|
||||
Returns:
|
||||
An initialized a2a.client.A2AClient instance
|
||||
|
|
@ -819,7 +829,8 @@ async def create_a2a_client(
|
|||
await resolver.get_agent_card(
|
||||
relative_card_path=relative_card_path,
|
||||
http_kwargs=_card_http_kwargs(extra_headers),
|
||||
)
|
||||
),
|
||||
protocol_version=protocol_version,
|
||||
)
|
||||
|
||||
a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall]
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
|
||||
from litellm.a2a_protocol import exception_mapping_utils as emu
|
||||
from litellm.a2a_protocol.exceptions import A2ALocalhostURLError
|
||||
from litellm.a2a_protocol.exceptions import A2AError, A2ALocalhostURLError
|
||||
|
||||
|
||||
def _localhost_error() -> A2ALocalhostURLError:
|
||||
|
|
@ -17,6 +17,20 @@ def _localhost_error() -> A2ALocalhostURLError:
|
|||
)
|
||||
|
||||
|
||||
def test_map_a2a_exception_adds_protocol_version_hint_for_kind_field_error():
|
||||
with pytest.raises(A2AError) as error:
|
||||
emu.map_a2a_exception(RuntimeError('Message type "lf.a2a.v1.Task" has no field named "kind"'))
|
||||
|
||||
assert "a2a_protocol_version: " in str(error.value)
|
||||
|
||||
|
||||
def test_map_a2a_exception_does_not_add_protocol_version_hint_for_other_errors():
|
||||
with pytest.raises(A2AError) as error:
|
||||
emu.map_a2a_exception(RuntimeError("unrelated agent error"))
|
||||
|
||||
assert "a2a_protocol_version" not in str(error.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_localhost_retry_reuses_stashed_httpx_client():
|
||||
"""The retry must reuse the httpx client LiteLLM attached at creation (it carries
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.a2a_protocol.card_resolver import (
|
|||
fix_agent_card_url,
|
||||
is_localhost_or_internal_url,
|
||||
normalize_agent_card_interfaces,
|
||||
resolve_a2a_protocol_version,
|
||||
set_agent_card_url,
|
||||
)
|
||||
from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError
|
||||
|
|
@ -139,6 +140,77 @@ def test_normalize_agent_card_interfaces_downgrades_miscased_interfaces_to_the_0
|
|||
assert card.supported_interfaces[0].protocol_version == "1.0"
|
||||
|
||||
|
||||
def test_normalize_agent_card_interfaces_applies_override_to_canonical_bindings_without_mutating_input():
|
||||
pb2 = pytest.importorskip("a2a.types.a2a_pb2")
|
||||
card = pb2.AgentCard(
|
||||
name="langgraph",
|
||||
supported_interfaces=[
|
||||
pb2.AgentInterface(url="http://a/", protocol_binding="jsonrpc", protocol_version="1.0"),
|
||||
pb2.AgentInterface(url="http://b/", protocol_binding="JSONRPC", protocol_version="1.0"),
|
||||
pb2.AgentInterface(url="http://c/", protocol_binding="HTTP+JSON", protocol_version="1.0"),
|
||||
pb2.AgentInterface(url="http://d/", protocol_binding="GRPC", protocol_version="1.0"),
|
||||
pb2.AgentInterface(url="http://e/", protocol_binding="websocket", protocol_version="1.0"),
|
||||
],
|
||||
)
|
||||
|
||||
normalized = normalize_agent_card_interfaces(card, protocol_version="0.3")
|
||||
|
||||
assert [(item.protocol_binding, item.protocol_version) for item in normalized.supported_interfaces] == [
|
||||
("JSONRPC", "0.3"),
|
||||
("JSONRPC", "0.3"),
|
||||
("HTTP+JSON", "0.3"),
|
||||
("GRPC", "0.3"),
|
||||
("websocket", "1.0"),
|
||||
]
|
||||
assert card.supported_interfaces[0].protocol_binding == "jsonrpc"
|
||||
assert [(item.protocol_binding, item.protocol_version) for item in card.supported_interfaces] == [
|
||||
("jsonrpc", "1.0"),
|
||||
("JSONRPC", "1.0"),
|
||||
("HTTP+JSON", "1.0"),
|
||||
("GRPC", "1.0"),
|
||||
("websocket", "1.0"),
|
||||
]
|
||||
|
||||
|
||||
def test_normalize_agent_card_interfaces_override_opts_out_of_mis_cased_downgrade():
|
||||
pb2 = pytest.importorskip("a2a.types.a2a_pb2")
|
||||
card = pb2.AgentCard(
|
||||
name="langgraph",
|
||||
supported_interfaces=[
|
||||
pb2.AgentInterface(url="http://a/", protocol_binding="jsonrpc", protocol_version="1.0"),
|
||||
],
|
||||
)
|
||||
|
||||
normalized = normalize_agent_card_interfaces(card, protocol_version="1.0")
|
||||
|
||||
assert normalized.supported_interfaces[0].protocol_binding == "JSONRPC"
|
||||
assert normalized.supported_interfaces[0].protocol_version == "1.0"
|
||||
assert card.supported_interfaces[0].protocol_binding == "jsonrpc"
|
||||
assert card.supported_interfaces[0].protocol_version == "1.0"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("value", "expected"),
|
||||
[
|
||||
("0.3", "0.3"),
|
||||
(" 0.3 ", "0.3"),
|
||||
("0.3.0", "0.3"),
|
||||
(0.3, "0.3"),
|
||||
(1.0, "1.0"),
|
||||
("1.0.0", "1.0"),
|
||||
(None, None),
|
||||
(True, None),
|
||||
(1, None),
|
||||
("2.0", None),
|
||||
("abc", None),
|
||||
("", None),
|
||||
(["0.3"], None),
|
||||
],
|
||||
)
|
||||
def test_resolve_a2a_protocol_version(value, expected):
|
||||
assert resolve_a2a_protocol_version(value) == expected
|
||||
|
||||
|
||||
_FOUNDRY_BASE_URL: Final = "https://foundry.example.com/a2a"
|
||||
|
||||
_FOUNDRY_CARD_JSON: Final = {
|
||||
|
|
|
|||
|
|
@ -356,6 +356,7 @@ async def test_handle_streaming_keeps_agent_card_path_out_of_the_completion_call
|
|||
"custom_llm_provider": "langgraph",
|
||||
"model": "agent",
|
||||
"agent_card_path": "agentCard/v1.0",
|
||||
"a2a_protocol_version": "0.3",
|
||||
},
|
||||
api_base="http://localhost:2024",
|
||||
)
|
||||
|
|
@ -363,3 +364,4 @@ async def test_handle_streaming_keeps_agent_card_path_out_of_the_completion_call
|
|||
|
||||
assert len(events) == 4
|
||||
assert "agent_card_path" not in mock_acompletion.call_args.kwargs
|
||||
assert "a2a_protocol_version" not in mock_acompletion.call_args.kwargs
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""Tests for litellm/a2a_protocol/main.py non-streaming send behavior."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -23,6 +24,7 @@ from litellm.a2a_protocol.main import (
|
|||
asend_message,
|
||||
create_a2a_client,
|
||||
)
|
||||
from litellm.a2a_protocol.exceptions import A2AError
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -230,6 +232,16 @@ _LOWERCASE_BINDING_CARD = {
|
|||
"supportedInterfaces": [{"url": "http://127.0.0.1:9/", "protocolBinding": "jsonrpc", "protocolVersion": "1.0"}],
|
||||
}
|
||||
|
||||
_UPPERCASE_BINDING_CARD = {
|
||||
"name": "langgraph-agent",
|
||||
"version": "1.0.0",
|
||||
"capabilities": {"streaming": True},
|
||||
"defaultInputModes": ["text/plain"],
|
||||
"defaultOutputModes": ["text/plain"],
|
||||
"skills": [],
|
||||
"supportedInterfaces": [{"url": "http://127.0.0.1:9/", "protocolBinding": "JSONRPC", "protocolVersion": "1.0"}],
|
||||
}
|
||||
|
||||
|
||||
class _RequestRecorder:
|
||||
"""Records the headers httpx put on the wire, per outbound request."""
|
||||
|
|
@ -240,6 +252,7 @@ class _RequestRecorder:
|
|||
self.card_requests = []
|
||||
self.card_urls = []
|
||||
self.rpc_requests = []
|
||||
self.rpc_bodies = []
|
||||
self.client = None
|
||||
|
||||
def __call__(self, request: httpx.Request) -> httpx.Response:
|
||||
|
|
@ -249,6 +262,7 @@ class _RequestRecorder:
|
|||
self.card_urls.append(str(request.url))
|
||||
return httpx.Response(200, json=self.card)
|
||||
self.rpc_requests.append(headers)
|
||||
self.rpc_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(200, json=self.rpc_reply)
|
||||
|
||||
|
||||
|
|
@ -392,6 +406,79 @@ async def test_lowercase_protocol_binding_card_round_trips_the_langgraph_dialect
|
|||
assert interface.protocol_version == "0.3"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_uppercase_binding_card_without_override_fails_with_actionable_hint(isolated_client_cache):
|
||||
await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_reply=_LANGGRAPH_TASK_REPLY)
|
||||
|
||||
with pytest.raises(A2AError) as error:
|
||||
await asend_message(request=_send_request("uppercase-default"), api_base="http://127.0.0.1:9")
|
||||
|
||||
assert 'has no field named "kind"' in str(error.value)
|
||||
assert "a2a_protocol_version" in str(error.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("card", [_LOWERCASE_BINDING_CARD, _UPPERCASE_BINDING_CARD])
|
||||
@pytest.mark.asyncio
|
||||
async def test_protocol_version_override_round_trips_the_langgraph_dialect(card, isolated_client_cache):
|
||||
recorder = await _seed_shared_a2a_client(card=card, rpc_reply=_LANGGRAPH_TASK_REPLY)
|
||||
|
||||
a2a_client = await create_a2a_client(base_url="http://127.0.0.1:9", protocol_version="0.3")
|
||||
response = await _send_message(a2a_client, _send_request("override-round-trip"))
|
||||
|
||||
assert type(response.root.result).__name__ == "Task"
|
||||
assert response.root.result.artifacts[0].parts[0].root.text == "langgraph echo: hi"
|
||||
assert a2a_client._litellm_agent_card.supported_interfaces[0].protocol_version == "0.3"
|
||||
assert recorder.rpc_bodies[-1]["method"] == "message/send"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_resolves_float_protocol_version_from_agent_params(isolated_client_cache):
|
||||
recorder = await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_reply=_LANGGRAPH_TASK_REPLY)
|
||||
|
||||
response = await asend_message(
|
||||
request=_send_request("float-protocol-override"),
|
||||
api_base="http://127.0.0.1:9",
|
||||
litellm_params={"a2a_protocol_version": 0.3},
|
||||
)
|
||||
|
||||
assert response.result["artifacts"][0]["parts"][0]["text"] == "langgraph echo: hi"
|
||||
assert recorder.rpc_bodies[-1]["method"] == "message/send"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("litellm_params", "expected_protocol_version"),
|
||||
[({"a2a_protocol_version": "0.3"}, "0.3"), (None, None)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_passes_agent_protocol_version_to_client(litellm_params, expected_protocol_version):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.a2a_protocol import main as a2a_main
|
||||
|
||||
request = SendStreamingMessageRequest(
|
||||
id="stream-protocol-version",
|
||||
params=MessageSendParams(
|
||||
message={"messageId": "m1", "role": "user", "parts": [{"kind": "text", "text": "hi"}]}
|
||||
),
|
||||
)
|
||||
captured = {}
|
||||
|
||||
async def _capture(*, base_url, extra_headers=None, streaming=False, protocol_version=None, **_):
|
||||
captured["protocol_version"] = protocol_version
|
||||
raise RuntimeError("stop")
|
||||
|
||||
with patch.object(a2a_main, "create_a2a_client", new=AsyncMock(side_effect=_capture)):
|
||||
with pytest.raises(RuntimeError, match="stop"):
|
||||
async for _ in a2a_main.asend_message_streaming(
|
||||
request=request,
|
||||
api_base="http://upstream.local",
|
||||
litellm_params=litellm_params,
|
||||
):
|
||||
pass
|
||||
|
||||
assert captured["protocol_version"] == expected_protocol_version
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_card_fetch_carries_the_callers_headers(isolated_client_cache):
|
||||
"""Agent cards can sit behind the same auth as the agent, so the card fetch must stay
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue