From 1b094d516101bbd4ca23bdede58ab2b3809f0d80 Mon Sep 17 00:00:00 2001 From: hutiefang76 <137664623+hutiefang76@users.noreply.github.com> Date: Sun, 4 Oct 2026 14:50:03 +0800 Subject: [PATCH] fix: cancel in-flight HTTP futures on REST transport close Track active REST streaming futures in both protocol implementations, remove completed futures, and cancel pending futures when closing. Cover completion, shutdown races, idempotence, and JDK socket release. Related: #1194 --- .../client/transport/rest/RestTransport.java | 46 +++- .../rest/RestTransportCloseTest.java | 214 ++++++++++++++++++ .../transport/rest/RestTransport_v0_3.java | 46 +++- .../rest/RestTransportClose_v0_3_Test.java | 212 +++++++++++++++++ 4 files changed, 508 insertions(+), 10 deletions(-) create mode 100644 client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportCloseTest.java create mode 100644 compat-0.3/client/transport/rest/src/test/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransportClose_v0_3_Test.java diff --git a/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java b/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java index 404ab7a3b..522a8b8f7 100644 --- a/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java +++ b/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java @@ -16,9 +16,12 @@ import java.io.IOException; import java.net.URLEncoder; import java.nio.charset.StandardCharsets; +import java.util.ArrayList; import java.util.Collections; +import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; @@ -67,6 +70,8 @@ public class RestTransport implements ClientTransport { private static final Logger LOGGER = Logger.getLogger(RestTransport.class.getName()); private final A2AHttpClient httpClient; + private final Set> activeStreams = new HashSet<>(); + private boolean closed; private final AgentInterface agentInterface; private @Nullable final List interceptors; private final AgentCard agentCard; @@ -117,12 +122,12 @@ public void sendMessageStreaming(MessageSendParams messageSendParams, Consumer sseEventListener.onMessage(event, ref.get()), throwable -> sseEventListener.onError(throwable, ref.get()), () -> { // We don't need to do anything special on completion - })); + }))); } catch (IOException e) { throw new A2AClientException("Failed to send streaming message request: " + e, e); } catch (InterruptedException e) { @@ -377,12 +382,12 @@ public void subscribeToTask(TaskIdParams request, Consumer e try { String url = Utils.buildBaseUrl(agentInterface, request.tenant()) + String.format("/tasks/%1s:subscribe", request.id()); A2AHttpClient.PostBuilder postBuilder = createPostBuilder(url, payloadAndHeaders); - ref.set(postBuilder.postAsyncSSE( + ref.set(trackStream(postBuilder.postAsyncSSE( event -> sseEventListener.onMessage(event, ref.get()), throwable -> sseEventListener.onError(throwable, ref.get()), () -> { // We don't need to do anything special on completion - })); + }))); } catch (IOException e) { throw new A2AClientException("Failed to send streaming message request: " + e, e); } catch (InterruptedException e) { @@ -414,7 +419,38 @@ public AgentCard getExtendedAgentCard(GetExtendedAgentCardParams params, @Nullab @Override public void close() { - // no-op + List> streams; + synchronized (activeStreams) { + if (closed) { + return; + } + closed = true; + streams = new ArrayList<>(activeStreams); + activeStreams.clear(); + } + for (CompletableFuture stream : streams) { + stream.cancel(true); + } + } + + private CompletableFuture trackStream(CompletableFuture stream) { + boolean cancel; + synchronized (activeStreams) { + cancel = closed; + if (!cancel) { + activeStreams.add(stream); + } + } + stream.whenComplete((result, error) -> { + synchronized (activeStreams) { + activeStreams.remove(stream); + } + }); + if (cancel) { + // close() can run while the HTTP client is setting up the stream. + stream.cancel(true); + } + return stream; } private PayloadAndHeaders applyInterceptors(String methodName, @Nullable MessageOrBuilder payload, diff --git a/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportCloseTest.java b/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportCloseTest.java new file mode 100644 index 000000000..bb00033fc --- /dev/null +++ b/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportCloseTest.java @@ -0,0 +1,214 @@ +package org.a2aproject.sdk.client.transport.rest; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.BufferedReader; +import java.io.IOException; +import java.io.InputStreamReader; +import java.net.InetAddress; +import java.net.ServerSocket; +import java.net.Socket; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; + +import org.a2aproject.sdk.client.http.A2AHttpClient; +import org.a2aproject.sdk.client.http.A2AHttpResponse; +import org.a2aproject.sdk.client.http.JdkA2AHttpClient; +import org.a2aproject.sdk.client.http.ServerSentEvent; +import org.a2aproject.sdk.spec.AgentCapabilities; +import org.a2aproject.sdk.spec.AgentCard; +import org.a2aproject.sdk.spec.AgentInterface; +import org.a2aproject.sdk.spec.Message; +import org.a2aproject.sdk.spec.MessageSendParams; +import org.a2aproject.sdk.spec.StreamingEventKind; +import org.a2aproject.sdk.spec.TaskIdParams; +import org.a2aproject.sdk.spec.TextPart; +import org.a2aproject.sdk.spec.TransportProtocol; +import org.junit.jupiter.api.Test; + +class RestTransportCloseTest { + + private static final AgentInterface INTERFACE = new AgentInterface(TransportProtocol.HTTP_JSON.asString(), "http://localhost:4001"); + private static final AgentCard CARD = AgentCard.builder().name("Test Agent").description("Test Agent") + .version("1.0").supportedInterfaces(List.of(INTERFACE)) + .capabilities(AgentCapabilities.builder().streaming(true).build()) + .defaultInputModes(List.of("text")).defaultOutputModes(List.of("text")).skills(List.of()).build(); + + @Test + void closeCancelsAllActiveStreamsAndIsIdempotent() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport transport = new RestTransport(http, CARD, INTERFACE, null); + sendStreaming(transport, event -> {}); + transport.subscribeToTask(new TaskIdParams("task-1234"), event -> {}, error -> {}, null); + + transport.close(); + transport.close(); + + assertEquals(2, http.requests.size()); + for (CountingFuture request : http.requests) { + assertTrue(request.isCancelled(), "Closing REST must cancel each active HTTP stream"); + assertEquals(1, request.cancellations, "Repeated close must not cancel a request again"); + } + } + + @Test + void closeDoesNotCancelCompletedStreams() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport transport = new RestTransport(http, CARD, INTERFACE, null); + sendStreaming(transport, event -> {}); + CountingFuture request = http.requests.get(0); + request.complete(null); + + transport.close(); + + assertEquals(0, request.cancellations, "Completed streams must be removed from the active requests"); + } + + @Test + void closeDoesNotRetainStreamsCompletedDuringSetup() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport transport = new RestTransport(http, CARD, INTERFACE, null); + http.duringSetup = () -> http.requests.get(0).complete(null); + + sendStreaming(transport, event -> {}); + transport.close(); + + assertEquals(0, http.requests.get(0).cancellations, "Synchronous completion must not leave a tracked stream"); + } + + @Test + void closeDuringRequestSetupCancelsTheLateRegisteredStream() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport transport = new RestTransport(http, CARD, INTERFACE, null); + http.duringSetup = () -> CompletableFuture.runAsync(transport::close).orTimeout(5, TimeUnit.SECONDS).join(); + + sendStreaming(transport, event -> {}); + + assertTrue(http.requests.get(0).isCancelled(), "Close must not lose a future returned after close began"); + } + + @Test + void closeReleasesTheRealJdkHttpConnection() throws Exception { + ExecutorService executor = Executors.newSingleThreadExecutor(); + AtomicReference accepted = new AtomicReference<>(); + try (ServerSocket server = new ServerSocket(0, 1, InetAddress.getByName("127.0.0.1"))) { + server.setSoTimeout(10000); + Future connectionEnd = executor.submit(() -> { + try (Socket socket = server.accept()) { + accepted.set(socket); + socket.setSoTimeout(10000); + BufferedReader input = new BufferedReader(new InputStreamReader(socket.getInputStream(), StandardCharsets.UTF_8)); + int contentLength = 0; + for (String line; (line = input.readLine()) != null && !line.isEmpty();) { + if (line.toLowerCase(Locale.ROOT).startsWith("content-length:")) { + contentLength = Integer.parseInt(line.substring(line.indexOf(':') + 1).trim()); + } + } + for (int i = 0; i < contentLength; i++) { + if (input.read() == -1) { + throw new IOException("Request ended before its body"); + } + } + String event = "data: {\"task\":{\"id\":\"task-1234\",\"contextId\":\"context-1234\"," + + "\"status\":{\"state\":\"TASK_STATE_WORKING\"}}}\n\n"; + byte[] data = event.getBytes(StandardCharsets.UTF_8); + String headers = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\n\r\n"; + socket.getOutputStream().write(headers.getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().write((Integer.toHexString(data.length) + "\r\n").getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().write(data); + socket.getOutputStream().write("\r\n".getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().flush(); + // Keep the response open. EOF here proves the client released the upstream socket. + return input.read(); + } + }); + AgentInterface endpoint = new AgentInterface(TransportProtocol.HTTP_JSON.asString(), "http://127.0.0.1:" + server.getLocalPort()); + RestTransport transport = new RestTransport(new JdkA2AHttpClient(), CARD, endpoint, null); + CountDownLatch received = new CountDownLatch(1); + try { + sendStreaming(transport, event -> received.countDown()); + assertTrue(received.await(5, TimeUnit.SECONDS), "The stream must be active before close"); + + transport.close(); + + assertEquals(-1, connectionEnd.get(5, TimeUnit.SECONDS), "Close must terminate the real HTTP connection"); + } finally { + transport.close(); + Socket socket = accepted.get(); + if (socket != null) { + socket.close(); + } + } + } finally { + executor.shutdownNow(); + } + } + + private static void sendStreaming(RestTransport transport, Consumer consumer) + throws Exception { + Message message = Message.builder().role(Message.Role.ROLE_USER).messageId("message-1234") + .parts(List.of(new TextPart("hello"))).build(); + transport.sendMessageStreaming(MessageSendParams.builder().message(message).build(), consumer, error -> {}, null); + } + + private static class CountingFuture extends CompletableFuture { + private int cancellations; + + @Override + public boolean cancel(boolean mayInterruptIfRunning) { + cancellations++; + return super.cancel(mayInterruptIfRunning); + } + } + + private static class PendingHttpClient extends JdkA2AHttpClient { + private final List requests = new ArrayList<>(); + private Runnable duringSetup = () -> {}; + + @Override + public PostBuilder createPost() { + return new A2AHttpClient.PostBuilder() { + @Override + public PostBuilder url(String url) { + return this; + } + @Override + public PostBuilder addHeaders(Map headers) { + return this; + } + @Override + public PostBuilder addHeader(String name, String value) { + return this; + } + @Override + public PostBuilder body(String body) { + return this; + } + @Override + public A2AHttpResponse post() { + throw new UnsupportedOperationException(); + } + @Override + public CompletableFuture postAsyncSSE(Consumer messages, + Consumer errors, Runnable complete) { + CountingFuture request = new CountingFuture(); + requests.add(request); + duringSetup.run(); + return request; + } + }; + } + } +} diff --git a/compat-0.3/client/transport/rest/src/main/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransport_v0_3.java b/compat-0.3/client/transport/rest/src/main/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransport_v0_3.java index 7a8e6a1ec..f34a3dcfc 100644 --- a/compat-0.3/client/transport/rest/src/main/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransport_v0_3.java +++ b/compat-0.3/client/transport/rest/src/main/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransport_v0_3.java @@ -48,11 +48,14 @@ import org.a2aproject.sdk.compat03.spec.SetTaskPushNotificationConfigRequest_v0_3; import org.a2aproject.sdk.compat03.json.JsonUtil_v0_3; import java.io.IOException; +import java.util.ArrayList; import java.util.Collections; +import java.util.HashSet; import java.util.List; import java.util.logging.Level; import java.util.logging.Logger; import java.util.Map; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; @@ -63,6 +66,8 @@ public class RestTransport_v0_3 implements ClientTransport_v0_3 { private static final Logger LOGGER = Logger.getLogger(RestTransport_v0_3.class.getName()); private final A2AHttpClient httpClient; + private final Set> activeStreams = new HashSet<>(); + private boolean closed; private final String agentUrl; private @Nullable final List interceptors; private AgentCard_v0_3 agentCard; @@ -115,12 +120,12 @@ public void sendMessageStreaming(MessageSendParams_v0_3 messageSendParams, Consu RestSSEEventListener_v0_3 sseEventListener = new RestSSEEventListener_v0_3(eventConsumer, errorConsumer); try { A2AHttpClient.PostBuilder postBuilder = createPostBuilder(agentUrl + "/v1/message:stream", payloadAndHeaders); - ref.set(postBuilder.postAsyncSSE( + ref.set(trackStream(postBuilder.postAsyncSSE( event -> sseEventListener.onMessage(event.data(), ref.get()), throwable -> sseEventListener.onError(throwable, ref.get()), () -> { // We don't need to do anything special on completion - })); + }))); } catch (IOException e) { throw new A2AClientException_v0_3("Failed to send streaming message request: " + e, e); } catch (InterruptedException e) { @@ -306,12 +311,12 @@ public void resubscribe(TaskIdParams_v0_3 request, Consumer sseEventListener.onMessage(event.data(), ref.get()), throwable -> sseEventListener.onError(throwable, ref.get()), () -> { // We don't need to do anything special on completion - })); + }))); } catch (IOException e) { throw new A2AClientException_v0_3("Failed to send streaming message request: " + e, e); } catch (InterruptedException e) { @@ -359,7 +364,38 @@ public AgentCard_v0_3 getAgentCard(@Nullable ClientCallContext_v0_3 context) thr @Override public void close() { - // no-op + List> streams; + synchronized (activeStreams) { + if (closed) { + return; + } + closed = true; + streams = new ArrayList<>(activeStreams); + activeStreams.clear(); + } + for (CompletableFuture stream : streams) { + stream.cancel(true); + } + } + + private CompletableFuture trackStream(CompletableFuture stream) { + boolean cancel; + synchronized (activeStreams) { + cancel = closed; + if (!cancel) { + activeStreams.add(stream); + } + } + stream.whenComplete((result, error) -> { + synchronized (activeStreams) { + activeStreams.remove(stream); + } + }); + if (cancel) { + // close() can run while the HTTP client is setting up the stream. + stream.cancel(true); + } + return stream; } private PayloadAndHeaders_v0_3 applyInterceptors(String methodName, @Nullable MessageOrBuilder payload, diff --git a/compat-0.3/client/transport/rest/src/test/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransportClose_v0_3_Test.java b/compat-0.3/client/transport/rest/src/test/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransportClose_v0_3_Test.java new file mode 100644 index 000000000..3d43caa3c --- /dev/null +++ b/compat-0.3/client/transport/rest/src/test/java/org/a2aproject/sdk/compat03/client/transport/rest/RestTransportClose_v0_3_Test.java @@ -0,0 +1,212 @@ +package org.a2aproject.sdk.compat03.client.transport.rest; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.BufferedReader; +import java.io.IOException; +import java.io.InputStreamReader; +import java.net.InetAddress; +import java.net.ServerSocket; +import java.net.Socket; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; + +import org.a2aproject.sdk.client.http.A2AHttpClient; +import org.a2aproject.sdk.client.http.A2AHttpResponse; +import org.a2aproject.sdk.client.http.JdkA2AHttpClient; +import org.a2aproject.sdk.client.http.ServerSentEvent; +import org.a2aproject.sdk.compat03.spec.AgentCapabilities_v0_3; +import org.a2aproject.sdk.compat03.spec.AgentCard_v0_3; +import org.a2aproject.sdk.compat03.spec.MessageSendParams_v0_3; +import org.a2aproject.sdk.compat03.spec.Message_v0_3; +import org.a2aproject.sdk.compat03.spec.StreamingEventKind_v0_3; +import org.a2aproject.sdk.compat03.spec.TaskIdParams_v0_3; +import org.a2aproject.sdk.compat03.spec.TextPart_v0_3; +import org.junit.jupiter.api.Test; + +class RestTransportClose_v0_3_Test { + + private static final String INTERFACE = "http://localhost:4001"; + private static final AgentCard_v0_3 CARD = new AgentCard_v0_3.Builder().name("Test Agent").description("Test Agent") + .version("1.0").url(INTERFACE) + .capabilities(new AgentCapabilities_v0_3.Builder().streaming(true).build()) + .defaultInputModes(List.of("text")).defaultOutputModes(List.of("text")).skills(List.of()).build(); + + @Test + void closeCancelsAllActiveStreamsAndIsIdempotent() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport_v0_3 transport = new RestTransport_v0_3(http, CARD, INTERFACE, null); + sendStreaming(transport, event -> {}); + transport.resubscribe(new TaskIdParams_v0_3("task-1234"), event -> {}, error -> {}, null); + + transport.close(); + transport.close(); + + assertEquals(2, http.requests.size()); + for (CountingFuture request : http.requests) { + assertTrue(request.isCancelled(), "Closing REST must cancel each active HTTP stream"); + assertEquals(1, request.cancellations, "Repeated close must not cancel a request again"); + } + } + + @Test + void closeDoesNotCancelCompletedStreams() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport_v0_3 transport = new RestTransport_v0_3(http, CARD, INTERFACE, null); + sendStreaming(transport, event -> {}); + CountingFuture request = http.requests.get(0); + request.complete(null); + + transport.close(); + + assertEquals(0, request.cancellations, "Completed streams must be removed from the active requests"); + } + + @Test + void closeDoesNotRetainStreamsCompletedDuringSetup() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport_v0_3 transport = new RestTransport_v0_3(http, CARD, INTERFACE, null); + http.duringSetup = () -> http.requests.get(0).complete(null); + + sendStreaming(transport, event -> {}); + transport.close(); + + assertEquals(0, http.requests.get(0).cancellations, "Synchronous completion must not leave a tracked stream"); + } + + @Test + void closeDuringRequestSetupCancelsTheLateRegisteredStream() throws Exception { + PendingHttpClient http = new PendingHttpClient(); + RestTransport_v0_3 transport = new RestTransport_v0_3(http, CARD, INTERFACE, null); + http.duringSetup = () -> CompletableFuture.runAsync(transport::close).orTimeout(5, TimeUnit.SECONDS).join(); + + sendStreaming(transport, event -> {}); + + assertTrue(http.requests.get(0).isCancelled(), "Close must not lose a future returned after close began"); + } + + @Test + void closeReleasesTheRealJdkHttpConnection() throws Exception { + ExecutorService executor = Executors.newSingleThreadExecutor(); + AtomicReference accepted = new AtomicReference<>(); + try (ServerSocket server = new ServerSocket(0, 1, InetAddress.getByName("127.0.0.1"))) { + server.setSoTimeout(10000); + Future connectionEnd = executor.submit(() -> { + try (Socket socket = server.accept()) { + accepted.set(socket); + socket.setSoTimeout(10000); + BufferedReader input = new BufferedReader(new InputStreamReader(socket.getInputStream(), StandardCharsets.UTF_8)); + int contentLength = 0; + for (String line; (line = input.readLine()) != null && !line.isEmpty();) { + if (line.toLowerCase(Locale.ROOT).startsWith("content-length:")) { + contentLength = Integer.parseInt(line.substring(line.indexOf(':') + 1).trim()); + } + } + for (int i = 0; i < contentLength; i++) { + if (input.read() == -1) { + throw new IOException("Request ended before its body"); + } + } + String event = "data: {\"task\":{\"id\":\"task-1234\",\"contextId\":\"context-1234\"," + + "\"status\":{\"state\":\"TASK_STATE_WORKING\"}}}\n\n"; + byte[] data = event.getBytes(StandardCharsets.UTF_8); + String headers = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\n\r\n"; + socket.getOutputStream().write(headers.getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().write((Integer.toHexString(data.length) + "\r\n").getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().write(data); + socket.getOutputStream().write("\r\n".getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().flush(); + // Keep the response open. EOF here proves the client released the upstream socket. + return input.read(); + } + }); + String endpoint = "http://127.0.0.1:" + server.getLocalPort(); + RestTransport_v0_3 transport = new RestTransport_v0_3(new JdkA2AHttpClient(), CARD, endpoint, null); + CountDownLatch received = new CountDownLatch(1); + try { + sendStreaming(transport, event -> received.countDown()); + assertTrue(received.await(5, TimeUnit.SECONDS), "The stream must be active before close"); + + transport.close(); + + assertEquals(-1, connectionEnd.get(5, TimeUnit.SECONDS), "Close must terminate the real HTTP connection"); + } finally { + transport.close(); + Socket socket = accepted.get(); + if (socket != null) { + socket.close(); + } + } + } finally { + executor.shutdownNow(); + } + } + + private static void sendStreaming(RestTransport_v0_3 transport, Consumer consumer) + throws Exception { + Message_v0_3 message = new Message_v0_3.Builder().role(Message_v0_3.Role.USER).messageId("message-1234") + .parts(List.of(new TextPart_v0_3("hello"))).build(); + transport.sendMessageStreaming(new MessageSendParams_v0_3.Builder().message(message).build(), consumer, error -> {}, null); + } + + private static class CountingFuture extends CompletableFuture { + private int cancellations; + + @Override + public boolean cancel(boolean mayInterruptIfRunning) { + cancellations++; + return super.cancel(mayInterruptIfRunning); + } + } + + private static class PendingHttpClient extends JdkA2AHttpClient { + private final List requests = new ArrayList<>(); + private Runnable duringSetup = () -> {}; + + @Override + public PostBuilder createPost() { + return new A2AHttpClient.PostBuilder() { + @Override + public PostBuilder url(String url) { + return this; + } + @Override + public PostBuilder addHeaders(Map headers) { + return this; + } + @Override + public PostBuilder addHeader(String name, String value) { + return this; + } + @Override + public PostBuilder body(String body) { + return this; + } + @Override + public A2AHttpResponse post() { + throw new UnsupportedOperationException(); + } + @Override + public CompletableFuture postAsyncSSE(Consumer messages, + Consumer errors, Runnable complete) { + CountingFuture request = new CountingFuture(); + requests.add(request); + duringSetup.run(); + return request; + } + }; + } + } +}