diff --git a/CHANGELOG.md b/CHANGELOG.md index ffb0f3ae..88b151c6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -28,6 +28,21 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 Micronaut client configuration hard-coded the three as not advertised, so a Micronaut client could not advertise elicitation or boolean config options from its settings. +- **Headers on the remote client transports: `WebSocketAcpClientTransport.webSocketCustomizer(...)` + and `StreamableHttpAcpClientTransport.requestCustomizer(...)`.** The JDK's `HttpClient` has no + default headers, so neither transport could send an `Authorization` header, an API key or any + other header an agent's endpoint requires; an application had to write its own transport. + `goose serve`, for one, refuses every connection without its `X-Secret-Key`. The WebSocket + customizer receives the `WebSocket.Builder` of each connect attempt; the HTTP one receives the + `HttpRequest.Builder` of every request the transport sends (the cleartext probe, `initialize`, + each POST, every SSE stream it opens or reopens, and the closing `DELETE`), so a token that + expires is read again for each. The HTTP transport keeps the method, body, URI and its own + headers (Content-Type, Accept, Acp-Connection-Id, Acp-Session-Id): a customizer's values for + them are dropped. A customizer that throws, or sets a header the JDK restricts, fails that + connect or that request's Mono rather than the caller, and a failed connect or `initialize` may + be tried again. The HTTP transport now also builds each request inside its Mono, so a failure + building one is reported there too. + - **Transport constructors without a JSON mapper:** `StreamableHttpAcpClientTransport(URI)`, `WebSocketAcpClientTransport(URI)`, `StreamableHttpAcpAgentTransport(int port, AcpAgentFactory)` and `StreamableHttpAcpServlet(AcpAgentFactory)` use `AcpJsonMapper.createDefault()`, as the stdio diff --git a/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransport.java b/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransport.java index f4ac4e20..a5d020f5 100644 --- a/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransport.java +++ b/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransport.java @@ -61,7 +61,16 @@ * the server does not. The default client keeps cookies in a cookie manager of its own and * runs on a bounded pool of daemon threads; {@link StreamableHttpAcpClientTransportOptions} * sets its sizes and the number of SSE streams. Pass an {@link HttpClient} of your own for TLS, - * proxy or authentication settings. + * proxy or cookie settings. + * + *

