diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 638cb22..9cddc89 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -71,6 +71,14 @@ uv run pytest tests/contracts/test_extension_contract_consistency.py uv run mypy src/opencode_a2a ``` +If you change outbound authentication, tracing, protocol-version, or extension headers, run the SDK transport regression directly: + +```bash +uv run pytest --no-cov tests/client/test_client_http_headers.py +``` + +Assertions against a mocked SDK client or `ClientCallContext` alone do not prove that headers reached the HTTP request. Keep at least one regression test across Agent Card resolution, the real SDK transport, and the final `httpx.Request` boundary. + ## Change Expectations - Keep code, comments, and docs in English. diff --git a/src/opencode_a2a/client/request_context.py b/src/opencode_a2a/client/request_context.py index 570b685..5d5694f 100644 --- a/src/opencode_a2a/client/request_context.py +++ b/src/opencode_a2a/client/request_context.py @@ -79,16 +79,13 @@ def build_call_context( if extra_headers: merged_headers.update(extra_headers) normalized_extensions = [value for value in (extensions or ()) if value] - service_parameters = None - if normalized_extensions: - service_parameters = ServiceParametersFactory.create_from( - None, - [with_a2a_extensions(normalized_extensions)], - ) - return ClientCallContext( - state={ - "headers": dict(merged_headers), - "http_kwargs": {"headers": dict(merged_headers)}, - }, - service_parameters=service_parameters, + service_parameter_updates = ( + [with_a2a_extensions(normalized_extensions)] if normalized_extensions else [] ) + # a2a-sdk transports serialize service parameters as HTTP headers or gRPC + # metadata. This is also the channel used by the SDK's AuthInterceptor. + service_parameters = ServiceParametersFactory.create_from( + merged_headers, + service_parameter_updates, + ) + return ClientCallContext(service_parameters=service_parameters) diff --git a/tests/client/test_client_facade.py b/tests/client/test_client_facade.py index 2a24004..ebc0f75 100644 --- a/tests/client/test_client_facade.py +++ b/tests/client/test_client_facade.py @@ -333,7 +333,7 @@ async def test_send_message_adds_bearer_token_from_settings( request, _, kwargs = fake_client.send_message_inputs[0] assert request.metadata == {} assert kwargs["context"] is not None - assert kwargs["context"].state["headers"]["Authorization"] == "Bearer peer-token" + assert kwargs["context"].service_parameters["Authorization"] == "Bearer peer-token" @pytest.mark.asyncio @@ -354,7 +354,7 @@ async def test_send_message_adds_basic_auth_from_settings( request, _, kwargs = fake_client.send_message_inputs[0] assert request.metadata == {} assert kwargs["context"] is not None - assert kwargs["context"].state["headers"]["Authorization"] == ( + assert kwargs["context"].service_parameters["Authorization"] == ( f"Basic {b64encode(b'user:pass').decode()}" ) @@ -382,7 +382,7 @@ async def test_send_message_preserves_explicit_authorization_metadata( assert result[0].HasField("message") request, _, kwargs = fake_client.send_message_inputs[0] assert request.metadata == {"trace_id": "trace-1"} - assert kwargs["context"].state["headers"]["Authorization"] == "Bearer explicit-token" + assert kwargs["context"].service_parameters["Authorization"] == "Bearer explicit-token" @pytest.mark.asyncio @@ -405,7 +405,8 @@ async def test_send_message_negotiates_extensions_via_service_parameters( assert len(result) == 1 payload, _, kwargs = fake_client.send_message_inputs[0] assert kwargs["context"].service_parameters == { - "A2A-Extensions": "https://example.com/ext-a,https://example.com/ext-b" + "A2A-Version": "1.0", + "A2A-Extensions": "https://example.com/ext-a,https://example.com/ext-b", } @@ -428,7 +429,7 @@ async def test_send_message_prefers_explicit_authorization_without_default_token assert result[0].HasField("message") request, _, kwargs = fake_client.send_message_inputs[0] assert request.metadata == {} - assert kwargs["context"].state["headers"]["Authorization"] == "Bearer explicit-token" + assert kwargs["context"].service_parameters["Authorization"] == "Bearer explicit-token" @pytest.mark.asyncio @@ -532,7 +533,7 @@ async def test_get_task_uses_authorization_header_context( params, kwargs = fake_client.task_inputs[0] assert params.id == "task-id" - assert kwargs["context"].state["headers"]["Authorization"] == "Bearer explicit-token" + assert kwargs["context"].service_parameters["Authorization"] == "Bearer explicit-token" assert "request_metadata" not in kwargs @@ -551,7 +552,8 @@ async def test_get_task_negotiates_extensions_via_service_parameters( _params, kwargs = fake_client.task_inputs[0] assert kwargs["context"].service_parameters == { - "A2A-Extensions": "https://example.com/ext-a,https://example.com/ext-b" + "A2A-Version": "1.0", + "A2A-Extensions": "https://example.com/ext-a,https://example.com/ext-b", } @@ -570,7 +572,7 @@ async def test_cancel_task_uses_authorization_header_context( params, kwargs = fake_client.cancel_inputs[0] assert params.metadata == {"trace_id": "trace-1"} - assert kwargs["context"].state["headers"]["Authorization"] == "Bearer explicit-token" + assert kwargs["context"].service_parameters["Authorization"] == "Bearer explicit-token" @pytest.mark.asyncio @@ -621,7 +623,7 @@ async def test_subscribe_to_task_uses_authorization_header_context( assert [event.task.status.state for event in result] == [TaskState.TASK_STATE_WORKING] params, kwargs = fake_client.subscribe_inputs[0] assert params.id == "task-id" - assert kwargs["context"].state["headers"]["Authorization"] == "Bearer explicit-token" + assert kwargs["context"].service_parameters["Authorization"] == "Bearer explicit-token" assert "request_metadata" not in kwargs diff --git a/tests/client/test_client_http_headers.py b/tests/client/test_client_http_headers.py new file mode 100644 index 0000000..5d2e145 --- /dev/null +++ b/tests/client/test_client_http_headers.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +import json +from base64 import b64encode + +import httpx +import pytest +from a2a.types import ( + AgentCapabilities, + AgentCard, + AgentInterface, + Message, + Part, + Role, + SendMessageResponse, + Task, + TaskState, + TaskStatus, +) +from google.protobuf.json_format import MessageToDict + +from opencode_a2a.client import A2AClient +from opencode_a2a.client.config import A2AClientSettings + +_PEER_URL = "https://peer.example.com" +_TRACEPARENT = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" + + +def _agent_card(protocol_binding: str) -> AgentCard: + return AgentCard( + name="HTTP stub peer", + description="Exercises the real SDK HTTP transports.", + version="1.0", + supported_interfaces=[ + AgentInterface( + url=f"{_PEER_URL}/", + protocol_binding=protocol_binding, + protocol_version="1.0", + ) + ], + capabilities=AgentCapabilities(streaming=False), + default_input_modes=["text/plain"], + default_output_modes=["text/plain"], + skills=[], + ) + + +def _http_stub( + requests: dict[str, httpx.Request], + protocol_binding: str, +) -> httpx.MockTransport: + card = _agent_card(protocol_binding) + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path == "/.well-known/agent-card.json": + requests["GetAgentCard"] = request + return httpx.Response(200, json=MessageToDict(card)) + + payload = json.loads(request.content) if request.content else None + if protocol_binding == "JSONRPC": + method = payload["method"] + elif request.url.path == "/message:send": + method = "SendMessage" + elif request.url.path == "/tasks/task-1": + method = "GetTask" + else: # pragma: no cover - keeps unexpected SDK calls visible + raise AssertionError(f"Unexpected REST request: {request.method} {request.url}") + requests[method] = request + if method == "SendMessage": + result = MessageToDict( + SendMessageResponse( + message=Message( + message_id="reply-1", + role=Role.ROLE_AGENT, + parts=[Part(text="ok")], + ) + ) + ) + elif method == "GetTask": + result = MessageToDict( + Task( + id="task-1", + context_id="context-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + ) + else: # pragma: no cover - keeps unexpected SDK calls visible + raise AssertionError(f"Unexpected JSON-RPC method: {method}") + if protocol_binding == "JSONRPC": + result = {"jsonrpc": "2.0", "id": payload["id"], "result": result} + return httpx.Response(200, json=result) + + return httpx.MockTransport(handle) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "protocol_binding", + ("JSONRPC", "HTTP+JSON"), + ids=("jsonrpc", "rest"), +) +@pytest.mark.parametrize( + ("bearer_token", "basic_auth", "expected_authorization"), + [ + ("peer-token", None, "Bearer peer-token"), + (None, "user:pass", f"Basic {b64encode(b'user:pass').decode()}"), + ], + ids=("bearer", "basic"), +) +async def test_sdk_http_requests_send_configured_headers( + protocol_binding: str, + bearer_token: str | None, + basic_auth: str | None, + expected_authorization: str, +) -> None: + requests: dict[str, httpx.Request] = {} + settings = A2AClientSettings( + bearer_token=bearer_token, + basic_auth=basic_auth, + supported_transports=(protocol_binding,), + ) + async with httpx.AsyncClient(transport=_http_stub(requests, protocol_binding)) as http_client: + client = A2AClient(_PEER_URL, settings=settings, httpx_client=http_client) + + await client.send("hello", metadata={"traceparent": _TRACEPARENT}) + await client.get_task("task-1", extensions=["https://example.com/ext"]) + + for method in ("GetAgentCard", "SendMessage", "GetTask"): + assert requests[method].headers["Authorization"] == expected_authorization + assert requests[method].headers["A2A-Version"] == "1.0" + assert requests["SendMessage"].headers["traceparent"] == _TRACEPARENT + assert requests["GetTask"].headers["A2A-Extensions"] == "https://example.com/ext" + + +@pytest.mark.asyncio +async def test_sdk_jsonrpc_request_preserves_explicit_authorization_override() -> None: + requests: dict[str, httpx.Request] = {} + settings = A2AClientSettings( + bearer_token="default-token", + supported_transports=("JSONRPC",), + ) + async with httpx.AsyncClient(transport=_http_stub(requests, "JSONRPC")) as http_client: + client = A2AClient(_PEER_URL, settings=settings, httpx_client=http_client) + + await client.send( + "hello", + metadata={"authorization": "Bearer explicit-token"}, + ) + + assert requests["SendMessage"].headers["Authorization"] == "Bearer explicit-token" diff --git a/tests/client/test_request_context.py b/tests/client/test_request_context.py index 1d0ea67..f85d103 100644 --- a/tests/client/test_request_context.py +++ b/tests/client/test_request_context.py @@ -81,8 +81,8 @@ def test_build_call_context_includes_fixed_protocol_version() -> None: context = build_call_context(None, None) assert context is not None - assert context.state["headers"] == {"A2A-Version": "1.0"} - assert context.state["http_kwargs"]["headers"] == {"A2A-Version": "1.0"} + assert context.service_parameters == {"A2A-Version": "1.0"} + assert context.state == {} def test_build_call_context_includes_current_trace_headers() -> None: @@ -96,7 +96,7 @@ def test_build_call_context_includes_current_trace_headers() -> None: context = build_call_context(None, None) assert context is not None - assert context.state["headers"] == { + assert context.service_parameters == { "A2A-Version": "1.0", "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", "tracestate": "vendor=value", @@ -120,7 +120,7 @@ def test_build_call_context_preserves_explicit_trace_headers_over_current_contex ) assert context is not None - assert context.state["headers"] == { + assert context.service_parameters == { "Authorization": "Bearer peer-token", "A2A-Version": "1.0", "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", @@ -128,16 +128,11 @@ def test_build_call_context_preserves_explicit_trace_headers_over_current_contex } -def test_build_call_context_carries_default_headers_without_interceptor_layer() -> None: +def test_build_call_context_carries_headers_via_service_parameters() -> None: context = build_call_context("peer-token", {"X-Trace-Id": "trace-1"}) assert isinstance(context, ClientCallContext) - assert context.state["headers"] == { - "Authorization": "Bearer peer-token", - "A2A-Version": "1.0", - "X-Trace-Id": "trace-1", - } - assert context.state["http_kwargs"]["headers"] == { + assert context.service_parameters == { "Authorization": "Bearer peer-token", "A2A-Version": "1.0", "X-Trace-Id": "trace-1", @@ -153,7 +148,10 @@ def test_build_call_context_merges_extension_service_parameters() -> None: assert isinstance(context, ClientCallContext) assert context.service_parameters == { - "A2A-Extensions": "https://example.com/ext-a,https://example.com/ext-b" + "Authorization": "Bearer peer-token", + "A2A-Version": "1.0", + "X-Trace-Id": "trace-1", + "A2A-Extensions": "https://example.com/ext-a,https://example.com/ext-b", } diff --git a/tests/execution/test_opencode_agent_session_binding.py b/tests/execution/test_opencode_agent_session_binding.py index 4e90de5..88ca690 100644 --- a/tests/execution/test_opencode_agent_session_binding.py +++ b/tests/execution/test_opencode_agent_session_binding.py @@ -588,7 +588,7 @@ async def test_agent_a2a_call_uses_server_side_basic_auth_headers( assert results[0]["output"] == "remote response" _, _, kwargs = fake_sdk_client.send_message_inputs[0] assert kwargs["context"] is not None - assert kwargs["context"].state["headers"]["Authorization"] == ( + assert kwargs["context"].service_parameters["Authorization"] == ( f"Basic {b64encode(b'user:pass').decode()}" ) @@ -658,10 +658,10 @@ async def test_agent_a2a_call_propagates_current_trace_headers( assert results[0]["output"] == "remote response" _, _, kwargs = fake_sdk_client.send_message_inputs[0] assert kwargs["context"] is not None - assert kwargs["context"].state["headers"]["traceparent"] == ( + assert kwargs["context"].service_parameters["traceparent"] == ( "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" ) - assert kwargs["context"].state["headers"]["tracestate"] == "vendor=value" + assert kwargs["context"].service_parameters["tracestate"] == "vendor=value" await manager.close_all()