diff --git a/src/a2a/client/client_factory.py b/src/a2a/client/client_factory.py index a59189ade..e43207ca7 100644 --- a/src/a2a/client/client_factory.py +++ b/src/a2a/client/client_factory.py @@ -55,6 +55,12 @@ ] +def major_minor(version: str) -> tuple[int, int]: + """Returns the `Major.Minor` release of a protocol version.""" + v = Version(version) + return (v.major, v.minor) + + class ClientFactory: """Factory for creating clients that communicate with A2A agents. @@ -229,10 +235,10 @@ def _find_best_interface( if not candidates: return None - # Prefer interface with version 1.0 - for i in candidates: - if i.protocol_version == PROTOCOL_VERSION_1_0: - return i + # Prefer interface with version 1.0. Patch versions must not affect + # negotiation, so '1.0.0' counts as '1.0'. + v1_0 = major_minor(PROTOCOL_VERSION_1_0) + v0_3 = major_minor(PROTOCOL_VERSION_0_3) best_gt_1_0 = None best_ge_0_3 = None @@ -245,10 +251,12 @@ def _find_best_interface( continue try: - v = Version(i.protocol_version) - if best_gt_1_0 is None and v > Version(PROTOCOL_VERSION_1_0): + v = major_minor(i.protocol_version) + if v == v1_0: + return i + if best_gt_1_0 is None and v > v1_0: best_gt_1_0 = i - if best_ge_0_3 is None and v >= Version(PROTOCOL_VERSION_0_3): + if best_ge_0_3 is None and v >= v0_3: best_ge_0_3 = i except InvalidVersion: pass diff --git a/tests/client/test_client_factory.py b/tests/client/test_client_factory.py index d211a7331..869881086 100644 --- a/tests/client/test_client_factory.py +++ b/tests/client/test_client_factory.py @@ -115,6 +115,36 @@ def test_client_factory_server_preference(base_agent_card: AgentCard): assert client._transport.url == 'http://primary-url.com' # type: ignore[attr-defined] +@pytest.mark.parametrize('v1_version', ['1.0', '1.0.0']) +def test_client_factory_prefers_1_0_ignoring_patch_version( + base_agent_card: AgentCard, v1_version: str +): + """A patch version must not affect protocol negotiation (spec: Versioning).""" + del base_agent_card.supported_interfaces[:] + base_agent_card.supported_interfaces.extend( + [ + AgentInterface( + protocol_binding=TransportProtocol.JSONRPC, + url='http://legacy-url.com', + protocol_version='0.3.0', + ), + AgentInterface( + protocol_binding=TransportProtocol.JSONRPC, + url='http://v1-url.com', + protocol_version=v1_version, + ), + ] + ) + config = ClientConfig( + httpx_client=httpx.AsyncClient(), + supported_protocol_bindings=[TransportProtocol.JSONRPC], + ) + client = ClientFactory(config).create(base_agent_card) + + assert isinstance(client._transport, JsonRpcTransport) # type: ignore[attr-defined] + assert client._transport.url == 'http://v1-url.com' # type: ignore[attr-defined] + + def test_client_factory_no_compatible_transport(base_agent_card: AgentCard): """Verify that the factory raises an error if no compatible transport is found.""" config = ClientConfig(