An endpoint that requires authentication, such as one that expects an API key or a + * bearer token in a header, gets it through {@link #requestCustomizer}, which every request + * the transport sends passes through: + * + *

{@code
+ * var transport = new StreamableHttpAcpClientTransport(URI.create("https://agents.example.com/acp"))
+ *     .requestCustomizer(builder -> builder.header("Authorization", "Bearer " + tokens.current()));
+ * }
* *

{@link #closeGracefully()} closes the streams and sends {@code DELETE} for the connection, * waiting at most five seconds for the answer; {@link #close()} does the same and blocks for up @@ -201,6 +210,31 @@ private static HttpClientBundle createDefaultHttpClient(StreamableHttpAcpClientT return HttpClientBundle.createDefault(options); } + /** + * Customizes every HTTP request the transport sends, typically to add the headers the + * endpoint requires: an {@code Authorization} header, an API key, a tenant. The JDK's + * {@link HttpClient} has no default headers, so this is the only way to send one. It is + * applied to each request as it is built, the cleartext probe, {@code initialize}, every + * POST, every SSE stream (re)opened and the closing {@code DELETE}, so a header whose + * value changes, such as a token that expires, is read again for each one. It runs on the + * thread that sends the request; keep it quick, and refresh a token elsewhere. + * + *

The transport owns the method, the body, the URI and its own headers (Content-Type, + * Accept, Acp-Connection-Id, Acp-Session-Id): values the customizer sets for any of them + * are dropped or replaced, so a customizer cannot break the protocol by accident. Other + * settings, such as a per-request timeout, are kept. The JDK refuses restricted headers + * such as {@code Host} or {@code Connection}; setting one, or any exception the customizer + * throws, fails that request's Mono. Call it before the first message is sent. + * @param customizer applied to the builder of each request + * @return this transport + * @throws IllegalArgumentException if {@code customizer} is null + */ + public StreamableHttpAcpClientTransport requestCustomizer(Consumer customizer) { + Assert.notNull(customizer, "The request customizer can not be null"); + this.requests.requestCustomizer(customizer); + return this; + } + /** * {@inheritDoc} *

It contacts nothing: it registers the handler and completes at once. The connection @@ -249,17 +283,21 @@ private Mono initialize(AcpSchema.JSONRPCRequest request) { return Mono.error(new IllegalStateException("Transport is already initialized")); } - HttpRequest httpRequest; + String json; try { - httpRequest = requests.jsonPost(RouteScope.bootstrap(), jsonMapper.writeValueAsString(request)); + json = jsonMapper.writeValueAsString(request); } catch (IOException e) { initialized.set(false); return Mono.error(new AcpConnectionException("Failed to serialize initialize request", e)); } + // Built after the probe, so it carries the HTTP version the probe settled on; built + // inside defer, so a request customizer that throws fails this Mono rather than the + // caller. return requests.upgradeCleartextToHttp2() - .then(Mono.defer(() -> requests.sendAsync(requests.pinned(httpRequest), HttpResponse.BodyHandlers.ofString()))) + .then(Mono.defer(() -> requests.sendAsync(requests.jsonPost(RouteScope.bootstrap(), json), + HttpResponse.BodyHandlers.ofString()))) .flatMap(response -> readInitializeResponse(request, response)) .flatMap(responseMessage -> streams.openConnectionStream().then(inbound.emit(responseMessage))) .doOnError(error -> { diff --git a/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpRequests.java b/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpRequests.java index e7e87dd4..68d4505c 100644 --- a/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpRequests.java +++ b/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpRequests.java @@ -13,6 +13,7 @@ import java.nio.charset.StandardCharsets; import java.time.Duration; import java.util.Locale; +import java.util.Set; import java.util.concurrent.ArrayBlockingQueue; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutorService; @@ -20,6 +21,7 @@ import java.util.concurrent.ThreadFactory; import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.TimeUnit; +import java.util.function.Consumer; import com.agentclientprotocol.sdk.error.AcpConnectionException; import com.agentclientprotocol.sdk.util.AcpSchedulers; @@ -30,8 +32,8 @@ /** * The HTTP exchanges of a Streamable HTTP client connection: the scope headers, the HTTP - * version settled by the cleartext probe, and the POST, GET and DELETE requests with the - * status and content type each must answer with. Completion signals are delivered on a + * version settled by the cleartext probe, the application's request customizer, and the + * POST, GET and DELETE requests with the status and content type each must answer with. Completion signals are delivered on a * bounded executor of their own, never on the HTTP client's. */ final class StreamableHttpRequests { @@ -48,6 +50,14 @@ final class StreamableHttpRequests { static final Duration PROBE_TIMEOUT = Duration.ofSeconds(5); + /** + * The headers this transport owns, lower-cased. A request customizer's values for them are + * dropped, including on the bootstrap {@code initialize}, which carries no connection id + * of the transport's own for one set by the customizer to be replaced by. + */ + private static final Set PROTOCOL_HEADERS = Set.of("content-type", "accept", + HEADER_CONNECTION_ID.toLowerCase(Locale.ROOT), HEADER_SESSION_ID.toLowerCase(Locale.ROOT)); + /** An HTTP client and the executor this transport created for it, if any. */ record HttpClientBundle(HttpClient httpClient, @Nullable ExecutorService ownedExecutor) { @@ -75,6 +85,9 @@ static HttpClientBundle createDefault(StreamableHttpAcpClientTransportOptions op private volatile @Nullable String connectionId; + private volatile Consumer requestCustomizer = builder -> { + }; + /** * HTTP version pinned for every request after the cleartext probe, or {@code null} to * use the client's own setting. Set to HTTP/1.1 when an {@code http://} server does not @@ -106,6 +119,10 @@ static ThreadFactory daemonThreadFactory(String threadName) { }; } + void requestCustomizer(Consumer requestCustomizer) { + this.requestCustomizer = requestCustomizer; + } + @Nullable String connectionId() { return connectionId; } @@ -129,9 +146,10 @@ Mono upgradeCleartextToHttp2() { if (!"http".equalsIgnoreCase(endpointUri.getScheme()) || httpClient.version() != HttpClient.Version.HTTP_2) { return Mono.empty(); } - HttpRequest probe = HttpRequest.newBuilder(endpointUri).GET().build(); - // Cancelling the Mono on timeout cancels the HTTP exchange (see sendAsync). - return sendAsync(probe, HttpResponse.BodyHandlers.discarding()) + // Customized like every other request: a gateway in front of the agent may refuse the + // probe without the application's credentials. Cancelling the Mono on timeout cancels + // the HTTP exchange (see sendAsync). + return Mono.defer(() -> sendAsync(customized().GET().build(), HttpResponse.BodyHandlers.discarding())) .timeout(PROBE_TIMEOUT, AcpSchedulers.timeouts()) .doOnNext(response -> { logger.debug("Cleartext probe to {} negotiated {}", endpointUri, response.version()); @@ -148,17 +166,9 @@ Mono upgradeCleartextToHttp2() { }); } - /** Re-stamps a request built before the probe with the version the probe settled on. */ - HttpRequest pinned(HttpRequest request) { - HttpClient.Version version = this.pinnedVersion; - if (version == null) { - return request; - } - return HttpRequest.newBuilder(request, (name, value) -> true).version(version).build(); - } - + /** A customized builder for the endpoint, with the version the probe settled on. */ private HttpRequest.Builder newRequest() { - HttpRequest.Builder requestBuilder = HttpRequest.newBuilder(endpointUri); + HttpRequest.Builder requestBuilder = customized(); HttpClient.Version version = this.pinnedVersion; if (version != null) { requestBuilder.version(version); @@ -166,25 +176,46 @@ private HttpRequest.Builder newRequest() { return requestBuilder; } - /** A JSON POST of {@code json} in {@code scope}, not yet sent. */ + /** + * A builder for the endpoint carrying what the application's customizer set, less the + * protocol headers and with the endpoint's URI whatever the customizer did. The + * customizer runs on a builder of its own and the result is copied, because a builder + * can replace a header but never remove one. The transport sets the method, the body and + * its own headers afterwards. + */ + private HttpRequest.Builder customized() { + HttpRequest.Builder scratch = HttpRequest.newBuilder(endpointUri); + requestCustomizer.accept(scratch); + return HttpRequest + .newBuilder(scratch.build(), (name, value) -> !PROTOCOL_HEADERS.contains(name.toLowerCase(Locale.ROOT))) + .uri(endpointUri); + } + + /** + * A JSON POST of {@code json} in {@code scope}, not yet sent. Builds the request, so it + * throws whatever the request customizer throws; the Monos below build inside + * {@code defer} and fail with it instead. + */ HttpRequest jsonPost(RouteScope scope, String json) { - HttpRequest.Builder builder = newRequest().header("Content-Type", CONTENT_TYPE_JSON) - .header("Accept", CONTENT_TYPE_JSON); + HttpRequest.Builder builder = newRequest().setHeader("Content-Type", CONTENT_TYPE_JSON) + .setHeader("Accept", CONTENT_TYPE_JSON); addScopeHeaders(builder, scope); return builder.POST(HttpRequest.BodyPublishers.ofString(json, StandardCharsets.UTF_8)).build(); } /** Posts a message that the server must accept with 202. */ Mono postAccepted(RouteScope scope, String json) { - return sendAsync(jsonPost(scope, json), HttpResponse.BodyHandlers.discarding()) + return Mono.defer(() -> sendAsync(jsonPost(scope, json), HttpResponse.BodyHandlers.discarding())) .flatMap(response -> expectStatus(response, 202, "for POST")); } /** Opens the SSE stream of {@code scope}; emits its body or an error, never completes empty. */ Mono openEventStream(RouteScope scope) { - HttpRequest.Builder builder = newRequest().GET().header("Accept", CONTENT_TYPE_EVENT_STREAM); - addScopeHeaders(builder, scope); - return sendAsync(builder.build(), HttpResponse.BodyHandlers.ofInputStream()) + return Mono.defer(() -> { + HttpRequest.Builder builder = newRequest().GET().setHeader("Accept", CONTENT_TYPE_EVENT_STREAM); + addScopeHeaders(builder, scope); + return sendAsync(builder.build(), HttpResponse.BodyHandlers.ofInputStream()); + }) .flatMap(response -> expectStatus(response, 200, "when opening SSE stream") .then(expectContentType(response, CONTENT_TYPE_EVENT_STREAM, "response")) .then(Mono.fromSupplier(response::body))); @@ -192,8 +223,9 @@ Mono openEventStream(RouteScope scope) { /** Deletes the connection; the server must accept with 202. */ Mono deleteConnection(String id) { - HttpRequest request = newRequest().DELETE().header(HEADER_CONNECTION_ID, id).build(); - return sendAsync(request, HttpResponse.BodyHandlers.discarding()) + return Mono + .defer(() -> sendAsync(newRequest().DELETE().setHeader(HEADER_CONNECTION_ID, id).build(), + HttpResponse.BodyHandlers.discarding())) .flatMap(response -> expectStatus(response, 202, "for DELETE")); } @@ -215,10 +247,10 @@ static Mono expectContentType(HttpResponse response, String expected, S private void addScopeHeaders(HttpRequest.Builder builder, RouteScope scope) { if (!scope.isBootstrap()) { - builder.header(HEADER_CONNECTION_ID, requireConnectionId()); + builder.setHeader(HEADER_CONNECTION_ID, requireConnectionId()); } if (scope.isSession()) { - builder.header(HEADER_SESSION_ID, scope.boundSessionId()); + builder.setHeader(HEADER_SESSION_ID, scope.boundSessionId()); } } diff --git a/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransport.java b/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransport.java index 6f9be582..d72bea83 100644 --- a/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransport.java +++ b/acp-core/src/main/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransport.java @@ -56,6 +56,14 @@ * skipped. When the agent closes the connection, or it fails, {@link #awaitTermination()} ends * and the client's pending requests fail. * + *

An endpoint that requires authentication, such as one that expects an API key or a + * bearer token in a header, gets it through {@link #webSocketCustomizer}: + * + *

{@code
+ * var transport = new WebSocketAcpClientTransport(URI.create("wss://agents.example.com/acp"))
+ *     .webSocketCustomizer(builder -> builder.header("Authorization", "Bearer " + token));
+ * }
+ * *

The transport is thread-safe: messages may be sent from any thread, and one daemon thread * of its own ({@code acp-ws-client-outbound}) sends them one frame at a time. The default * HTTP client runs on a pool of daemon threads named {@code acp-ws-client}. @@ -108,6 +116,9 @@ public class WebSocketAcpClientTransport implements AcpClientTransport { private Duration connectTimeout = Duration.ofSeconds(30); + private Consumer webSocketCustomizer = builder -> { + }; + /** * Creates a transport for the WebSocket endpoint at {@code serverUri}, with an HTTP client * of its own and the default JSON mapper ({@link AcpJsonMapper#createDefault()}). @@ -186,6 +197,26 @@ public WebSocketAcpClientTransport connectTimeout(Duration timeout) { return this; } + /** + * Customizes the WebSocket handshake before it is sent, typically to add the headers the + * endpoint requires: an {@code Authorization} header, an API key, a tenant. The JDK's + * {@link java.net.http.HttpClient} has no default headers, so this is the only way to send + * one. Runs on every {@link #connect} attempt, after the transport has set its connect + * timeout, so the customizer may replace that too; a header whose value changes, such as + * a token that expires, is read again by each attempt. The JDK refuses headers that belong + * to the handshake itself ({@code Connection}, {@code Upgrade}, {@code Host}, + * {@code Sec-WebSocket-*}): setting one, or any exception the customizer throws, fails + * that connect, which may then be tried again. Call it before connecting. + * @param customizer applied to the builder of each handshake + * @return this transport + * @throws IllegalArgumentException if {@code customizer} is null + */ + public WebSocketAcpClientTransport webSocketCustomizer(Consumer customizer) { + Assert.notNull(customizer, "The WebSocket customizer can not be null"); + this.webSocketCustomizer = customizer; + return this; + } + /** * {@inheritDoc} *

Opens the WebSocket connection when the returned Mono is subscribed, and completes @@ -204,9 +235,9 @@ public Mono connect(Function, Mono> h // Build WebSocket connection with listener; frames that arrive before the handshake // completes wait in the inbound sink. - return httpClient.newWebSocketBuilder() - .connectTimeout(connectTimeout) - .buildAsync(serverUri, new AcpWebSocketListener()); + WebSocket.Builder builder = httpClient.newWebSocketBuilder().connectTimeout(connectTimeout); + webSocketCustomizer.accept(builder); + return builder.buildAsync(serverUri, new AcpWebSocketListener()); }).doOnSuccess(ws -> { this.webSocket = ws; // Only an open connection takes the inbound sink's one subscriber, so a connect that diff --git a/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransportTest.java b/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransportTest.java index 3e59af3e..0b7c5724 100644 --- a/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransportTest.java +++ b/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/StreamableHttpAcpClientTransportTest.java @@ -27,7 +27,9 @@ import java.util.concurrent.Future; import java.util.concurrent.LinkedBlockingQueue; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; import java.util.stream.Collectors; import com.agentclientprotocol.sdk.AcpTestFixtures; @@ -1278,6 +1280,124 @@ public void onComplete() { return bytes.toString(StandardCharsets.UTF_8); } + /** + * The JDK's HttpClient has no default headers, so before the request customizer there was + * no way to send an API key or a bearer token to an agent behind authentication. Every + * request passes through it, it is read again for each one (a refreshed token reaches the + * next request), and the transport's own headers replace any the customizer set. + */ + @Test + void theRequestCustomizerReachesEveryRequestWithoutOverridingProtocolHeaders() throws Exception { + HttpClient httpClient = mock(HttpClient.class); + List sent = new CopyOnWriteArrayList<>(); + PipedInputStream connectionStreamBody = new PipedInputStream(); + PipedOutputStream connectionStreamWriter = new PipedOutputStream(connectionStreamBody); + CountDownLatch connectionStreamOpened = new CountDownLatch(1); + when(httpClient.sendAsync(any(), any())).thenAnswer(invocation -> { + HttpRequest request = invocation.getArgument(0); + sent.add(request); + if ("POST".equals(request.method()) && sent.size() == 1) { + String initializeResponse = jsonMapper.writeValueAsString(AcpTestFixtures + .createJsonRpcResponse("init-1", AcpTestFixtures.createInitializeResponse())); + return CompletableFuture.completedFuture(response(200, + Map.of("Content-Type", "application/json", "Acp-Connection-Id", "conn-1"), initializeResponse)); + } + if ("GET".equals(request.method())) { + connectionStreamOpened.countDown(); + return CompletableFuture.completedFuture( + response(200, Map.of("Content-Type", "text/event-stream"), connectionStreamBody)); + } + return CompletableFuture.completedFuture(response(202, Map.of(), null)); + }); + AtomicReference token = new AtomicReference<>("first"); + StreamableHttpAcpClientTransport transport = new StreamableHttpAcpClientTransport( + URI.create("https://localhost:8443/acp"), jsonMapper, httpClient) + .requestCustomizer(builder -> builder.header("Authorization", "Bearer " + token.get()) + .header("Acp-Connection-Id", "forged") + .header("Accept", "text/plain")); + try { + transport.setExceptionHandler(error -> { + }); + transport.connect(message -> Mono.empty()).block(); + transport.sendMessage(AcpTestFixtures.createJsonRpcRequest(AcpSchema.METHOD_INITIALIZE, "init-1", + AcpTestFixtures.createInitializeRequest())) + .block(); + assertThat(connectionStreamOpened.await(1, TimeUnit.SECONDS)).isTrue(); + + token.set("second"); + transport.sendMessage(AcpTestFixtures.createJsonRpcRequest(AcpSchema.METHOD_SESSION_NEW, "new-1", + AcpTestFixtures.createNewSessionRequest())) + .block(); + } + finally { + transport.close(); + connectionStreamWriter.close(); + } + + assertThat(sent).extracting(HttpRequest::method).containsExactly("POST", "GET", "POST", "DELETE"); + assertThat(sent).allSatisfy(request -> assertThat(request.headers().allValues("Authorization")).hasSize(1)); + assertThat(sent.get(0).headers().allValues("Authorization")).containsExactly("Bearer first"); + assertThat(sent.get(2).headers().allValues("Authorization")).containsExactly("Bearer second"); + // The bootstrap POST has no connection yet; every later request names the real one. + assertThat(sent.get(0).headers().allValues("Acp-Connection-Id")).isEmpty(); + assertThat(sent.subList(1, 4)) + .allSatisfy(request -> assertThat(request.headers().allValues("Acp-Connection-Id")).containsExactly("conn-1")); + assertThat(sent.get(0).headers().allValues("Accept")).containsExactly("application/json"); + assertThat(sent.get(1).headers().allValues("Accept")).containsExactly("text/event-stream"); + } + + /** + * A customizer that throws, say because no token is available yet, fails the request's + * Mono instead of throwing out of {@code sendMessage}, and a failed {@code initialize} may + * be sent again. + */ + @Test + void aRequestCustomizerThatThrowsFailsTheMonoAndInitializeMayBeRetried() throws Exception { + HttpClient httpClient = mock(HttpClient.class); + String body = jsonMapper.writeValueAsString( + AcpTestFixtures.createJsonRpcResponse("init-1", AcpTestFixtures.createInitializeResponse())); + when(httpClient.sendAsync(any(), any())).thenAnswer(invocation -> { + HttpRequest request = invocation.getArgument(0); + if ("GET".equals(request.method())) { + return CompletableFuture.completedFuture( + response(200, Map.of("Content-Type", "text/event-stream"), emptyBody())); + } + return CompletableFuture.completedFuture( + response(200, Map.of("Content-Type", "application/json", "Acp-Connection-Id", "conn-1"), body)); + }); + AtomicBoolean signedIn = new AtomicBoolean(); + StreamableHttpAcpClientTransport transport = new StreamableHttpAcpClientTransport( + URI.create("https://localhost:8443/acp"), jsonMapper, httpClient) + .requestCustomizer(builder -> { + if (!signedIn.get()) { + throw new IllegalStateException("not signed in"); + } + builder.header("Authorization", "Bearer token"); + }); + transport.setExceptionHandler(error -> { + }); + AcpSchema.JSONRPCRequest initialize = AcpTestFixtures.createJsonRpcRequest(AcpSchema.METHOD_INITIALIZE, + "init-1", AcpTestFixtures.createInitializeRequest()); + try { + Mono first = transport.sendMessage(initialize); + assertThatThrownBy(first::block).isInstanceOf(IllegalStateException.class).hasMessage("not signed in"); + + signedIn.set(true); + transport.sendMessage(initialize).block(); + } + finally { + transport.close(); + } + } + + @Test + void requestCustomizerRejectsNull() { + StreamableHttpAcpClientTransport transport = new StreamableHttpAcpClientTransport( + URI.create("https://localhost:8443/acp"), jsonMapper, mock(HttpClient.class)); + + assertThatThrownBy(() -> transport.requestCustomizer(null)).isInstanceOf(IllegalArgumentException.class); + } + private InputStream emptyBody() { return new ByteArrayInputStream(new byte[0]); } diff --git a/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransportTest.java b/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransportTest.java index 40161c10..882e2a65 100644 --- a/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransportTest.java +++ b/acp-core/src/test/java/com/agentclientprotocol/sdk/client/transport/WebSocketAcpClientTransportTest.java @@ -4,10 +4,16 @@ package com.agentclientprotocol.sdk.client.transport; +import java.net.InetAddress; +import java.net.InetSocketAddress; import java.net.URI; import java.time.Duration; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.atomic.AtomicBoolean; import com.agentclientprotocol.sdk.json.AcpJsonMapper; +import com.sun.net.httpserver.HttpServer; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import reactor.core.publisher.Mono; @@ -117,4 +123,92 @@ void closeShutsDownTheExecutorOfItsOwnHttpClient() throws Exception { assertThat(executor.isShutdown()).isTrue(); } + /** + * The JDK's HttpClient has no default headers, so before the customizer there was no way + * to send an API key or a bearer token with the handshake. A real handshake against a + * server that records it and refuses it: the header is on the wire. + */ + @Test + void theWebSocketCustomizerAddsHeadersToTheHandshake() throws Exception { + List authorization = new CopyOnWriteArrayList<>(); + HttpServer server = HttpServer.create(new InetSocketAddress(InetAddress.getLoopbackAddress(), 0), 0); + server.createContext("/acp", exchange -> { + authorization.addAll(exchange.getRequestHeaders().getOrDefault("Authorization", List.of())); + exchange.sendResponseHeaders(401, -1); + exchange.close(); + }); + server.start(); + WebSocketAcpClientTransport transport = new WebSocketAcpClientTransport( + URI.create("ws://127.0.0.1:" + server.getAddress().getPort() + "/acp"), jsonMapper) + .webSocketCustomizer(builder -> builder.header("Authorization", "Bearer token")); + transport.setExceptionHandler(error -> { + }); + try { + assertThatThrownBy(() -> transport.connect(msg -> Mono.empty()).block(Duration.ofSeconds(10))) + .isNotNull(); + + assertThat(authorization).containsExactly("Bearer token"); + } + finally { + transport.closeGracefully().block(Duration.ofSeconds(10)); + server.stop(0); + } + } + + /** + * A customizer that throws, or sets a header the JDK reserves for the handshake, fails + * that connect with its own error, and the connect may be tried again. + */ + @Test + void aWebSocketCustomizerThatThrowsFailsTheConnectWhichMayBeRetried() { + AtomicBoolean signedIn = new AtomicBoolean(); + WebSocketAcpClientTransport transport = new WebSocketAcpClientTransport(URI.create("ws://127.0.0.1:1/acp"), + jsonMapper) + .webSocketCustomizer(builder -> { + if (!signedIn.get()) { + throw new IllegalStateException("not signed in"); + } + }); + transport.setExceptionHandler(error -> { + }); + try { + assertThatThrownBy(() -> transport.connect(msg -> Mono.empty()).block(Duration.ofSeconds(10))) + .isInstanceOf(IllegalStateException.class) + .hasMessage("not signed in"); + + signedIn.set(true); + // Nothing listens on port 1, so this attempt fails too, but by trying to connect. + assertThatThrownBy(() -> transport.connect(msg -> Mono.empty()).block(Duration.ofSeconds(10))) + .satisfies(error -> assertThat(error.getMessage()).doesNotContain("Already connected") + .doesNotContain("not signed in")); + } + finally { + transport.closeGracefully().block(Duration.ofSeconds(10)); + } + } + + @Test + void aReservedHandshakeHeaderFailsTheConnect() { + WebSocketAcpClientTransport transport = new WebSocketAcpClientTransport(URI.create("ws://127.0.0.1:1/acp"), + jsonMapper) + .webSocketCustomizer(builder -> builder.header("Sec-WebSocket-Key", "forged")); + transport.setExceptionHandler(error -> { + }); + try { + assertThatThrownBy(() -> transport.connect(msg -> Mono.empty()).block(Duration.ofSeconds(10))) + .isInstanceOf(IllegalArgumentException.class); + } + finally { + transport.closeGracefully().block(Duration.ofSeconds(10)); + } + } + + @Test + void webSocketCustomizerRejectsNull() { + WebSocketAcpClientTransport transport = new WebSocketAcpClientTransport(URI.create("ws://127.0.0.1:1/acp"), + jsonMapper); + + assertThatThrownBy(() -> transport.webSocketCustomizer(null)).isInstanceOf(IllegalArgumentException.class); + } + }