From 618ff855ad4fd93a476b3504e8fe689ad2f50e60 Mon Sep 17 00:00:00 2001 From: NewPeople-star <232190515+NewPeople-star@users.noreply.github.com> Date: Wed, 30 Sep 2026 11:11:32 +0800 Subject: [PATCH 1/3] Fix Streamable HTTP: surface invalid JSON response as transport error In HttpClientStreamableHttpTransport.sendMessage, the application/json branch completed the delivery sink before deserializing the payload. When the server returned a body that is not valid JSON, the parse failure only travelled through the response Flux; the delivery sink had already completed, so McpClientSession.sendRequest never received the error, never removed its pending response entry, and the caller waited for the full request timeout and only saw a TimeoutException. The original parsing exception was missing from the terminal chain. Complete the sink only after the payload has been parsed successfully. Notifications keep completing before their early return since they have no response to parse. Fixes #1147 --- .../HttpClientStreamableHttpTransport.java | 15 +++- ...eHttpTransportInvalidJsonResponseTest.java | 86 +++++++++++++++++++ 2 files changed, 99 insertions(+), 2 deletions(-) create mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportInvalidJsonResponseTest.java diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java index 5517823b6..31bd4c3a6 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java @@ -645,16 +645,27 @@ else if (contentType.contains(TEXT_EVENT_STREAM)) { }); } else if (contentType.contains(APPLICATION_JSON)) { - deliveredSink.success(); String data = ((ResponseSubscribers.AggregateResponseEvent) responseEvent).data(); if (sentMessage instanceof McpSchema.JSONRPCNotification) { logger.warn("Notification: {} received non-compliant response: {}", sentMessage, Utils.hasText(data) ? data : "[empty]"); + deliveredSink.success(); return Mono.empty(); } try { - return Mono.just(McpSchema.deserializeJsonRpcMessage(jsonMapper, data)); + McpSchema.JSONRPCMessage message = McpSchema.deserializeJsonRpcMessage(jsonMapper, data); + // Signal delivery only after the payload has been parsed + // successfully. + // Completing the sink before deserialization would swallow a + // parse + // failure: McpClientSession relies on the error signal to + // remove the + // pending response, and without it the caller waits for the + // full + // request timeout and only sees a TimeoutException. + deliveredSink.success(); + return Mono.just(message); } catch (IOException e) { return Mono.error(new McpTransportException( diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportInvalidJsonResponseTest.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportInvalidJsonResponseTest.java new file mode 100644 index 000000000..c1b988624 --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportInvalidJsonResponseTest.java @@ -0,0 +1,86 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.IOException; +import java.net.InetSocketAddress; +import java.time.Duration; + +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; + +import com.sun.net.httpserver.HttpServer; + +import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.McpTransportException; +import io.modelcontextprotocol.spec.ProtocolVersions; +import io.modelcontextprotocol.server.transport.TomcatTestUtil; +import reactor.test.StepVerifier; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * Verifies that an {@code application/json} response whose body is not valid JSON fails + * the {@link HttpClientStreamableHttpTransport#sendMessage} mono with the parsing error + * instead of completing it successfully. + * + *

+ * Completing the delivery sink before deserialization used to swallow the parse failure: + * the {@code McpClientSession} then never received the error, kept the pending response + * entry and the caller only saw a {@code TimeoutException} once the request timeout + * elapsed. + * + * @see #1147 + */ +public class HttpClientStreamableHttpTransportInvalidJsonResponseTest { + + static int PORT = TomcatTestUtil.findAvailablePort(); + + static String host = "http://localhost:" + PORT; + + static HttpServer server; + + @BeforeAll + static void startServer() throws IOException { + server = HttpServer.create(new InetSocketAddress(PORT), 0); + + // 200 OK with an invalid JSON body for the /mcp endpoint + server.createContext("/mcp", exchange -> { + byte[] body = "{broken".getBytes(); + exchange.getResponseHeaders().set("Content-Type", "application/json"); + exchange.sendResponseHeaders(200, body.length); + exchange.getResponseBody().write(body); + exchange.close(); + }); + + server.setExecutor(null); + server.start(); + } + + @AfterAll + static void stopServer() { + server.stop(1); + } + + @Test + @Timeout(10) + void testInvalidJsonResponseFailsWithParseError() { + var transport = HttpClientStreamableHttpTransport.builder(host).build(); + + var initializeRequest = McpSchema.InitializeRequest + .builder(ProtocolVersions.MCP_2025_03_26, McpSchema.ClientCapabilities.builder().roots(true).build(), + McpSchema.Implementation.builder("MCP Client", "0.3.1").build()) + .build(); + var testMessage = new McpSchema.JSONRPCRequest(McpSchema.METHOD_INITIALIZE, "test-id", initializeRequest); + + StepVerifier.create(transport.sendMessage(testMessage)).expectErrorSatisfies(error -> { + // The parse failure must surface as the delivery error, not a timeout + assertThat(error).isInstanceOf(McpTransportException.class); + }).verify(Duration.ofSeconds(5)); + } + +} From 1cf7903935ac6a99ea8920b4ce85a43e4f1d0689 Mon Sep 17 00:00:00 2001 From: Daniel Garnier-Moiroux Date: Wed, 30 Sep 2026 15:05:36 +0200 Subject: [PATCH 2/3] Refactor HttpClient-based transports to use Publisher instead of Subscriber (#1079) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit HttpClient-based transports used to capture the enclosing sseSink in the subscriber, leading to HttpClient leaks when the client was closed. This PR addresses this, and adds many other improvements to HttpClient-based transports. This has no public API change. Improved transport reliability: - Closed transports no longer keep their `HttpClient` alive, so selector threads and memory stop piling up - `closeGracefully()` now releases open connections even when the session DELETE fails, e.g. when the server is down - `connect()` on the legacy SSE transport no longer hangs when it gets the stream ends (error, stream closed,`closeGracefully()`, ...) before the first event or - `sendMessage()` on Streamable HTTP no longer hangs when the SSE stream is closed without response or before the response arrives - Responses the client never reads are always released (e.g. `DELETE`), so connections go back to the pool. Errors surface immediately instead of as timeouts: - On Streamable HTTP, a JSON response that can't be read (malformed, or over maxResponseSize) now fails the request immediately with the real cause, instead of a TimeoutException after requestTimeout` - A server that answers a request with an empty JSON body now makes that request fail instead of silently timing out. An empty body in reply to a notification is still tolerated. - Server-caused errors are now McpTransportException instead of a plain RuntimeException, and the message includes the response body the server sent. - Errors that happen after connect() or sendMessage() has already completed now reach the transport's exception handler instead of Reactor "onErrorDropped" logs. - A 404 or 400 invalidates the session only if the request that got it carried a session id. Fixes a race condition where are reconnect got a new session while another request was already in flight (with the old session). The old request could have ended up invalidating the new session. Performance - Large SSE responses, such as multi-MB tool results, are no longer slow to receive. SSE parsing spec compliance - Unknown fields such as retry: are ignored instead of failing the stream with "Invalid SSE response". - The event type resets after each event, so a message that follows a named event is no longer misclassified and dropped. - A data: line containing U+2028, U+2029 or U+0085 is no longer truncated. - The legacy SSE transport skips empty "primer" events and unknown event types instead of failing. - An empty id: clears the last event id. Fixes #547 Fixes #620 Fixes #1042 Fixes #1047 Fixes #1147 Signed-off-by: Daniel Garnier-Moiroux Signed-off-by: Dariusz Jędrzejczyk --- .../HttpClientSseClientTransport.java | 150 +++-- .../HttpClientStreamableHttpTransport.java | 519 ++++++--------- .../transport/ResponseBodyHandlers.java | 576 ++++++++++++++++ .../client/transport/ResponseSubscribers.java | 619 ------------------ .../spec/DefaultMcpTransportSession.java | 1 + .../transport/BoundedBodySubscriberTests.java | 334 ---------- .../HttpClientHttpTransportLeakTests.java | 107 +++ ...pClientSseClientTransportConnectTests.java | 109 +++ ...reamableHttpTransportSendMessageTests.java | 121 ++++ .../transport/LargeSseEventDecodingTests.java | 219 +++++++ .../transport/LoopbackMcpHttpServer.java | 136 ++++ .../ResponseBodyHandlersSendAsyncTests.java | 78 +++ .../client/transport/SseEventParserTests.java | 185 ++++++ .../transport/Utf8LineDecoderBoundTests.java | 131 ++++ .../transport/Utf8LineDecoderTests.java | 256 ++++++++ .../spec/DefaultMcpTransportSessionTests.java | 26 + .../HttpClientBoundedReadTestSupport.java | 98 ++- ...entSseClientTransportBoundedReadTests.java | 28 +- ...reamableHttpTransportBoundedReadTests.java | 36 +- ...bleHttpTransportEmptyJsonResponseTest.java | 94 --- ...amableHttpTransportEmptyResponseTests.java | 147 +++++ ...amableHttpTransportLargeResponseTests.java | 295 +++++++++ 22 files changed, 2809 insertions(+), 1456 deletions(-) create mode 100644 mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java delete mode 100644 mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java delete mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/BoundedBodySubscriberTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientHttpTransportLeakTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportConnectTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportSendMessageTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LargeSseEventDecodingTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LoopbackMcpHttpServer.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlersSendAsyncTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderBoundTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderTests.java create mode 100644 mcp-core/src/test/java/io/modelcontextprotocol/spec/DefaultMcpTransportSessionTests.java delete mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyJsonResponseTest.java create mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyResponseTests.java create mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportLargeResponseTests.java diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransport.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransport.java index 3e4b613ff..9ed5c5cd4 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransport.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransport.java @@ -11,12 +11,11 @@ import java.net.http.HttpResponse; import java.time.Duration; import java.util.List; -import java.util.concurrent.CompletableFuture; +import java.util.Optional; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; import java.util.function.Function; -import io.modelcontextprotocol.client.transport.ResponseSubscribers.ResponseEvent; import io.modelcontextprotocol.client.transport.customizer.McpAsyncHttpClientRequestCustomizer; import io.modelcontextprotocol.client.transport.customizer.McpSyncHttpClientRequestCustomizer; import io.modelcontextprotocol.common.McpTransportContext; @@ -390,70 +389,93 @@ public Mono connect(Function, Mono> h var transportContext = ctx.getOrDefault(McpTransportContext.KEY, McpTransportContext.EMPTY); return Mono.from(this.httpRequestCustomizer.customize(builder, "GET", uri, null, transportContext)); }).flatMap(requestBuilder -> Mono.create(sink -> { - Disposable connection = Flux.create( - sseSink -> this.httpClient - .sendAsync(requestBuilder.build(), - responseInfo -> ResponseSubscribers.sseToBodySubscriber(responseInfo, sseSink, - this.maxResponseSize)) - .exceptionallyCompose(e -> { - sseSink.error(e); - return CompletableFuture.failedFuture(e); - })) - .map(responseEvent -> (ResponseSubscribers.SseResponseEvent) responseEvent) - .flatMap(responseEvent -> { + Disposable connection = ResponseBodyHandlers.sendAsync(this.httpClient, requestBuilder.build()) + .flatMapMany(response -> { if (isClosing) { - return Mono.empty(); + // The body is handed over as a publisher and the connection is + // only released once it is subscribed to. It is an SSE stream + // that may never end, so it is cancelled rather than drained. + return ResponseBodyHandlers.cancel(response.body()); } - int statusCode = responseEvent.responseInfo().statusCode(); + int statusCode = response.statusCode(); if (statusCode >= 200 && statusCode < 300) { - try { - if (ENDPOINT_EVENT_TYPE.equals(responseEvent.sseEvent().event())) { - String messageEndpointUri = responseEvent.sseEvent().data(); - try { - messageEndpointValidator.validate(uri, messageEndpointUri); - } - catch (InvalidSseMessageEndpointException e) { - sink.error(e); - this.messageEndpointSink.tryEmitError(e); - return Flux.error(e); - } - if (this.messageEndpointSink.tryEmitValue(messageEndpointUri).isSuccess()) { - sink.success(); - return Flux.empty(); // No further processing needed - } - else { - sink.error(new RuntimeException("Failed to handle SSE endpoint event")); - } + Flux lines = ResponseBodyHandlers.decodeLines(response.body(), this.maxResponseSize); + return ResponseBodyHandlers.decodeSseResponse(lines, this.maxResponseSize); + } + else { + return ResponseBodyHandlers.readThenError(response.body(), this.maxResponseSize, + "Failed to connect to SSE stream: " + statusCode); + } + }) + // Every successfully processed event yields exactly one element, empty + // when it carries no message, so that the first one can mark the + // connection as established. + .>handle((sseEvent, events) -> { + try { + if (ENDPOINT_EVENT_TYPE.equals(sseEvent.event())) { + String messageEndpointUri = sseEvent.data(); + try { + messageEndpointValidator.validate(uri, messageEndpointUri); + } + catch (InvalidSseMessageEndpointException e) { + this.messageEndpointSink.tryEmitError(e); + events.error(e); + return; } - else if (MESSAGE_EVENT_TYPE.equals(responseEvent.sseEvent().event())) { - JSONRPCMessage message = McpSchema.deserializeJsonRpcMessage(jsonMapper, - responseEvent.sseEvent().data()); - sink.success(); - return Flux.just(message); + if (this.messageEndpointSink.tryEmitValue(messageEndpointUri).isSuccess()) { + events.next(Optional.empty()); } else { - logger.debug("Received unrecognized SSE event type: {}", responseEvent.sseEvent()); - sink.success(); + events.error(new McpTransportException("Failed to handle SSE endpoint event")); } } - catch (IOException e) { - sink.error(new McpTransportException("Error processing SSE event", e)); + else if (MESSAGE_EVENT_TYPE.equals(sseEvent.event())) { + String data = sseEvent.data(); + if (data == null || data.isBlank()) { + logger.debug("Skipping SSE event with empty data (stream primer)"); + events.next(Optional.empty()); + } + else { + events.next(Optional.of(McpSchema.deserializeJsonRpcMessage(jsonMapper, data))); + } + } + else { + logger.debug("Received unrecognized SSE event type: {}", sseEvent); + events.next(Optional.empty()); } } - return Flux.error( - new RuntimeException("Failed to send message: " + responseEvent)); - + catch (IOException e) { + events.error(new McpTransportException("Error processing SSE event", e)); + } }) - .flatMap(jsonRpcMessage -> handler.apply(Mono.just(jsonRpcMessage))) + // connect() is resolved by the first signal only: any later failure is + // merely logged below, as connect() has already completed by then. + .switchOnFirst((first, events) -> { + if (first.hasValue()) { + sink.success(); + } + else if (first.isOnError()) { + sink.error(first.getThrowable()); + } + else if (first.isOnComplete()) { + sink.error(new McpTransportException("SSE stream closed before any event was received")); + } + return events; + }) + .handle((message, messages) -> message.ifPresent(messages::next)) + .flatMap(message -> handler.apply(Mono.just(message))) .onErrorComplete(t -> { if (!isClosing) { logger.warn("SSE stream observed an error", t); - sink.error(t); } return true; }) + // A closeGracefully() before the first signal cancels the stream: + // complete + // connect() instead of leaving it pending. A no-op once it has resolved. + .doOnCancel(sink::success) .doFinally(s -> { Disposable ref = this.sseSubscription.getAndSet(null); if (ref != null && !ref.isDisposed()) { @@ -486,17 +508,7 @@ public Mono sendMessage(JSONRPCMessage message) { } return this.serializeMessage(message) - .flatMap(body -> sendHttpPost(messageEndpointUri, body).handle((response, sink) -> { - if (response.statusCode() != 200 && response.statusCode() != 201 && response.statusCode() != 202 - && response.statusCode() != 206) { - sink.error(new RuntimeException("Sending message failed with a non-OK HTTP code: " - + response.statusCode() + " - " + response.body())); - } - else { - sink.next(response); - sink.complete(); - } - })) + .flatMap(body -> sendHttpPost(messageEndpointUri, body)) .doOnError(error -> { if (!isClosing) { logger.error("Error sending message: {}", error.getMessage()); @@ -517,7 +529,16 @@ private Mono serializeMessage(final JSONRPCMessage message) { }); } - private Mono> sendHttpPost(final String endpoint, final String body) { + /** + * POSTs {@code body} to {@code endpoint} and consumes the response, failing if the + * server did not accept the message. + * + *

+ * The response body is streamed rather than aggregated: it is only read as text when + * a non-OK status makes it part of the failure message, and discarded otherwise. + * Either way it has to be consumed, or the connection is never released. + */ + private Mono sendHttpPost(final String endpoint, final String body) { final URI requestUri = Utils.resolveUri(baseUri, endpoint); return Mono.deferContextual(ctx -> { var builder = this.requestBuilder.copy() @@ -529,8 +550,15 @@ private Mono> sendHttpPost(final String endpoint, final Str return Mono.from(this.httpRequestCustomizer.customize(builder, "POST", requestUri, body, transportContext)); }).flatMap(customizedBuilder -> { var request = customizedBuilder.build(); - return Mono.fromFuture( - httpClient.sendAsync(request, ResponseSubscribers.boundedStringBodyHandler(this.maxResponseSize))); + return ResponseBodyHandlers.sendAsync(this.httpClient, request).flatMap(response -> { + int statusCode = response.statusCode(); + if (statusCode == 200 || statusCode == 201 || statusCode == 202 || statusCode == 206) { + return ResponseBodyHandlers.drain(response.body(), this.maxResponseSize).then(); + } + return ResponseBodyHandlers.decodeAggregateResponse(response.body(), this.maxResponseSize) + .flatMap(text -> Mono.error(new McpTransportException( + "Sending message failed with a non-OK HTTP code: " + statusCode + " - " + text))); + }); }); } diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java index 5517823b6..9d55e816c 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java @@ -9,19 +9,19 @@ import java.net.http.HttpClient; import java.net.http.HttpRequest; import java.net.http.HttpResponse; -import java.net.http.HttpResponse.BodyHandler; +import java.nio.ByteBuffer; import java.time.Duration; import java.util.Collections; import java.util.Comparator; import java.util.List; import java.util.Optional; import java.util.concurrent.CompletionException; +import java.util.concurrent.Flow; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; import java.util.function.Function; import io.modelcontextprotocol.client.McpAsyncClient; -import io.modelcontextprotocol.client.transport.ResponseSubscribers.ResponseEvent; import io.modelcontextprotocol.client.transport.customizer.McpAsyncHttpClientRequestCustomizer; import io.modelcontextprotocol.client.transport.customizer.McpHttpClientAuthorizationErrorHandler; import io.modelcontextprotocol.client.transport.customizer.McpHttpClientTransportAuthorizationErrorHandler; @@ -49,7 +49,6 @@ import org.slf4j.LoggerFactory; import reactor.core.Disposable; import reactor.core.publisher.Flux; -import reactor.core.publisher.FluxSink; import reactor.core.publisher.Mono; import reactor.util.function.Tuple2; import reactor.util.function.Tuples; @@ -207,7 +206,7 @@ public static Builder builder(String baseUri) { @Override public Mono connect(Function, Mono> handler) { - return Mono.deferContextual(ctx -> { + return Mono.defer(() -> { this.handler.set(handler); if (this.openConnectionOnStartup) { logger.debug("Eagerly opening connection on startup"); @@ -240,11 +239,13 @@ private Publisher createDelete(String sessionId) { .DELETE(); var transportContext = ctx.getOrDefault(McpTransportContext.KEY, McpTransportContext.EMPTY); return Mono.from(this.httpRequestCustomizer.customize(builder, "DELETE", uri, null, transportContext)); - }).flatMap(requestBuilder -> { - var request = requestBuilder.build(); - return Mono.fromFuture(() -> this.httpClient.sendAsync(request, - ResponseSubscribers.boundedStringBodyHandler(this.maxResponseSize))); - }).then(); + }) + .flatMap(requestBuilder -> ResponseBodyHandlers.sendAsync(this.httpClient, requestBuilder.build()) + // The response is not inspected, but the body still has to be consumed + // to release the connection. + .flatMapMany(response -> ResponseBodyHandlers.drain(response.body(), this.maxResponseSize)) + .then()) + .then(); } @Override @@ -254,7 +255,8 @@ public void setExceptionHandler(Consumer handler) { } private void handleException(Throwable t) { - logger.debug("Handling exception for session {}", sessionIdOrPlaceholder(this.activeSession.get()), t); + logger.debug("Handling exception for session {}", sessionIdOrPlaceholder( + activeSession.get() != null ? activeSession.get().sessionId() : Optional.empty()), t); if (t instanceof McpTransportSessionNotFoundException) { McpTransportSession invalidSession = this.activeSession.getAndSet(createTransportSession()); logger.warn("Server does not recognize session {}. Invalidating.", invalidSession.sessionId()); @@ -266,6 +268,15 @@ private void handleException(Throwable t) { } } + private void handleExceptionSafely(Throwable t) { + try { + handleException(t); + } + catch (Exception e) { + logger.error("Error handling exception {}", t.getMessage(), e); + } + } + @Override public Mono closeGracefully() { return Mono.defer(() -> { @@ -279,6 +290,38 @@ public Mono closeGracefully() { }); } + /** + * Every successfully processed event yields exactly one element, empty when it + * carries no message, so that callers can tell when the first one has arrived. + */ + private Flux> consumeSseStream(Flow.Publisher> body, + McpTransportStream existingStream) { + Flux lines = ResponseBodyHandlers.decodeLines(body, this.maxResponseSize); + return ResponseBodyHandlers.decodeSseResponse(lines, this.maxResponseSize).flatMap(sseEvent -> { + if (!isMessageEvent(sseEvent.event())) { + logger.debug("Received SSE event with type: {}", sseEvent); + return Flux.just(Optional.empty()); + } + String data = sseEvent.data(); + if (data == null || data.isBlank()) { + logger.debug("Skipping SSE event with empty data (stream primer)"); + return Flux.just(Optional.empty()); + } + try { + McpSchema.JSONRPCMessage message = McpSchema.deserializeJsonRpcMessage(this.jsonMapper, data); + Tuple2, Iterable> idWithMessages = Tuples + .of(Optional.ofNullable(sseEvent.id()), List.of(message)); + McpTransportStream sessionStream = existingStream != null ? existingStream + : new DefaultMcpTransportStream<>(this.resumableStreams, this::reconnect); + return Flux.from(sessionStream.consumeSseStream(Flux.just(idWithMessages))).map(Optional::of); + } + catch (IOException e) { + return Flux.>error( + new McpTransportException("Error parsing JSON-RPC message: " + data, e)); + } + }); + } + private Mono reconnect(McpTransportStream stream) { return Mono.deferContextual(ctx -> { var rh = this.handler.get(); @@ -304,9 +347,8 @@ private Mono reconnect(McpTransportStream stream) { final AtomicReference disposableRef = new AtomicReference<>(); - var uri = Utils.resolveUri(this.baseUri, this.endpoint); - Disposable connection = Mono.deferContextual(connectionCtx -> { + var uri = Utils.resolveUri(this.baseUri, this.endpoint); HttpRequest.Builder requestBuilder = this.requestBuilder.copy(); if (transportSession != null && transportSession.sessionId().isPresent()) { @@ -327,124 +369,54 @@ private Mono reconnect(McpTransportStream stream) { .GET(); var transportContext = connectionCtx.getOrDefault(McpTransportContext.KEY, McpTransportContext.EMPTY); return Mono.from(this.httpRequestCustomizer.customize(builder, "GET", uri, null, transportContext)); + }).flatMapMany(requestBuilder -> { + var request = requestBuilder.build(); + return ResponseBodyHandlers.sendAsync(this.httpClient, request).flatMapMany(httpResponse -> { + int statusCode = httpResponse.statusCode(); + if (statusCode == 401 || statusCode == 403) { + logger.debug("Authorization error in reconnect with code {}", statusCode); + var requestSnapshot = new HttpRequestSnapshot(request.uri(), request.method(), + request.headers()); + return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize, + new McpHttpClientTransportAuthorizationException( + "Authorization error connecting to SSE stream", requestSnapshot, + toResponseInfo(httpResponse))); + } + if (statusCode == METHOD_NOT_ALLOWED) { + logger.debug("The server does not support SSE streams, using request-response mode."); + return ResponseBodyHandlers.drain(httpResponse.body(), this.maxResponseSize); + } + if (statusCode < 200 || statusCode >= 300) { + return statusError(request, httpResponse); + } + String contentType = httpResponse.headers() + .firstValue(HttpHeaders.CONTENT_TYPE) + .orElse("") + .toLowerCase(); + if (!contentType.contains(TEXT_EVENT_STREAM)) { + return ResponseBodyHandlers.readThenError(httpResponse.body(), this.maxResponseSize, + "Unrecognized server error when connecting to SSE stream, status code: " + statusCode); + } + logger.debug("SSE connection established successfully"); + return consumeSseStream(httpResponse.body(), stream); + }); }) - .flatMapMany(requestBuilder -> Flux.create(sseSink -> this.httpClient - .sendAsync(requestBuilder.build(), this.toSendMessageBodySubscriber(sseSink)) - .whenComplete((response, throwable) -> { - if (throwable != null) { - sseSink.error(throwable); - } - else { - logger.debug("SSE connection established successfully"); - } - })).flatMap(responseEvent -> { - int statusCode = responseEvent.responseInfo().statusCode(); - if (statusCode == 401 || statusCode == 403) { - logger.debug("Authorization error in reconnect with code {}", statusCode); - var request = requestBuilder.build(); - var requestSnapshot = new HttpRequestSnapshot(request.uri(), request.method(), - request.headers()); - return Mono.error( - new McpHttpClientTransportAuthorizationException( - "Authorization error connecting to SSE stream", requestSnapshot, - responseEvent.responseInfo())); - } - else if (statusCode == METHOD_NOT_ALLOWED) { - logger.debug("The server does not support SSE streams, using request-response mode."); - return Flux.empty(); - } - - if (!(responseEvent instanceof ResponseSubscribers.SseResponseEvent sseResponseEvent)) { - return Flux.error(new McpTransportException( - "Unrecognized server error when connecting to SSE stream, status code: " - + statusCode)); - } - else if (statusCode >= 200 && statusCode < 300) { - if (isMessageEvent(sseResponseEvent.sseEvent().event())) { - String data = sseResponseEvent.sseEvent().data(); - // Per 2025-11-25 spec (SEP-1699), servers may - // send SSE events - // with empty data to prime the client for - // reconnection. - // Skip these events as they contain no JSON-RPC - // message. - if (data == null || data.isBlank()) { - logger.debug("Skipping SSE event with empty data (stream primer)"); - return Flux.empty(); - } - try { - // We don't support batching ATM and probably - // won't since the next version considers - // removing it. - McpSchema.JSONRPCMessage message = McpSchema - .deserializeJsonRpcMessage(this.jsonMapper, data); - - Tuple2, Iterable> idWithMessages = Tuples - .of(Optional.ofNullable(sseResponseEvent.sseEvent().id()), List.of(message)); - - McpTransportStream sessionStream = stream != null ? stream - : new DefaultMcpTransportStream<>(this.resumableStreams, this::reconnect); - logger.debug("Connected stream {}", sessionStream.streamId()); - - return Flux.from(sessionStream.consumeSseStream(Flux.just(idWithMessages))); - - } - catch (IOException ioException) { - return Flux.error(new McpTransportException( - "Error parsing JSON-RPC message: " + responseEvent, ioException)); - } - } - else { - logger.debug("Received SSE event with type: {}", sseResponseEvent.sseEvent()); - return Flux.empty(); - } - } - else if (statusCode == NOT_FOUND) { - - if (transportSession != null && transportSession.sessionId().isPresent()) { - // only if the request was sent with a session id - // and the response is 404, we consider it a - // session not found error. - logger.debug("Session not found for session ID: {}", - transportSession.sessionId().get()); - String sessionIdRepresentation = sessionIdOrPlaceholder(transportSession); - McpTransportSessionNotFoundException exception = new McpTransportSessionNotFoundException( - "Session not found for session ID: " + sessionIdRepresentation); - return Flux.error(exception); - } - return Flux.error( - new McpTransportException("Server Not Found. Status code:" + statusCode - + ", response-event:" + responseEvent)); - } - else if (statusCode == BAD_REQUEST) { - if (transportSession != null && transportSession.sessionId().isPresent()) { - // only if the request was sent with a session id - // and thre response is 404, we consider it a - // session not found error. - String sessionIdRepresentation = sessionIdOrPlaceholder(transportSession); - McpTransportSessionNotFoundException exception = new McpTransportSessionNotFoundException( - "Session not found for session ID: " + sessionIdRepresentation); - return Flux.error(exception); - } - return Flux.error(new McpTransportException( - "Bad Request. Status code:" + statusCode + ", response-event:" + responseEvent)); - } - return Flux.error(new McpTransportException( - "Received unrecognized SSE event type: " + sseResponseEvent.sseEvent().event())); - }) - .retryWhen(authorizationErrorRetrySpec()) - .flatMap(jsonrpcMessage -> requestHandler.apply(Mono.just(jsonrpcMessage))) - .onErrorMap(CompletionException.class, t -> t.getCause()) - .doFinally(s -> { - Disposable ref = disposableRef.getAndSet(null); - if (ref != null) { - transportSession.removeConnection(ref); - } - })) + .retryWhen(authorizationErrorRetrySpec()).handle((message, messages) -> message.ifPresent(messages::next)) + .flatMap(jsonrpcMessage -> requestHandler.apply(Mono.just(jsonrpcMessage))) .onErrorComplete(t -> { + if (t instanceof CompletionException) { + t = t.getCause(); + } this.handleException(t); return true; }) + .doFinally(s -> { + Disposable ref = disposableRef.getAndSet(null); + if (ref != null) { + transportSession.removeConnection(ref); + } + }) .contextWrite(ctx) .subscribe(); @@ -455,6 +427,14 @@ else if (statusCode == BAD_REQUEST) { } + private static HttpResponse.ResponseInfo toResponseInfo(HttpResponse>> response) { + return new HttpClientResponseInfo(response.statusCode(), response.headers(), response.version()); + } + + private record HttpClientResponseInfo(int statusCode, java.net.http.HttpHeaders headers, + HttpClient.Version version) implements HttpResponse.ResponseInfo { + } + private Retry authorizationErrorRetrySpec() { return Retry.from(companion -> companion.flatMap(retrySignal -> { if (!(retrySignal.failure() instanceof McpHttpClientTransportAuthorizationException authException)) { @@ -475,31 +455,6 @@ private Retry authorizationErrorRetrySpec() { })); } - private BodyHandler toSendMessageBodySubscriber(FluxSink sink) { - - BodyHandler responseBodyHandler = responseInfo -> { - - String contentType = responseInfo.headers().firstValue(HttpHeaders.CONTENT_TYPE).orElse("").toLowerCase(); - - if (contentType.contains(TEXT_EVENT_STREAM)) { - // For SSE streams, use line subscriber that returns Void - logger.debug("Received SSE stream response, using line subscriber"); - return ResponseSubscribers.sseToBodySubscriber(responseInfo, sink, this.maxResponseSize); - } - else if (contentType.contains(APPLICATION_JSON)) { - // For JSON responses and others, use string subscriber - logger.debug("Received response, using string subscriber"); - return ResponseSubscribers.aggregateBodySubscriber(responseInfo, sink, this.maxResponseSize); - } - - logger.debug("Received Bodyless response, using discarding subscriber"); - return ResponseSubscribers.bodilessBodySubscriber(responseInfo, sink, this.maxResponseSize); - }; - - return responseBodyHandler; - - } - public String toString(McpSchema.JSONRPCMessage message) { try { return this.jsonMapper.writeValueAsString(message); @@ -529,9 +484,6 @@ public Mono sendMessage(McpSchema.JSONRPCMessage sentMessage) { final AtomicReference disposableRef = new AtomicReference<>(); - var uri = Utils.resolveUri(this.baseUri, this.endpoint); - String jsonBody = this.toString(sentMessage); - Disposable connection = Mono.deferContextual(ctx -> { HttpRequest.Builder requestBuilder = this.requestBuilder.copy(); @@ -540,6 +492,8 @@ public Mono sendMessage(McpSchema.JSONRPCMessage sentMessage) { transportSession.sessionId().get()); } + String jsonBody = this.toString(sentMessage); + var uri = Utils.resolveUri(this.baseUri, this.endpoint); var builder = requestBuilder.uri(uri) .header(HttpHeaders.ACCEPT, APPLICATION_JSON + ", " + TEXT_EVENT_STREAM) .header(HttpHeaders.CONTENT_TYPE, APPLICATION_JSON_UTF8) @@ -551,179 +505,114 @@ public Mono sendMessage(McpSchema.JSONRPCMessage sentMessage) { var transportContext = ctx.getOrDefault(McpTransportContext.KEY, McpTransportContext.EMPTY); return Mono .from(this.httpRequestCustomizer.customize(builder, "POST", uri, jsonBody, transportContext)); - }).flatMapMany(requestBuilder -> Flux.create(responseEventSink -> { - // Create the async request with proper body subscriber selection - Mono.fromFuture(this.httpClient - .sendAsync(requestBuilder.build(), this.toSendMessageBodySubscriber(responseEventSink)) - .whenComplete((response, throwable) -> { - if (throwable != null) { - responseEventSink.error(throwable); - } - else { - logger.debug("SSE connection established successfully"); - } - })).onErrorMap(CompletionException.class, t -> t.getCause()).onErrorComplete().subscribe(); - - }).flatMap(responseEvent -> { - int statusCode = responseEvent.responseInfo().statusCode(); - if (statusCode == 401 || statusCode == 403) { - var request = requestBuilder.build(); - var requestSnapshot = new HttpRequestSnapshot(request.uri(), request.method(), request.headers()); - logger.debug("Authorization error in sendMessage with code {}", statusCode); - return Mono.error(new McpHttpClientTransportAuthorizationException( - "Authorization error when sending message", requestSnapshot, responseEvent.responseInfo())); - } - - if (transportSession.markInitialized( - responseEvent.responseInfo().headers().firstValue("mcp-session-id").orElseGet(() -> null))) { - // Once we have a session, we try to open an async stream for - // the server to send notifications and requests out-of-band. - - reconnect(null).contextWrite(deliveredSink.contextView()).subscribe(); - } + }).flatMapMany(requestBuilder -> { + var request = requestBuilder.build(); + return ResponseBodyHandlers.sendAsync(this.httpClient, request).flatMapMany(httpResponse -> { + int statusCode = httpResponse.statusCode(); + if (statusCode == 401 || statusCode == 403) { + logger.debug("Authorization error in sendMessage with code {}", statusCode); + var requestSnapshot = new HttpRequestSnapshot(request.uri(), request.method(), + request.headers()); + return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize, + new McpHttpClientTransportAuthorizationException( + "Authorization error when sending message", requestSnapshot, + toResponseInfo(httpResponse))); + } - String sessionRepresentation = sessionIdOrPlaceholder(transportSession); + if (transportSession + .markInitialized(httpResponse.headers().firstValue("mcp-session-id").orElse(null))) { + // Fails only when the transport has been closed in the meantime, + // in which case there is no stream left to open. + reconnect(null).contextWrite(deliveredSink.contextView()).subscribe(ignored -> { + }, t -> logger.debug("Not opening the SSE stream: {}", t.getMessage())); + } - if (statusCode >= 200 && statusCode < 300) { + if (statusCode < 200 || statusCode >= 300) { + return statusError(request, httpResponse); + } - String contentType = responseEvent.responseInfo() - .headers() + String sessionRepresentation = sessionIdOrPlaceholder( + request.headers().firstValue(HttpHeaders.MCP_SESSION_ID)); + String contentType = httpResponse.headers() .firstValue(HttpHeaders.CONTENT_TYPE) .orElse("") .toLowerCase(); + String contentLength = httpResponse.headers().firstValue(HttpHeaders.CONTENT_LENGTH).orElse(null); - String contentLength = responseEvent.responseInfo() - .headers() - .firstValue(HttpHeaders.CONTENT_LENGTH) - .orElse(null); - - // For empty content or HTTP code 202 (ACCEPTED), assume success if (contentType.isBlank() || "0".equals(contentLength) || statusCode == 202) { - // if (contentType.isBlank() || "0".equals(contentLength)) { logger.debug("No body returned for POST in session {}", sessionRepresentation); - // No content type means no response body, so we can just - // return an empty stream - deliveredSink.success(); - return Flux.empty(); + return ResponseBodyHandlers.>drain(httpResponse.body(), + this.maxResponseSize) + .startWith(Optional.empty()); } else if (contentType.contains(TEXT_EVENT_STREAM)) { - return Flux.just(((ResponseSubscribers.SseResponseEvent) responseEvent).sseEvent()) - .flatMap(sseEvent -> { - String data = sseEvent.data(); - // Per 2025-11-25 spec (SEP-1699), servers may send SSE - // events - // with empty data to prime the client for reconnection. - // Skip these events as they contain no JSON-RPC message. - if (data == null || data.isBlank()) { - logger.debug("Skipping SSE event with empty data (stream primer)"); - return Flux.empty(); - } - try { - // We don't support batching ATM and probably - // won't - // since the - // next version considers removing it. - McpSchema.JSONRPCMessage message = McpSchema - .deserializeJsonRpcMessage(this.jsonMapper, data); - - Tuple2, Iterable> idWithMessages = Tuples - .of(Optional.ofNullable(sseEvent.id()), List.of(message)); - - McpTransportStream sessionStream = new DefaultMcpTransportStream<>( - this.resumableStreams, this::reconnect); - - logger.debug("Connected stream {}", sessionStream.streamId()); - - deliveredSink.success(); - - return Flux.from(sessionStream.consumeSseStream(Flux.just(idWithMessages))); - } - catch (IOException ioException) { - return Flux.error(new McpTransportException( - "Error parsing JSON-RPC message: " + responseEvent, ioException)); - } - }); + return consumeSseStream(httpResponse.body(), null); } else if (contentType.contains(APPLICATION_JSON)) { - deliveredSink.success(); - String data = ((ResponseSubscribers.AggregateResponseEvent) responseEvent).data(); - if (sentMessage instanceof McpSchema.JSONRPCNotification) { - logger.warn("Notification: {} received non-compliant response: {}", sentMessage, - Utils.hasText(data) ? data : "[empty]"); - return Mono.empty(); - } - - try { - return Mono.just(McpSchema.deserializeJsonRpcMessage(jsonMapper, data)); - } - catch (IOException e) { - return Mono.error(new McpTransportException( - "Error deserializing JSON-RPC message: " + responseEvent, e)); - } + return ResponseBodyHandlers.decodeAggregateResponse(httpResponse.body(), + this.maxResponseSize).>handle((data, messages) -> { + if (sentMessage instanceof McpSchema.JSONRPCNotification) { + logger.warn("Notification: {} received non-compliant response: {}", sentMessage, + Utils.hasText(data) ? data : "[empty]"); + messages.next(Optional.empty()); + return; + } + try { + messages + .next(Optional.of(McpSchema.deserializeJsonRpcMessage(jsonMapper, data))); + } + catch (IOException e) { + messages.error(new McpTransportException( + "Error deserializing JSON-RPC message: " + data, e)); + } + }) + .flux(); } + logger.warn("Unknown media type {} returned for POST in session {}", contentType, sessionRepresentation); - - return Flux.error( - new RuntimeException("Unknown media type returned: " + contentType)); - } - else if (statusCode == NOT_FOUND) { - if (transportSession != null && transportSession.sessionId().isPresent()) { - // only if the request was sent with a session id and the - // response is 404, we consider it a session not found error. - logger.debug("Session not found for session ID: {}", transportSession.sessionId().get()); - McpTransportSessionNotFoundException exception = new McpTransportSessionNotFoundException( - "Session not found for session ID: " + sessionRepresentation); - return Flux.error(exception); - } - return Flux.error(new McpTransportException( - "Server Not Found. Status code:" + statusCode + ", response-event:" + responseEvent)); - } - else if (statusCode == BAD_REQUEST) { - // Some implementations can return 400 when presented with a - // session id that it doesn't know about, so we will - // invalidate the session - // https://github.com/modelcontextprotocol/typescript-sdk/issues/389 - - if (transportSession != null && transportSession.sessionId().isPresent()) { - // only if the request was sent with a session id and the - // response is 404, we consider it a session not found error. - McpTransportSessionNotFoundException exception = new McpTransportSessionNotFoundException( - "Session not found for session ID: " + sessionRepresentation); - return Flux.error(exception); - } - return Flux.error(new McpTransportException( - "Bad Request. Status code:" + statusCode + ", response-event:" + responseEvent)); - } - else if (statusCode >= 400 && statusCode < 500) { - return Flux.error( - new McpTransportException("Invalid request. Status code: " + statusCode)); - } - - return Flux.error( - new RuntimeException("Failed to send message: " + responseEvent)); + return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize, + new McpTransportException("Unknown media type returned: " + contentType)); + }); }) .retryWhen(authorizationErrorRetrySpec()) - .flatMap(jsonRpcMessage -> requestHandler.apply(Mono.just(jsonRpcMessage))) .onErrorMap(CompletionException.class, t -> t.getCause()) + // sendMessage() is resolved by the first signal only: any later failure + // is + // merely handled below, as sendMessage() has already completed by then. + // An exchange ending without any event still means the server accepted + // the message, so completion resolves it successfully too. + .switchOnFirst((first, messages) -> { + if (first.isOnError()) { + // Handled before failing sendMessage(), so that a session the + // server does not recognise is already invalidated by the time + // the caller learns about it. Consumed here so that it is not + // handled a second time below. + handleExceptionSafely(first.getThrowable()); + deliveredSink.error(first.getThrowable()); + return Flux.empty(); + } + deliveredSink.success(); + return messages; + }).handle((message, messages) -> message.ifPresent(messages::next)) + .flatMap(jsonRpcMessage -> requestHandler.apply(Mono.just(jsonRpcMessage))) .doFinally(s -> { - logger.debug("SendMessage finally: {}", s); Disposable ref = disposableRef.getAndSet(null); if (ref != null) { transportSession.removeConnection(ref); } - })).onErrorComplete(t -> { - // handle the error first - try { - this.handleException(t); - } - catch (Exception e) { - logger.error("Error handling exception {}", t.getMessage(), e); - } - // inform the caller of sendMessage - deliveredSink.error(t); + }) + .onErrorComplete(t -> { + handleExceptionSafely(t); return true; - }).contextWrite(deliveredSink.contextView()).subscribe(); + }) + // Closing the session before the first signal cancels the exchange: + // complete sendMessage() instead of leaving it pending. A no-op once it + // has resolved. + .doOnCancel(deliveredSink::success) + .contextWrite(deliveredSink.contextView()) + .subscribe(); disposableRef.set(connection); transportSession.addConnection(connection); @@ -731,8 +620,32 @@ else if (statusCode >= 400 && statusCode < 500) { } - private static String sessionIdOrPlaceholder(McpTransportSession transportSession) { - return transportSession.sessionId().orElse("[missing_session_id]"); + /** + * Fails the exchange over a response with an error status. A session id the server + * does not recognise invalidates the session; any other failure carries the response + * body, which is what the server said about it. + */ + private Flux statusError(HttpRequest request, HttpResponse>> response) { + int statusCode = response.statusCode(); + // Classify the response against the session id that this very request carried, + // rather than the one currently held by the session, which can be established + // concurrently. Some implementations return 400 rather than 404 for a session id + // they do not know about. + // https://github.com/modelcontextprotocol/typescript-sdk/issues/389 + Optional sessionId = request.headers().firstValue(HttpHeaders.MCP_SESSION_ID); + if ((statusCode == NOT_FOUND || statusCode == BAD_REQUEST) && sessionId.isPresent()) { + logger.debug("Session not found for session ID: {}", sessionId.get()); + return ResponseBodyHandlers.drainThenError(response.body(), this.maxResponseSize, + new McpTransportSessionNotFoundException(sessionId.get())); + } + String failure = statusCode == NOT_FOUND ? "Server Not Found. Status code:" + statusCode + : statusCode == BAD_REQUEST ? "Bad Request. Status code:" + statusCode + : "Received unexpected status code: " + statusCode; + return ResponseBodyHandlers.readThenError(response.body(), this.maxResponseSize, failure); + } + + private static String sessionIdOrPlaceholder(Optional sessionId) { + return sessionId.orElse("[missing_session_id]"); } @Override diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java new file mode 100644 index 000000000..01e243da5 --- /dev/null +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java @@ -0,0 +1,576 @@ +/* + * Copyright 2024 - 2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.ByteBuffer; +import java.nio.CharBuffer; +import java.nio.charset.CharacterCodingException; +import java.nio.charset.CharsetDecoder; +import java.nio.charset.CoderResult; +import java.nio.charset.CodingErrorAction; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import java.util.concurrent.CancellationException; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; +import java.util.concurrent.Flow; +import java.util.concurrent.Flow.Publisher; + +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +import io.modelcontextprotocol.spec.McpTransportException; +import io.modelcontextprotocol.util.Utils; +import reactor.adapter.JdkFlowAdapter; +import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; + +/** + * Utility class providing various operations for handling different types of HTTP + * response bodies in the context of Model Context Protocol (MCP) clients. + * + *

+ * Defines Flux operators for processing Server-Sent Events (SSE), aggregate responses, + * and bodiless responses. + * + * @author Christian Tzolov + * @author Dariusz Jędrzejczyk + * @author Daniel Garnier-Moiroux + */ +class ResponseBodyHandlers { + + /** + * Bytes of SSE field framing a single line may carry on top of the message payload: + * {@code "event: "} is the longest field prefix this parser recognises. Line + * terminators are not counted, as they reset the running line length. Without this + * allowance, an event carrying exactly the maximum message size would be rejected + * because of the bytes the SSE wire format adds around it. + */ + private static final int SSE_FRAMING_OVERHEAD = "event: ".length(); + + /** + * The type of an SSE event that does not name one with an {@code event:} field. + */ + private static final String DEFAULT_EVENT_TYPE = "message"; + + record SseEvent(String id, String event, String data) { + } + + /** + * Adds {@link #SSE_FRAMING_OVERHEAD} to {@code maxSize}, saturating at + * {@link Integer#MAX_VALUE} rather than overflowing into a negative bound that would + * reject everything. + */ + private static int plusFramingOverhead(int maxSize) { + return maxSize > Integer.MAX_VALUE - SSE_FRAMING_OVERHEAD ? Integer.MAX_VALUE : maxSize + SSE_FRAMING_OVERHEAD; + } + + /** + * Converts a publisher of byte-buffer chunks into a flux of decoded string lines, + * bounding how much memory a single line may occupy. + * + *

+ * The decoder buffers characters until it encounters a line terminator, so a peer + * that never terminates a line (or sends an enormous one) would force the transport + * to buffer it in memory. Exceeding the bound fails the flux, which cancels the + * subscription and so closes the connection. + * + *

+ * The bound is allowed {@link #SSE_FRAMING_OVERHEAD} extra characters so that the SSE + * framing around a payload does not count against the payload's own budget: only SSE + * streams are read line by line, so every caller of this method is parsing one. + * + *

+ * This only bounds a single line. What accumulates across lines is bounded where it + * accumulates: see {@link #decodeSseResponse} for multi-line SSE events and + * {@link #decodeAggregateResponse} for whole response bodies. + * @param publisher the response body + * @param maxSize the maximum number of bytes read for a single inbound message + */ + static Flux decodeLines(Publisher> publisher, int maxSize) { + return Flux.defer(() -> { + Utf8LineDecoder dec = new Utf8LineDecoder(plusFramingOverhead(maxSize)); + return JdkFlowAdapter.flowPublisherToFlux(publisher) + .concatMapIterable(dec::decode) + .concatWith(Flux.defer(() -> Flux.fromIterable(dec.flush()))); + }); + } + + /** + * Parses a flux of SSE-formatted lines into a flux of {@link SseEvent}, bounding how + * much memory a single event may occupy. + * @param lines the SSE-formatted lines to parse + * @param maxSize the maximum number of bytes that may accumulate for a single SSE + * event + */ + static Flux decodeSseResponse(Flux lines, int maxSize) { + return Flux.defer(() -> { + SseEventParser parser = new SseEventParser(maxSize); + return lines.handle((line, sink) -> parser.feed(line).ifPresent(sink::next)) + .concatWith(Mono.defer(() -> parser.flush().map(Mono::just).orElseGet(Mono::empty))); + }); + } + + /** + * Collects all byte-buffer chunks from the publisher into a single UTF-8 decoded + * string, bounding how much memory it may occupy. A peer sending a body larger than + * {@code maxSize} has its response aborted instead of forcing the transport to buffer + * it in memory. + * @param publisher the response body + * @param maxSize the maximum number of bytes read for the response body + */ + static Mono decodeAggregateResponse(Publisher> publisher, int maxSize) { + return boundTotalBytes(publisher, maxSize).collectList().map(buffers -> { + int totalSize = buffers.stream().mapToInt(ByteBuffer::remaining).sum(); + ByteBuffer combined = ByteBuffer.allocate(totalSize); + buffers.forEach(combined::put); + combined.flip(); + return StandardCharsets.UTF_8.decode(combined).toString(); + }).defaultIfEmpty(""); + } + + /** + * Subscribes to the body publisher to release the underlying connection, discarding + * all bytes, then propagates the given error. + * + *

+ * Nothing accumulates here, so the bound is not protecting memory: it stops a peer + * from making the transport read an unbounded body only to throw it away. Should the + * body outgrow {@code maxSize}, or fail to be read, that failure is dropped and + * {@code error} is propagated all the same, as it is the reason the body is being + * discarded in the first place. + * @param body the response body + * @param maxSize the maximum number of bytes read for the response body + * @param error the error to propagate once the body has been discarded + */ + static Flux drainThenError(Publisher> body, int maxSize, Throwable error) { + return boundTotalBytes(body, maxSize).onErrorComplete().thenMany(Mono.error(error)); + } + + /** + * Reads the body as text, then propagates a {@link McpTransportException} carrying + * {@code message} followed by that text, so that what the server said about the + * failure reaches the caller. The body is read under the same bound as any other. + * @param body the response body + * @param maxSize the maximum number of bytes read for the response body + * @param message describes the failure the body explains + */ + static Flux readThenError(Publisher> body, int maxSize, String message) { + return decodeAggregateResponse(body, maxSize).flatMapMany(text -> Flux + .error(new McpTransportException(Utils.hasText(text) ? message + ", response body: " + text : message))); + } + + /** + * Subscribes to the body publisher to release the underlying connection, discarding + * all bytes, then completes empty. + * + *

+ * As in {@link #drainThenError}, the bound caps what a peer can make the transport + * read rather than what it can make it hold. + * @param body the response body + * @param maxSize the maximum number of bytes read for the response body + */ + static Flux drain(Publisher> body, int maxSize) { + return boundTotalBytes(body, maxSize).thenMany(Flux.empty()); + } + + /** + * Subscribes to the body publisher only to cancel it, which releases the underlying + * connection without reading the body, then completes empty. + * + *

+ * Unlike {@link #drain}, this suits a body that may never end, such as an SSE stream + * whose content is of no further interest. + * @param body the response body + */ + static Flux cancel(Publisher> body) { + return Flux.defer(() -> { + cancelBody(body); + return Flux.empty(); + }); + } + + static Mono>>> sendAsync(HttpClient httpClient, HttpRequest request) { + // Not Mono.fromFuture: cancelling aborts the exchange, and the HttpClient then + // fails the future with a CompletionException wrapping a CancellationException, + // which fromFuture reports as a dropped error. Only this method cna cancel the + // future, so that failure is ignored here. Replace with a plain fromFuture, + // keeping + // the doOnDiscard, once https://github.com/reactor/reactor-core/issues/4415 is + // resolved. + return Mono.>>>create(sink -> { + CompletableFuture>>> exchange = httpClient.sendAsync(request, + HttpResponse.BodyHandlers.ofPublisher()); + sink.onCancel(() -> exchange.cancel(true)); + exchange.whenComplete((response, error) -> { + if (error == null) { + // Emit the response so the body can be consumed. + // If the surrounding Mono was cancelled though and due to a race + // the headers were already parsed, the below call will simply + // discard the response. + sink.success(response); + return; + } + Throwable cause = error instanceof CompletionException && error.getCause() != null ? error.getCause() + : error; + if (cause instanceof CancellationException) { + sink.success(); + } + else { + sink.error(cause); + } + }); + }) + // A body that is never subscribed to never releases its connection. + .doOnDiscard(HttpResponse.class, response -> { + if (response.body() instanceof Publisher body) { + cancelBody(body); + } + }); + } + + private static void cancelBody(Publisher body) { + body.subscribe(CancellingSubscriber.INSTANCE); + } + + /** + * Flattens the body into its individual byte buffers, failing once more than + * {@code maxSize} bytes have passed through. Failing cancels the subscription, which + * closes the connection and so stops the peer from streaming any more. + */ + private static Flux boundTotalBytes(Publisher> body, int maxSize) { + return Flux.defer(() -> { + // Held in an array because the handle callback below cannot mutate a + // captured local. The enclosing defer gives each subscriber its own. + long[] totalBytes = new long[1]; + return JdkFlowAdapter.flowPublisherToFlux(body) + .flatMapIterable(list -> list) + .handle((buffer, sink) -> { + totalBytes[0] += buffer.remaining(); + if (totalBytes[0] > maxSize) { + sink.error(new McpTransportException( + "Inbound response body exceeds the maximum allowed size of " + maxSize + " bytes")); + return; + } + sink.next(buffer); + }); + }); + } + + /** + * Stateful UTF-8 decoder that splits a stream of byte-buffer chunks into complete + * lines. Handles multi-byte characters split across chunk boundaries, and terminates + * a line on {@code "\r\n"}, {@code "\r"} or {@code "\n"} alike, as the SSE wire + * format does. Bytes that do not decode are replaced rather than reported, so a peer + * sending one does not cost the stream. + */ + static final class Utf8LineDecoder { + + /** + * Undecodable input costs one replacement character rather than the stream: a + * decoder left on the default {@link CodingErrorAction#REPORT} fails the whole + * response over a single byte a peer mangled, and takes with it the lines already + * decoded from the same chunk, because {@link #decode(List)} throws instead of + * returning them. A body cut short mid-character is enough to hit it. This + * matches {@link java.net.http.HttpResponse.BodySubscribers#fromLineSubscriber}, + * the path this decoder replaces, which configured the same two actions. + */ + private final CharsetDecoder decoder = StandardCharsets.UTF_8.newDecoder() + .onMalformedInput(CodingErrorAction.REPLACE) + .onUnmappableCharacter(CodingErrorAction.REPLACE); + + private final CharBuffer charBuffer = CharBuffer.allocate(4096); + + private final StringBuilder leftover = new StringBuilder(); + + /** + * The maximum number of bytes a single line may occupy. Measured against + * {@link #leftover}'s length in characters, which for UTF-8 is never more than + * the number of bytes those characters were decoded from, so a line is only ever + * rejected once it has genuinely exceeded the bound in bytes. + */ + private final int maxSize; + + /** + * How many leading characters of {@link #leftover} are already known to hold no + * line terminator, so that the search for one resumes where the previous search + * ended instead of restarting at the beginning of the buffer. Without it, a long + * line is searched again in full for every chunk that arrives, which makes + * reading an event cost time proportional to the square of its length. + * @see #1042 + */ + private int scannedForLineTerminator = 0; + + /** + * Whether the line just emitted was terminated by a CR, so that a LF opening what + * follows completes that terminator instead of ending a line of its own. A CR is + * emitted on as soon as it arrives, before it is known whether a LF follows it, + * and the two may be split across chunks. + */ + private boolean crTerminatedPreviousLine = false; + + // Holds partial UTF-8 sequences left over from a previous chunk (max 3 bytes + // for a BMP code point; 4 bytes for a supplementary one). + private ByteBuffer pendingBytes = ByteBuffer.allocate(0); + + Utf8LineDecoder(int maxSize) { + this.maxSize = maxSize; + } + + List decode(List chunk) { + List lines = new ArrayList<>(); + for (ByteBuffer bb : chunk) { + ByteBuffer input = bb; + if (pendingBytes.hasRemaining()) { + ByteBuffer merged = ByteBuffer.allocate(pendingBytes.remaining() + bb.remaining()); + merged.put(pendingBytes).put(bb); + merged.flip(); + pendingBytes = ByteBuffer.allocate(0); + input = merged; + } + while (true) { + CoderResult result = decoder.decode(input, charBuffer, false); + drainCharBuffer(); + extractCompletedLines(lines); + // Unreachable while the decoder replaces undecodable input, but kept + // so that an error result cannot spin this loop: it is neither an + // underflow nor an overflow. + if (result.isError()) { + try { + result.throwException(); + } + catch (CharacterCodingException e) { + throw new RuntimeException(e); + } + } + if (result.isUnderflow()) { + if (input.hasRemaining()) { + pendingBytes = ByteBuffer.allocate(input.remaining()); + pendingBytes.put(input).flip(); + } + break; + } + } + } + return lines; + } + + List flush() { + ByteBuffer tail = pendingBytes.hasRemaining() ? pendingBytes : ByteBuffer.allocate(0); + CoderResult result = decoder.decode(tail, charBuffer, true); + while (result.isOverflow()) { + drainCharBuffer(); + result = decoder.decode(tail, charBuffer, true); + } + drainCharBuffer(); + if (result.isError()) { + try { + result.throwException(); + } + catch (CharacterCodingException e) { + throw new RuntimeException(e); + } + } + + result = decoder.flush(charBuffer); + while (result.isOverflow()) { + drainCharBuffer(); + result = decoder.flush(charBuffer); + } + drainCharBuffer(); + pendingBytes = ByteBuffer.allocate(0); + + List lines = new ArrayList<>(); + extractCompletedLines(lines); + if (leftover.length() > 0) { + String last = leftover.toString(); + leftover.setLength(0); + this.scannedForLineTerminator = 0; + lines.add(last); + } + this.crTerminatedPreviousLine = false; + return lines; + } + + private void drainCharBuffer() { + charBuffer.flip(); + leftover.append(charBuffer); + charBuffer.clear(); + } + + private void extractCompletedLines(List out) { + while (true) { + if (this.crTerminatedPreviousLine) { + if (leftover.length() == 0) { + // The LF, if there is one, is in a chunk that has not arrived. + return; + } + if (leftover.charAt(0) == '\n') { + leftover.delete(0, 1); + } + this.crTerminatedPreviousLine = false; + } + int terminatorIdx = indexOfLineTerminator(this.scannedForLineTerminator); + if (terminatorIdx == -1) { + this.scannedForLineTerminator = leftover.length(); + if (leftover.length() > this.maxSize) { + throw new McpTransportException( + "Inbound line exceeds the maximum allowed size of " + this.maxSize + " bytes"); + } + return; + } + out.add(leftover.substring(0, terminatorIdx)); + this.crTerminatedPreviousLine = leftover.charAt(terminatorIdx) == '\r'; + leftover.delete(0, terminatorIdx + 1); + // What is left starts after the terminator, so none of it has been + // searched yet. + this.scannedForLineTerminator = 0; + } + } + + /** + * Index of the first CR or LF in {@link #leftover} at or after {@code from}, or + * {@code -1} when there is none. + */ + private int indexOfLineTerminator(int from) { + for (int i = from; i < leftover.length(); i++) { + char c = leftover.charAt(i); + if (c == '\n' || c == '\r') { + return i; + } + } + return -1; + } + + } + + /** + * Stateful SSE line parser. Accumulates {@code data:}, {@code id:} and {@code event:} + * fields until a blank line dispatches the event. Per the SSE spec, {@code id} is the + * last event ID and persists across events until re-set, with an empty value clearing + * it; {@code event} and {@code data} are reset by every blank line, so an event that + * does not name its type is a {@code message} event whatever preceded it. A blank + * line dispatches only when a {@code data:} field was seen, whether or not it carried + * a value. Comments and fields the parser does not handle, such as {@code retry:}, + * are ignored as the spec requires. + * + * @see Interpreting + * an event stream + */ + static final class SseEventParser { + + private static final Logger logger = LoggerFactory.getLogger(SseEventParser.class); + + private final StringBuilder data = new StringBuilder(); + + /** + * The maximum number of bytes that may accumulate for a single SSE event. A peer + * that never terminates an event (e.g. an endless stream of {@code data:} lines) + * has its stream aborted instead of exhausting memory. The accumulated data is + * measured in characters, which for UTF-8 is never more than the number of bytes + * it was decoded from. + */ + private final int maxSize; + + private String id; + + private String event; + + SseEventParser(int maxSize) { + this.maxSize = maxSize; + } + + Optional feed(String line) { + if (line.isEmpty()) { + return flush(); + } + if (line.startsWith("data:")) { + // Every data field appends its value followed by a separator, so a + // valueless `data:` line still marks the event as carrying data and gets + // dispatched with empty data. Servers send such an event to prime a + // stream, and dropping it leaves the request it answers hanging. + String value = line.substring(5).trim(); + // Measured before appending, so that an event carrying exactly + // maxSize of data is accepted: the trailing separator below is + // stripped again before the event is emitted. + if (data.length() + value.length() > this.maxSize) { + throw new McpTransportException( + "Inbound SSE event exceeds the maximum allowed size of " + this.maxSize + " bytes"); + } + data.append(value).append('\n'); + } + else if (line.startsWith("id:")) { + String value = line.substring(3).trim(); + // The spec ignores an id carrying a NULL, and an empty id resets the last + // event ID, which leaves nothing to resume from. + if (value.indexOf('\0') == -1) { + id = value.isEmpty() ? null : value; + } + } + else if (line.startsWith("event:")) { + String value = line.substring(6).trim(); + event = value.isEmpty() ? null : value; + } + else if (line.startsWith(":")) { + logger.debug("Ignoring comment line: {}", line); + } + else { + // The SSE spec mandates that fields the client does not know about, such + // as `retry:`, are ignored rather than treated as a protocol error. + logger.debug("Ignoring unknown SSE field line: {}", line); + } + return Optional.empty(); + } + + /** + * Emits the pending event, if a {@code data:} field was seen, and resets the + * per-event state. The event type is reset even when nothing is dispatched, as + * the spec requires, while the id is the last event ID and so survives. An event + * that did not name its type is emitted as a {@code message} event. + */ + Optional flush() { + String type = this.event; + this.event = null; + if (data.isEmpty()) { + return Optional.empty(); + } + SseEvent result = new SseEvent(id, type != null ? type : DEFAULT_EVENT_TYPE, data.toString().trim()); + data.setLength(0); + return Optional.of(result); + } + + } + + private static class CancellingSubscriber implements Flow.Subscriber { + + private static final CancellingSubscriber INSTANCE = new CancellingSubscriber(); + + @Override + public void onSubscribe(Flow.Subscription subscription) { + subscription.cancel(); + } + + @Override + public void onNext(Object item) { + } + + @Override + public void onError(Throwable throwable) { + } + + @Override + public void onComplete() { + } + + } + +} diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java deleted file mode 100644 index b19904de6..000000000 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java +++ /dev/null @@ -1,619 +0,0 @@ -/* -* Copyright 2024 - 2024 the original author or authors. -*/ - -package io.modelcontextprotocol.client.transport; - -import java.net.http.HttpResponse; -import java.net.http.HttpResponse.BodyHandler; -import java.net.http.HttpResponse.BodySubscriber; -import java.net.http.HttpResponse.ResponseInfo; -import java.nio.ByteBuffer; -import java.util.List; -import java.util.concurrent.CompletionStage; -import java.util.concurrent.Flow; -import java.util.concurrent.atomic.AtomicReference; -import java.util.regex.Pattern; - -import org.reactivestreams.FlowAdapters; -import org.reactivestreams.Subscription; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - -import io.modelcontextprotocol.spec.McpTransportException; -import reactor.core.publisher.BaseSubscriber; -import reactor.core.publisher.FluxSink; - -/** - * Utility class providing various {@link BodySubscriber} implementations for handling - * different types of HTTP response bodies in the context of Model Context Protocol (MCP) - * clients. - * - *

- * Defines subscribers for processing Server-Sent Events (SSE), aggregate responses, and - * bodiless responses. - * - * @author Christian Tzolov - * @author Dariusz Jędrzejczyk - * @author Daniel Garnier-Moiroux - */ -class ResponseSubscribers { - - private static final Logger logger = LoggerFactory.getLogger(ResponseSubscribers.class); - - /** - * Bytes of SSE field framing a single line may carry on top of the message payload: - * {@code "event: "} is the longest field prefix this parser recognises. Line - * terminators are not counted, as they reset the running line length. Without this - * allowance, an event carrying exactly the maximum message size would be rejected - * because of the bytes the SSE wire format adds around it. - */ - private static final int SSE_FRAMING_OVERHEAD = "event: ".length(); - - record SseEvent(String id, String event, String data) { - } - - sealed interface ResponseEvent permits SseResponseEvent, AggregateResponseEvent, DummyEvent { - - ResponseInfo responseInfo(); - - } - - record DummyEvent(ResponseInfo responseInfo) implements ResponseEvent { - - } - - record SseResponseEvent(ResponseInfo responseInfo, SseEvent sseEvent) implements ResponseEvent { - } - - record AggregateResponseEvent(ResponseInfo responseInfo, String data) implements ResponseEvent { - } - - /** - * Creates a {@link BodySubscriber} that parses a Server-Sent Events stream, bounding - * how much memory a single inbound message may occupy. Both the size of an individual - * line (as read off the wire before a terminator is seen) and the accumulated size of - * a multi-line SSE event are capped at {@code maxSize}; a peer exceeding either limit - * has its stream aborted instead of forcing the transport to buffer it in memory. The - * line bound is allowed {@link #SSE_FRAMING_OVERHEAD} extra bytes so that the SSE - * framing around a payload does not count against the payload's own budget. - * @param responseInfo the HTTP response information - * @param sink the sink to emit parsed events to - * @param maxSize the maximum number of bytes read for a single inbound message - */ - static BodySubscriber sseToBodySubscriber(ResponseInfo responseInfo, FluxSink sink, - int maxSize) { - BodySubscriber lineSubscriber = HttpResponse.BodySubscribers - .fromLineSubscriber(FlowAdapters.toFlowSubscriber(new SseLineSubscriber(responseInfo, sink, maxSize))); - return new BoundedLineBodySubscriber(lineSubscriber, plusFramingOverhead(maxSize)); - } - - /** - * Adds {@link #SSE_FRAMING_OVERHEAD} to {@code maxSize}, saturating at - * {@link Integer#MAX_VALUE} rather than overflowing into a negative bound that would - * reject everything. - */ - private static int plusFramingOverhead(int maxSize) { - return maxSize > Integer.MAX_VALUE - SSE_FRAMING_OVERHEAD ? Integer.MAX_VALUE : maxSize + SSE_FRAMING_OVERHEAD; - } - - /** - * Creates a {@link BodySubscriber} that aggregates the whole response body into a - * single event, bounding how much memory it may occupy. Both the size of an - * individual line (as read off the wire before a terminator is seen) and the total - * accumulated body are capped at {@code maxSize}; a peer exceeding either limit has - * its response aborted instead of forcing the transport to buffer it in memory. - * @param responseInfo the HTTP response information - * @param sink the sink to emit the aggregated event to - * @param maxSize the maximum number of bytes read for the response body - */ - static BodySubscriber aggregateBodySubscriber(ResponseInfo responseInfo, FluxSink sink, - int maxSize) { - BodySubscriber lineSubscriber = HttpResponse.BodySubscribers - .fromLineSubscriber(FlowAdapters.toFlowSubscriber(new AggregateSubscriber(responseInfo, sink, maxSize))); - return new BoundedLineBodySubscriber(lineSubscriber, maxSize); - } - - /** - * Creates a {@link BodySubscriber} that discards the response body, bounding how much - * memory reading it may occupy. The body is discarded as it arrives, but the - * underlying line subscriber still buffers each line before handing it over, so a - * peer sending a line longer than {@code maxSize} has its response aborted. - * @param responseInfo the HTTP response information - * @param sink the sink to emit the completion event to - * @param maxSize the maximum number of bytes read for a single line - */ - static BodySubscriber bodilessBodySubscriber(ResponseInfo responseInfo, FluxSink sink, - int maxSize) { - BodySubscriber lineSubscriber = HttpResponse.BodySubscribers - .fromLineSubscriber(FlowAdapters.toFlowSubscriber(new BodilessResponseLineSubscriber(responseInfo, sink))); - return new BoundedLineBodySubscriber(lineSubscriber, maxSize); - } - - /** - * Creates a {@link BodyHandler} that reads the response body into a string, bounding - * how much memory it may occupy. A peer sending more than {@code maxSize} bytes has - * its response aborted instead of forcing the transport to buffer it in memory. - * - *

- * Decoding matches {@link HttpResponse.BodyHandlers#ofString()}, including its - * handling of the charset declared in the {@code Content-Type} header. - * @param maxSize the maximum number of bytes read for the response body - */ - static BodyHandler boundedStringBodyHandler(int maxSize) { - BodyHandler delegate = HttpResponse.BodyHandlers.ofString(); - return responseInfo -> new BoundedTotalBodySubscriber<>(delegate.apply(responseInfo), maxSize); - } - - static class SseLineSubscriber extends BaseSubscriber { - - /** - * Pattern to extract data content from SSE "data:" lines. - */ - private static final Pattern EVENT_DATA_PATTERN = Pattern.compile("^data:(.+)$", Pattern.MULTILINE); - - /** - * Pattern to extract event ID from SSE "id:" lines. - */ - private static final Pattern EVENT_ID_PATTERN = Pattern.compile("^id:(.+)$", Pattern.MULTILINE); - - /** - * Pattern to extract event type from SSE "event:" lines. - */ - private static final Pattern EVENT_TYPE_PATTERN = Pattern.compile("^event:(.+)$", Pattern.MULTILINE); - - /** - * The sink for emitting parsed response events. - */ - private final FluxSink sink; - - /** - * StringBuilder for accumulating multi-line event data. - */ - private final StringBuilder eventBuilder; - - /** - * Current event's ID, if specified. - */ - private final AtomicReference currentEventId; - - /** - * Current event's type, if specified. - */ - private final AtomicReference currentEventType; - - /** - * The response information from the HTTP response. Send with each event to - * provide context. - */ - private ResponseInfo responseInfo; - - /** - * The maximum number of bytes that may accumulate for a single SSE event. A peer - * that never terminates an event (e.g. an endless stream of {@code data:} lines) - * has its stream aborted instead of exhausting memory. The accumulated data is - * measured in characters, which for UTF-8 is never more than the number of bytes - * it was decoded from. - */ - private final int maxSize; - - /** - * Creates a new LineSubscriber that will emit parsed SSE events to the provided - * sink. - * @param sink the {@link FluxSink} to emit parsed {@link ResponseEvent} objects - * to - * @param maxSize the maximum number of bytes that may accumulate for a single SSE - * event - */ - public SseLineSubscriber(ResponseInfo responseInfo, FluxSink sink, int maxSize) { - this.sink = sink; - this.eventBuilder = new StringBuilder(); - this.currentEventId = new AtomicReference<>(); - this.currentEventType = new AtomicReference<>(); - this.responseInfo = responseInfo; - this.maxSize = maxSize; - } - - @Override - protected void hookOnSubscribe(Subscription subscription) { - - sink.onRequest(n -> { - subscription.request(n); - }); - - // Register disposal callback to cancel subscription when Flux is disposed - sink.onDispose(() -> { - subscription.cancel(); - }); - } - - @Override - protected void hookOnNext(String line) { - if (line.isEmpty()) { - // Empty line means end of event - if (this.eventBuilder.length() > 0) { - String eventData = this.eventBuilder.toString(); - SseEvent sseEvent = new SseEvent(currentEventId.get(), currentEventType.get(), eventData.trim()); - - this.sink.next(new SseResponseEvent(responseInfo, sseEvent)); - this.eventBuilder.setLength(0); - } - } - else { - if (line.startsWith("data:")) { - var matcher = EVENT_DATA_PATTERN.matcher(line); - if (matcher.find()) { - String data = matcher.group(1).trim(); - // Measured before appending, so that an event carrying exactly - // maxSize of data is accepted: the trailing separator below is - // stripped again before the event is emitted. - if (this.eventBuilder.length() + data.length() > this.maxSize) { - upstream().cancel(); - this.sink.error( - new McpTransportException("Inbound SSE event exceeds the maximum allowed size of " - + this.maxSize + " bytes")); - return; - } - this.eventBuilder.append(data).append("\n"); - } - upstream().request(1); - } - else if (line.startsWith("id:")) { - var matcher = EVENT_ID_PATTERN.matcher(line); - if (matcher.find()) { - this.currentEventId.set(matcher.group(1).trim()); - } - upstream().request(1); - } - else if (line.startsWith("event:")) { - var matcher = EVENT_TYPE_PATTERN.matcher(line); - if (matcher.find()) { - this.currentEventType.set(matcher.group(1).trim()); - } - upstream().request(1); - } - else if (line.startsWith(":")) { - // Ignore comment lines starting with ":" - // This is a no-op, just to skip comments - logger.debug("Ignoring comment line: {}", line); - upstream().request(1); - } - else { - // If the response is not successful, emit an error - this.sink.error(new McpTransportException( - "Invalid SSE response. Status code: " + this.responseInfo.statusCode() + " Line: " + line)); - - } - } - } - - @Override - protected void hookOnComplete() { - if (this.eventBuilder.length() > 0) { - String eventData = this.eventBuilder.toString(); - SseEvent sseEvent = new SseEvent(currentEventId.get(), currentEventType.get(), eventData.trim()); - this.sink.next(new SseResponseEvent(responseInfo, sseEvent)); - } - this.sink.complete(); - } - - @Override - protected void hookOnError(Throwable throwable) { - this.sink.error(throwable); - } - - } - - static class AggregateSubscriber extends BaseSubscriber { - - /** - * The sink for emitting parsed response events. - */ - private final FluxSink sink; - - /** - * StringBuilder for accumulating multi-line event data. - */ - private final StringBuilder eventBuilder; - - /** - * The response information from the HTTP response. Send with each event to - * provide context. - */ - private ResponseInfo responseInfo; - - volatile boolean hasRequestedDemand = false; - - /** - * The maximum number of bytes that may accumulate for the aggregated response - * body. A peer that sends a larger body has its response aborted instead of - * exhausting memory. The accumulated body is measured in characters, which for - * UTF-8 is never more than the number of bytes it was decoded from. - */ - private final int maxSize; - - /** - * Creates a new JsonLineSubscriber that will emit parsed JSON-RPC messages. - * @param sink the {@link FluxSink} to emit parsed {@link ResponseEvent} objects - * to - * @param maxSize the maximum number of bytes that may accumulate for the - * aggregated response body - */ - public AggregateSubscriber(ResponseInfo responseInfo, FluxSink sink, int maxSize) { - this.sink = sink; - this.eventBuilder = new StringBuilder(); - this.responseInfo = responseInfo; - this.maxSize = maxSize; - } - - @Override - protected void hookOnSubscribe(Subscription subscription) { - - sink.onRequest(n -> { - if (!hasRequestedDemand) { - subscription.request(Long.MAX_VALUE); - } - hasRequestedDemand = true; - }); - - // Register disposal callback to cancel subscription when Flux is disposed - sink.onDispose(subscription::cancel); - } - - @Override - protected void hookOnNext(String line) { - // Measured before appending, so that a body of exactly maxSize is accepted. - // The separator this adds back for each line stands in for the terminator the - // peer sent, which the line subscriber has already stripped. - if (this.eventBuilder.length() + line.length() > this.maxSize) { - upstream().cancel(); - this.sink.error(new McpTransportException( - "Inbound response body exceeds the maximum allowed size of " + this.maxSize + " bytes")); - return; - } - this.eventBuilder.append(line).append("\n"); - } - - @Override - protected void hookOnComplete() { - - if (hasRequestedDemand) { - String data = this.eventBuilder.toString(); - this.sink.next(new AggregateResponseEvent(responseInfo, data)); - } - - this.sink.complete(); - } - - @Override - protected void hookOnError(Throwable throwable) { - this.sink.error(throwable); - } - - } - - static class BodilessResponseLineSubscriber extends BaseSubscriber { - - /** - * The sink for emitting parsed response events. - */ - private final FluxSink sink; - - private final ResponseInfo responseInfo; - - volatile boolean hasRequestedDemand = false; - - public BodilessResponseLineSubscriber(ResponseInfo responseInfo, FluxSink sink) { - this.sink = sink; - this.responseInfo = responseInfo; - } - - @Override - protected void hookOnSubscribe(Subscription subscription) { - - sink.onRequest(n -> { - if (!hasRequestedDemand) { - subscription.request(Long.MAX_VALUE); - } - hasRequestedDemand = true; - }); - - // Register disposal callback to cancel subscription when Flux is disposed - sink.onDispose(() -> { - subscription.cancel(); - }); - } - - @Override - protected void hookOnComplete() { - if (hasRequestedDemand) { - // emit dummy event to be able to inspect the response info - // this is a shortcut allowing for a more streamlined processing using - // operator composition instead of having to deal with the - // CompletableFuture along the Subscriber for inspecting the result - this.sink.next(new DummyEvent(responseInfo)); - } - this.sink.complete(); - } - - @Override - protected void hookOnError(Throwable throwable) { - this.sink.error(throwable); - } - - } - - /** - * Base for {@link BodySubscriber} wrappers that transparently forward the response - * body to a delegate, but abort it once the peer exceeds a size bound. - * - *

- * Aborting cancels the upstream subscription, which closes the connection, and - * signals a {@link McpTransportException} to the delegate so the failure surfaces - * both through the body's {@link CompletionStage} and through any sink the delegate - * feeds. - */ - abstract static class BoundedBodySubscriber implements BodySubscriber { - - private final BodySubscriber delegate; - - protected final int maxSize; - - /** - * What the bound applies to, e.g. {@code "Inbound line"}, used to build the - * failure message. - */ - private final String boundedEntity; - - private Flow.Subscription subscription; - - private volatile boolean done = false; - - BoundedBodySubscriber(BodySubscriber delegate, int maxSize, String boundedEntity) { - this.delegate = delegate; - this.maxSize = maxSize; - this.boundedEntity = boundedEntity; - } - - @Override - public CompletionStage getBody() { - return this.delegate.getBody(); - } - - @Override - public void onSubscribe(Flow.Subscription subscription) { - this.subscription = subscription; - this.delegate.onSubscribe(subscription); - } - - @Override - public void onNext(List buffers) { - if (this.done) { - return; - } - for (ByteBuffer buffer : buffers) { - if (!checkSize(buffer)) { - this.done = true; - this.subscription.cancel(); - this.delegate.onError(new McpTransportException( - this.boundedEntity + " exceeds the maximum allowed size of " + this.maxSize + " bytes")); - return; - } - } - this.delegate.onNext(buffers); - } - - /** - * Accounts for the bytes in {@code buffer}, which must be inspected with absolute - * reads only so the delegate still sees the original position. - * @param buffer the buffer about to be handed to the delegate - * @return {@code true} to accept the buffer, or {@code false} to abort the - * response because the bound has been exceeded - */ - protected abstract boolean checkSize(ByteBuffer buffer); - - @Override - public void onError(Throwable throwable) { - if (this.done) { - return; - } - this.done = true; - this.delegate.onError(throwable); - } - - @Override - public void onComplete() { - if (this.done) { - return; - } - this.done = true; - this.delegate.onComplete(); - } - - } - - /** - * A {@link BoundedBodySubscriber} that aborts the response once a single line (a run - * of bytes with no CR/LF terminator) exceeds {@code maxSize} bytes. - * - *

- * {@link HttpResponse.BodySubscribers#fromLineSubscriber} buffers characters until it - * encounters a line terminator, so a peer that never terminates a line (or sends an - * enormous one) would force the transport to buffer it in memory. This wrapper counts - * bytes as they arrive off the wire and cancels the subscription before that buffer - * can grow without bound. - */ - static final class BoundedLineBodySubscriber extends BoundedBodySubscriber { - - private long bytesSinceLineTerminator = 0; - - BoundedLineBodySubscriber(BodySubscriber delegate, int maxSize) { - super(delegate, maxSize, "Inbound line"); - } - - @Override - protected boolean checkSize(ByteBuffer buffer) { - int position = buffer.position(); - int limit = buffer.limit(); - if (position == limit) { - return true; - } - if (this.bytesSinceLineTerminator + (limit - position) <= this.maxSize) { - // No line ending in this buffer can exceed the limit, because there are - // not enough bytes since the last terminator for one to. Only the - // trailing (still unterminated) run matters, so scan back to the last - // terminator instead of walking every byte. - this.bytesSinceLineTerminator = lengthOfTrailingRun(buffer, position, limit); - return true; - } - // The limit is within reach, so account for every line exactly. - for (int i = position; i < limit; i++) { - byte b = buffer.get(i); - if (b == '\n' || b == '\r') { - this.bytesSinceLineTerminator = 0; - } - else if (++this.bytesSinceLineTerminator > this.maxSize) { - return false; - } - } - return true; - } - - /** - * Returns the number of bytes after the last line terminator in the buffer, or - * the whole span added to the running count when the buffer holds no terminator. - */ - private long lengthOfTrailingRun(ByteBuffer buffer, int position, int limit) { - for (int i = limit - 1; i >= position; i--) { - byte b = buffer.get(i); - if (b == '\n' || b == '\r') { - return limit - 1 - i; - } - } - return this.bytesSinceLineTerminator + (limit - position); - } - - } - - /** - * A {@link BoundedBodySubscriber} that aborts the response once the body as a whole - * exceeds {@code maxSize} bytes. Suitable for delegates that aggregate the entire - * body in memory, such as {@link HttpResponse.BodyHandlers#ofString()}. - */ - static final class BoundedTotalBodySubscriber extends BoundedBodySubscriber { - - private long totalBytes = 0; - - BoundedTotalBodySubscriber(BodySubscriber delegate, int maxSize) { - super(delegate, maxSize, "Inbound response body"); - } - - @Override - protected boolean checkSize(ByteBuffer buffer) { - this.totalBytes += buffer.remaining(); - return this.totalBytes <= this.maxSize; - } - - } - -} diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/spec/DefaultMcpTransportSession.java b/mcp-core/src/main/java/io/modelcontextprotocol/spec/DefaultMcpTransportSession.java index fdb7bfd89..bfd71549f 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/spec/DefaultMcpTransportSession.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/DefaultMcpTransportSession.java @@ -78,6 +78,7 @@ public void close() { @Override public Mono closeGracefully() { return Mono.from(this.onClose.apply(this.sessionId.get())) + .onErrorResume(error -> Mono.fromRunnable(this.openConnections::dispose).then(Mono.error(error))) .then(Mono.fromRunnable(this.openConnections::dispose)); } diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/BoundedBodySubscriberTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/BoundedBodySubscriberTests.java deleted file mode 100644 index 5f350319e..000000000 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/BoundedBodySubscriberTests.java +++ /dev/null @@ -1,334 +0,0 @@ -/* - * Copyright 2024-2026 the original author or authors. - */ - -package io.modelcontextprotocol.client.transport; - -import java.net.http.HttpResponse.BodySubscriber; -import java.nio.ByteBuffer; -import java.nio.charset.StandardCharsets; -import java.util.ArrayList; -import java.util.List; -import java.util.concurrent.CompletableFuture; -import java.util.concurrent.CompletionStage; -import java.util.concurrent.Flow; - -import io.modelcontextprotocol.client.transport.ResponseSubscribers.BoundedLineBodySubscriber; -import io.modelcontextprotocol.client.transport.ResponseSubscribers.BoundedTotalBodySubscriber; -import io.modelcontextprotocol.spec.McpTransportException; -import org.junit.jupiter.api.Test; - -import static org.assertj.core.api.Assertions.assertThat; - -/** - * Tests the size accounting in {@link ResponseSubscribers.BoundedBodySubscriber} and its - * two implementations. These bound how much of a response the transport will buffer, so - * the accounting is exercised directly rather than only through a live HTTP exchange: - * buffer boundaries, line terminators split across buffers, and the exact limit are all - * places where an off-by-one either lets a peer past the bound or rejects a legitimate - * message. - * - * @author Daniel Garnier-Moiroux - */ -class BoundedBodySubscriberTests { - - private static final int MAX_SIZE = 16; - - private final RecordingBodySubscriber delegate = new RecordingBodySubscriber(); - - private final RecordingSubscription subscription = new RecordingSubscription(); - - // --- BoundedLineBodySubscriber: per-line accounting ----------------------- - - @Test - void lineSubscriberAcceptsEmptyBuffer() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer(""))).isTrue(); - } - - @Test - void lineSubscriberAcceptsLineOfExactlyMaxSize() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(MAX_SIZE)))).isTrue(); - } - - @Test - void lineSubscriberRejectsLineOneByteOverMaxSize() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(MAX_SIZE + 1)))).isFalse(); - } - - @Test - void lineSubscriberAccumulatesAcrossBuffers() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(10)))).isTrue(); - assertThat(subscriber.checkSize(buffer("a".repeat(6)))).isTrue(); - // 17th byte of the same unterminated line. - assertThat(subscriber.checkSize(buffer("a"))).isFalse(); - } - - @Test - void lineSubscriberAcceptsUnboundedTotalOfTerminatedLines() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - // Far more than MAX_SIZE in total, but no single line comes close to it. - for (int i = 0; i < 100; i++) { - assertThat(subscriber.checkSize(buffer("aaaa\n"))).isTrue(); - } - } - - @Test - void lineSubscriberResetsOnTerminatorAtEndOfBuffer() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(10) + "\n"))).isTrue(); - // A fresh line, so the previous 10 bytes must not count towards it. - assertThat(subscriber.checkSize(buffer("a".repeat(MAX_SIZE)))).isTrue(); - } - - @Test - void lineSubscriberResetsOnTerminatorAtStartOfBuffer() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(MAX_SIZE)))).isTrue(); - assertThat(subscriber.checkSize(buffer("\n" + "a".repeat(MAX_SIZE)))).isTrue(); - } - - @Test - void lineSubscriberHandlesCrLfSplitAcrossBuffers() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(12) + "\r"))).isTrue(); - assertThat(subscriber.checkSize(buffer("\n" + "a".repeat(MAX_SIZE)))).isTrue(); - } - - @Test - void lineSubscriberAcceptsBufferLargerThanMaxSizeHoldingOnlyShortLines() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - // Forces the exact per-byte accounting path: the buffer alone is well over the - // limit, yet every line in it is legitimate. - assertThat(subscriber.checkSize(buffer("aaaa\n".repeat(20)))).isTrue(); - } - - @Test - void lineSubscriberRejectsRunSpanningManyBuffers() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - boolean accepted = true; - for (int i = 0; i < 10 && accepted; i++) { - accepted = subscriber.checkSize(buffer("aa")); - } - - assertThat(accepted).isFalse(); - } - - @Test - void lineSubscriberOnlyCountsFromTheBufferPosition() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - ByteBuffer partiallyConsumed = buffer("a".repeat(MAX_SIZE * 2)); - partiallyConsumed.position(MAX_SIZE * 2 - 4); - - assertThat(subscriber.checkSize(partiallyConsumed)).isTrue(); - } - - @Test - void lineSubscriberDoesNotConsumeTheBuffer() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - ByteBuffer buffer = buffer("aaaa\nbbbb"); - buffer.position(2); - - subscriber.checkSize(buffer); - - assertThat(buffer.position()).isEqualTo(2); - assertThat(buffer.limit()).isEqualTo(9); - } - - // --- BoundedTotalBodySubscriber: whole-body accounting -------------------- - - @Test - void totalSubscriberAcceptsBodyOfExactlyMaxSize() { - BoundedTotalBodySubscriber subscriber = totalSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(8)))).isTrue(); - assertThat(subscriber.checkSize(buffer("a".repeat(8)))).isTrue(); - } - - @Test - void totalSubscriberRejectsBodyOneByteOverMaxSize() { - BoundedTotalBodySubscriber subscriber = totalSubscriber(); - - assertThat(subscriber.checkSize(buffer("a".repeat(8)))).isTrue(); - assertThat(subscriber.checkSize(buffer("a".repeat(9)))).isFalse(); - } - - @Test - void totalSubscriberIsNotResetByLineTerminators() { - BoundedTotalBodySubscriber subscriber = totalSubscriber(); - - // Unlike the per-line bound, terminated lines still count towards the total. - assertThat(subscriber.checkSize(buffer("aaaa\n".repeat(3)))).isTrue(); - assertThat(subscriber.checkSize(buffer("aaaa\n"))).isFalse(); - } - - @Test - void totalSubscriberOnlyCountsFromTheBufferPosition() { - BoundedTotalBodySubscriber subscriber = totalSubscriber(); - ByteBuffer partiallyConsumed = buffer("a".repeat(MAX_SIZE * 2)); - partiallyConsumed.position(MAX_SIZE); - - assertThat(subscriber.checkSize(partiallyConsumed)).isTrue(); - } - - @Test - void totalSubscriberDoesNotConsumeTheBuffer() { - BoundedTotalBodySubscriber subscriber = totalSubscriber(); - ByteBuffer buffer = buffer("aaaa"); - - subscriber.checkSize(buffer); - - assertThat(buffer.position()).isZero(); - assertThat(buffer.remaining()).isEqualTo(4); - } - - // --- onNext: what a failed check does ------------------------------------ - - @Test - void forwardsBuffersWhileWithinBounds() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - List buffers = List.of(buffer("aaaa\n"), buffer("bbbb\n")); - - subscriber.onNext(buffers); - - assertThat(this.delegate.received).containsExactly(buffers); - assertThat(this.delegate.error).isNull(); - assertThat(this.subscription.cancellations).isZero(); - } - - @Test - void abortsTheResponseWhenTheLineBoundIsExceeded() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - subscriber.onNext(List.of(buffer("a".repeat(MAX_SIZE + 1)))); - - assertThat(this.subscription.cancellations).isEqualTo(1); - assertThat(this.delegate.received).isEmpty(); - assertThat(this.delegate.error).isInstanceOf(McpTransportException.class) - .hasMessage("Inbound line exceeds the maximum allowed size of " + MAX_SIZE + " bytes"); - } - - @Test - void abortsTheResponseWhenTheTotalBoundIsExceeded() { - BoundedTotalBodySubscriber subscriber = totalSubscriber(); - - subscriber.onNext(List.of(buffer("a".repeat(MAX_SIZE + 1)))); - - assertThat(this.subscription.cancellations).isEqualTo(1); - assertThat(this.delegate.received).isEmpty(); - assertThat(this.delegate.error).isInstanceOf(McpTransportException.class) - .hasMessage("Inbound response body exceeds the maximum allowed size of " + MAX_SIZE + " bytes"); - } - - @Test - void withholdsTheWholeListWhenALaterBufferExceedsTheBound() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - - subscriber.onNext(List.of(buffer("a".repeat(8)), buffer("a".repeat(9)))); - - assertThat(this.delegate.received).isEmpty(); - assertThat(this.delegate.error).isInstanceOf(McpTransportException.class); - } - - @Test - void ignoresFurtherSignalsOnceAborted() { - BoundedLineBodySubscriber subscriber = lineSubscriber(); - subscriber.onNext(List.of(buffer("a".repeat(MAX_SIZE + 1)))); - Throwable firstError = this.delegate.error; - - // The HTTP client may still signal after the subscription is cancelled. - subscriber.onNext(List.of(buffer("aaaa"))); - subscriber.onError(new RuntimeException("late failure")); - subscriber.onComplete(); - - assertThat(this.subscription.cancellations).isEqualTo(1); - assertThat(this.delegate.received).isEmpty(); - assertThat(this.delegate.error).isSameAs(firstError); - assertThat(this.delegate.completed).isFalse(); - } - - // --- fixtures ------------------------------------------------------------ - - private BoundedLineBodySubscriber lineSubscriber() { - BoundedLineBodySubscriber subscriber = new BoundedLineBodySubscriber(this.delegate, MAX_SIZE); - subscriber.onSubscribe(this.subscription); - return subscriber; - } - - private BoundedTotalBodySubscriber totalSubscriber() { - BoundedTotalBodySubscriber subscriber = new BoundedTotalBodySubscriber<>(this.delegate, MAX_SIZE); - subscriber.onSubscribe(this.subscription); - return subscriber; - } - - private static ByteBuffer buffer(String content) { - return ByteBuffer.wrap(content.getBytes(StandardCharsets.US_ASCII)); - } - - private static final class RecordingBodySubscriber implements BodySubscriber { - - private final List> received = new ArrayList<>(); - - private final CompletableFuture body = new CompletableFuture<>(); - - private Throwable error; - - private boolean completed; - - @Override - public CompletionStage getBody() { - return this.body; - } - - @Override - public void onSubscribe(Flow.Subscription subscription) { - } - - @Override - public void onNext(List item) { - this.received.add(item); - } - - @Override - public void onError(Throwable throwable) { - this.error = throwable; - this.body.completeExceptionally(throwable); - } - - @Override - public void onComplete() { - this.completed = true; - this.body.complete(null); - } - - } - - private static final class RecordingSubscription implements Flow.Subscription { - - private int cancellations; - - @Override - public void request(long n) { - } - - @Override - public void cancel() { - this.cancellations++; - } - - } - -} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientHttpTransportLeakTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientHttpTransportLeakTests.java new file mode 100644 index 000000000..28e2fb46a --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientHttpTransportLeakTests.java @@ -0,0 +1,107 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.util.List; +import java.util.Map; +import java.util.function.Function; +import java.util.stream.Stream; + +import io.modelcontextprotocol.spec.McpClientTransport; +import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.json.gson.GsonMcpJsonMapper; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import reactor.test.StepVerifier; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.junit.jupiter.api.Named.named; +import static org.junit.jupiter.params.provider.Arguments.arguments; + +class HttpClientHttpTransportLeakTests { + + static int selectorManagerThreadCount() { + return selectorManagerThreadNames().size(); + } + + static List selectorManagerThreadNames() { + return Thread.getAllStackTraces() + .keySet() + .stream() + .map(Thread::getName) + .filter(name -> name.contains("HttpClient") && name.contains("SelectorManager")) + .sorted() + .toList(); + } + + static int forceGcUntilStable() throws InterruptedException { + int previousCount = Integer.MAX_VALUE; + int stableIterations = 0; + int currentCount = previousCount; + + for (int i = 0; i < 40; i++) { + System.gc(); + System.runFinalization(); + Thread.sleep(250); + + currentCount = selectorManagerThreadCount(); + if (currentCount == previousCount) { + stableIterations++; + if (stableIterations >= 4) { + break; + } + } + else { + stableIterations = 0; + previousCount = currentCount; + } + } + + return currentCount; + } + + static void pauseForSelectorStartup() throws InterruptedException { + Thread.sleep(150); + } + + @ParameterizedTest + @MethodSource("httpTransports") + void closeDoesNotRetainOwnedHttpClient(Function httpTransportBuilder) throws Exception { + try (LoopbackMcpHttpServer server = LoopbackMcpHttpServer.start()) { + int selectorThreadsBefore = selectorManagerThreadCount(); + Function, reactor.core.publisher.Mono> handler = Function + .identity(); + + for (int i = 0; i < 12; i++) { + McpClientTransport transport = httpTransportBuilder.apply(server.baseUri().toString()); + + StepVerifier.create(transport.connect(handler)).verifyComplete(); + StepVerifier.create(transport.sendMessage( + new McpSchema.JSONRPCNotification(McpSchema.JSONRPC_VERSION, "ping", Map.of("iteration", i)))) + .verifyComplete(); + pauseForSelectorStartup(); + StepVerifier.create(transport.closeGracefully()).verifyComplete(); + } + + int selectorThreadsAfter = forceGcUntilStable(); + + assertThat(selectorThreadsAfter) + .describedAs( + "closed transports should not keep owned HttpClient instances alive, remaining threads: %s", + selectorManagerThreadNames()) + .isLessThanOrEqualTo(selectorThreadsBefore + 1); + } + } + + static Stream httpTransports() { + Function streamableHttp = ( + uri) -> HttpClientStreamableHttpTransport.builder(uri).jsonMapper(new GsonMcpJsonMapper()).build(); + Function sse = ( + uri) -> HttpClientSseClientTransport.builder(uri).jsonMapper(new GsonMcpJsonMapper()).build(); + return Stream.of(arguments(named("Streamable HTTP", streamableHttp)), arguments(named("SSE", sse))); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportConnectTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportConnectTests.java new file mode 100644 index 000000000..9d25cd4e0 --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportConnectTests.java @@ -0,0 +1,109 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.IOException; +import java.net.InetSocketAddress; +import java.time.Duration; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.function.Function; + +import com.sun.net.httpserver.HttpHandler; +import com.sun.net.httpserver.HttpServer; +import io.modelcontextprotocol.spec.McpTransportException; +import io.modelcontextprotocol.spec.json.gson.GsonMcpJsonMapper; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import reactor.test.StepVerifier; + +import static org.assertj.core.api.Assertions.assertThat; + +class HttpClientSseClientTransportConnectTests { + + // Only bounds a regression: every test resolves without waiting on it. + private static final Duration TIMEOUT = Duration.ofSeconds(5); + + private final ExecutorService executor = Executors.newCachedThreadPool(); + + private final CountDownLatch releaseResponse = new CountDownLatch(1); + + private HttpServer server; + + @AfterEach + void tearDown() { + this.releaseResponse.countDown(); + if (this.server != null) { + this.server.stop(0); + } + this.executor.shutdownNow(); + } + + @Test + void connectFailsWhenStreamEndsBeforeAnyEvent() throws IOException { + HttpClientSseClientTransport transport = transport(exchange -> { + exchange.getResponseHeaders().add("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, -1); + exchange.close(); + }); + + StepVerifier.create(transport.connect(Function.identity())) + .expectErrorSatisfies(e -> assertThat(e).isInstanceOf(McpTransportException.class) + .hasMessageContaining("before any event")) + .verify(TIMEOUT); + } + + @Test + void connectFailsWhenStreamErrorsBeforeAnyEventWhileClosing() throws IOException { + HttpClientSseClientTransport transport = transport(exchange -> { + // The server drops the connection without responding. + throw new IOException("dropped"); + }); + transport.closeGracefully().block(TIMEOUT); + + StepVerifier.create(transport.connect(Function.identity())).expectError().verify(TIMEOUT); + } + + @Test + void connectCompletesWhenClosedBeforeAnyEvent() throws IOException { + CountDownLatch streamOpened = new CountDownLatch(1); + HttpClientSseClientTransport transport = transport(exchange -> { + exchange.getResponseHeaders().add("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, 0); + exchange.getResponseBody().flush(); + streamOpened.countDown(); + try { + this.releaseResponse.await(); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + exchange.close(); + }); + + StepVerifier.create(transport.connect(Function.identity())).then(() -> { + try { + assertThat(streamOpened.await(TIMEOUT.toMillis(), TimeUnit.MILLISECONDS)).isTrue(); + } + catch (InterruptedException e) { + throw new IllegalStateException(e); + } + transport.closeGracefully().block(TIMEOUT); + }).expectComplete().verify(TIMEOUT); + } + + private HttpClientSseClientTransport transport(HttpHandler sseHandler) throws IOException { + this.server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + this.server.setExecutor(this.executor); + this.server.createContext("/sse", sseHandler); + this.server.start(); + return HttpClientSseClientTransport.builder("http://127.0.0.1:" + this.server.getAddress().getPort()) + .jsonMapper(new GsonMcpJsonMapper()) + .build(); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportSendMessageTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportSendMessageTests.java new file mode 100644 index 000000000..6c97963b0 --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportSendMessageTests.java @@ -0,0 +1,121 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.IOException; +import java.io.OutputStream; +import java.net.InetSocketAddress; +import java.nio.charset.StandardCharsets; +import java.time.Duration; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; + +import com.sun.net.httpserver.HttpExchange; +import com.sun.net.httpserver.HttpHandler; +import com.sun.net.httpserver.HttpServer; +import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.McpTransportException; +import io.modelcontextprotocol.spec.json.gson.GsonMcpJsonMapper; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import reactor.test.StepVerifier; + +import static org.assertj.core.api.Assertions.assertThat; + +class HttpClientStreamableHttpTransportSendMessageTests { + + // Only bounds a regression: every test resolves without waiting on it. + private static final Duration TIMEOUT = Duration.ofSeconds(5); + + private static final McpSchema.JSONRPCRequest REQUEST = new McpSchema.JSONRPCRequest(McpSchema.JSONRPC_VERSION, + "ping", "1", null); + + private final ExecutorService executor = Executors.newCachedThreadPool(); + + private final CountDownLatch releaseResponse = new CountDownLatch(1); + + private HttpServer server; + + @AfterEach + void tearDown() { + this.releaseResponse.countDown(); + if (this.server != null) { + this.server.stop(0); + } + this.executor.shutdownNow(); + } + + @Test + void sendMessageFailsWhenJsonResponseIsMalformed() throws IOException { + HttpClientStreamableHttpTransport transport = transport(exchange -> { + byte[] body = "{broken".getBytes(StandardCharsets.UTF_8); + exchange.getResponseHeaders().add("Content-Type", "application/json"); + exchange.sendResponseHeaders(200, body.length); + try (OutputStream outputStream = exchange.getResponseBody()) { + outputStream.write(body); + } + }); + + StepVerifier.create(transport.sendMessage(REQUEST)) + .expectErrorSatisfies(e -> assertThat(e).isInstanceOf(McpTransportException.class) + .hasMessageContaining("Error deserializing JSON-RPC message")) + .verify(TIMEOUT); + } + + @Test + void sendMessageCompletesWhenClosedBeforeAnyEvent() throws IOException { + CountDownLatch streamOpened = new CountDownLatch(1); + HttpClientStreamableHttpTransport transport = transport(exchange -> { + exchange.getResponseHeaders().add("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, 0); + exchange.getResponseBody().flush(); + streamOpened.countDown(); + try { + this.releaseResponse.await(); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + exchange.close(); + }); + + StepVerifier.create(transport.sendMessage(REQUEST)).then(() -> { + try { + assertThat(streamOpened.await(TIMEOUT.toMillis(), TimeUnit.MILLISECONDS)).isTrue(); + } + catch (InterruptedException e) { + throw new IllegalStateException(e); + } + transport.closeGracefully().block(TIMEOUT); + }).expectComplete().verify(TIMEOUT); + } + + private HttpClientStreamableHttpTransport transport(HttpHandler postHandler) throws IOException { + this.server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + this.server.setExecutor(this.executor); + this.server.createContext("/mcp", exchange -> { + if ("POST".equals(exchange.getRequestMethod())) { + postHandler.handle(exchange); + } + else { + // No standalone SSE stream, which keeps the POST the only exchange. + methodNotAllowed(exchange); + } + }); + this.server.start(); + return HttpClientStreamableHttpTransport.builder("http://127.0.0.1:" + this.server.getAddress().getPort()) + .jsonMapper(new GsonMcpJsonMapper()) + .build(); + } + + private static void methodNotAllowed(HttpExchange exchange) throws IOException { + try (exchange) { + exchange.sendResponseHeaders(405, -1); + } + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LargeSseEventDecodingTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LargeSseEventDecodingTests.java new file mode 100644 index 000000000..fe6bba9dd --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LargeSseEventDecodingTests.java @@ -0,0 +1,219 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.Flow; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import reactor.adapter.JdkFlowAdapter; +import reactor.core.publisher.Flux; + +import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEvent; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * Reproducer for the client-side SSE reading bottleneck reported in + * #1042: a + * tool response arriving as a single multi-megabyte {@code data:} line, which is what + * compact JSON looks like on the wire, took ~5s to read where the same bytes read with + * {@link java.net.http.HttpResponse.BodyHandlers#ofString()} took ~0.4s. + * + *

+ * What cost the time was the length of the line rather than the number of bytes, because + * the buffered characters were gone over again every time a chunk arrived. + * {@link #shouldReadOneLargeEventWithinBudgetOfManySmallOnes()} therefore measures the + * same payload twice, once as one long line and once split over short ones, and asserts a + * ratio rather than a duration, so that it keeps its meaning on a machine of any speed. + * + *

+ * Reading a single-line event in 16KiB chunks, best of three runs on the same machine: + * + *

+ * payload   2.0.0 (fromLineSubscriber)   rescanning the line   scanning each line once
+ *  1MiB                        210ms                    12ms                      2ms
+ *  2MiB                        810ms                    32ms                      6ms
+ *  4MiB                       3314ms                   138ms                     10ms
+ *  8MiB                      13140ms                   561ms                     18ms
+ * 
+ * + *

+ * The middle column read each chunk incrementally, which is ~25x quicker than what 2.0.0 + * shipped, but {@link ResponseBodyHandlers.Utf8LineDecoder} still searched its buffered + * characters for a line terminator from the start of the buffer on every chunk, so eight + * times the payload cost ~45x the time. Resuming that search where the previous one ended + * gives the third column, which scales with the payload rather than with its square and + * brings the ratio this test measures from ~12 to ~1.6. + * + *

+ * See {@code HttpClientStreamableHttpTransportLargeResponseTests} in {@code mcp-test} for + * the same comparison end to end, over a real connection. + * + * @author Daniel Garnier-Moiroux + */ +class LargeSseEventDecodingTests { + + private static final Logger logger = LoggerFactory.getLogger(LargeSseEventDecodingTests.class); + + /** + * Roughly what {@link java.net.http.HttpClient} hands to a body subscriber at a time. + * The cost the report is about was paid per chunk, so the chunking is part of the + * reproducer. + */ + private static final int CHUNK_SIZE = 16 * 1024; + + private static final int MIB = 1024 * 1024; + + /** + * The payload size in the report: ~4MiB of compact JSON, and therefore ~4MiB with no + * line terminator in it. + */ + private static final int PAYLOAD_SIZE = 4 * MIB; + + /** + * How much of {@link #PAYLOAD_SIZE} each event carries when the same total is split + * over many events. + */ + private static final int SMALL_EVENT_SIZE = 64 * 1024; + + /** + * How much longer decoding the payload as one long line may take than decoding the + * same bytes as short ones. A reader that goes over each line once is indifferent to + * how long the lines are, which measures ~1.6 here; the bound leaves headroom over + * that, and is far below what rescanning the line measured (~12x) or what 2.0.0 + * measured (~220x). + */ + private static final double MAX_SINGLE_EVENT_PENALTY = 4.0; + + private static final int MAX_SIZE = 64 * MIB; + + @Test + @Timeout(60) + void shouldDecodeMultiMegabyteSingleLineEventIntact() { + String payload = payloadOfSize(PAYLOAD_SIZE); + + List events = decode(oneLargeEvent(payload)); + + assertThat(events).hasSize(1); + assertThat(events.get(0).event()).isEqualTo("message"); + assertThat(events.get(0).data()).isEqualTo(payload); + } + + @Test + @Timeout(300) + void shouldReadOneLargeEventWithinBudgetOfManySmallOnes() { + byte[] oneEvent = oneLargeEvent(payloadOfSize(PAYLOAD_SIZE)); + byte[] manyEvents = manySmallEvents(PAYLOAD_SIZE, SMALL_EVENT_SIZE); + int smallEventCount = PAYLOAD_SIZE / SMALL_EVENT_SIZE; + + // The reporter measured a JIT effect at this payload size: the first few large + // reads spike before the hot loop settles. Warm up, then interleave the two + // shapes and take the best of each, so the comparison reflects steady state. + for (int i = 0; i < 5; i++) { + decode(manyEvents); + } + long single = Long.MAX_VALUE; + long split = Long.MAX_VALUE; + for (int i = 0; i < 3; i++) { + single = Math.min(single, timeDecode(oneEvent, 1)); + split = Math.min(split, timeDecode(manyEvents, smallEventCount)); + } + + double penalty = (double) single / Math.max(split, 1); + logger.info("decoded {}KiB as one event in {}ms and as {} events in {}ms: ratio {}", PAYLOAD_SIZE / 1024, + single / 1_000_000, smallEventCount, split / 1_000_000, String.format("%.1f", penalty)); + logScaling(); + + assertThat(penalty) + .as("decoding %dKiB as a single SSE event took %.1fx as long as decoding the same number of bytes as " + + "%dKiB events, so the cost of an event grows with the length of its line", PAYLOAD_SIZE / 1024, + penalty, SMALL_EVENT_SIZE / 1024) + .isLessThan(MAX_SINGLE_EVENT_PENALTY); + } + + /** + * Logs how decoding one long line scales with its length, which is the shape the + * report is about: doubling the payload should cost about twice the time, not four + * times it. Not asserted, because the ratio above covers the same ground with a + * baseline measured on the same machine. + */ + private void logScaling() { + for (int payloadSize : new int[] { MIB, 2 * MIB, 4 * MIB, 8 * MIB }) { + byte[] body = oneLargeEvent(payloadOfSize(payloadSize)); + long best = Math.min(timeDecode(body, 1), timeDecode(body, 1)); + logger.info("decoded a single-line event of {}KiB in {}ms", payloadSize / 1024, best / 1_000_000); + } + } + + /** + * Decodes the body once and returns how long it took, in nanoseconds. + */ + private static long timeDecode(byte[] body, int expectedEvents) { + long start = System.nanoTime(); + List events = decode(body); + long elapsed = System.nanoTime() - start; + assertThat(events).hasSize(expectedEvents); + return elapsed; + } + + /** + * Runs the body through the transport's SSE reading path, as + * {@code HttpClientStreamableHttpTransport} does, chunked the way the HTTP client + * chunks a response body. + */ + private static List decode(byte[] body) { + Flow.Publisher> publisher = JdkFlowAdapter + .publisherToFlowPublisher(Flux.fromIterable(chunk(body))); + Flux lines = ResponseBodyHandlers.decodeLines(publisher, Integer.MAX_VALUE); + return ResponseBodyHandlers.decodeSseResponse(lines, MAX_SIZE).collectList().block(); + } + + private static List> chunk(byte[] body) { + List> chunks = new ArrayList<>(); + for (int offset = 0; offset < body.length; offset += CHUNK_SIZE) { + int length = Math.min(CHUNK_SIZE, body.length - offset); + chunks.add(List.of(ByteBuffer.wrap(body, offset, length).asReadOnlyBuffer())); + } + return chunks; + } + + /** + * An SSE {@code message} event carrying the whole payload on a single {@code data:} + * line. + */ + private static byte[] oneLargeEvent(String payload) { + return ("event: message\ndata: " + payload + "\n\n").getBytes(StandardCharsets.UTF_8); + } + + /** + * The same {@code total} number of payload bytes, spread over events of + * {@code eachSize} each. + */ + private static byte[] manySmallEvents(int total, int eachSize) { + StringBuilder body = new StringBuilder(total + 4096); + for (int i = 0; i < total / eachSize; i++) { + body.append("event: message\ndata: ").append(payloadOfSize(eachSize)).append("\n\n"); + } + return body.toString().getBytes(StandardCharsets.UTF_8); + } + + /** + * A single-line JSON-RPC response of exactly {@code size} characters, none of them a + * line terminator. + */ + private static String payloadOfSize(int size) { + String prefix = "{\"jsonrpc\":\"2.0\",\"id\":\"test-id\",\"result\":{\"content\":\""; + String suffix = "\"}}"; + return prefix + "a".repeat(size - prefix.length() - suffix.length()) + suffix; + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LoopbackMcpHttpServer.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LoopbackMcpHttpServer.java new file mode 100644 index 000000000..01c7bb26d --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LoopbackMcpHttpServer.java @@ -0,0 +1,136 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.IOException; +import java.io.OutputStream; +import java.net.InetSocketAddress; +import java.net.URI; +import java.nio.charset.StandardCharsets; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.ThreadFactory; + +import com.sun.net.httpserver.Headers; +import com.sun.net.httpserver.HttpExchange; +import com.sun.net.httpserver.HttpHandler; +import com.sun.net.httpserver.HttpServer; + +final class LoopbackMcpHttpServer implements AutoCloseable { + + private static final byte[] EMPTY_BODY = new byte[0]; + + private static final byte[] STREAMABLE_PRIMER = """ + event: message + data: + + """.getBytes(StandardCharsets.UTF_8); + + private static final byte[] SSE_ENDPOINT_EVENT = """ + event: endpoint + data: /message + + """.getBytes(StandardCharsets.UTF_8); + + private final HttpServer server; + + private final ExecutorService executor; + + private LoopbackMcpHttpServer(HttpServer server, ExecutorService executor) { + this.server = server; + this.executor = executor; + } + + static LoopbackMcpHttpServer start() throws IOException { + HttpServer server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + ExecutorService executor = Executors.newCachedThreadPool(new LoopbackThreadFactory()); + server.setExecutor(executor); + server.createContext("/mcp", new StreamableHandler()); + server.createContext("/sse", new SseHandler()); + server.createContext("/message", new MessageHandler()); + server.start(); + return new LoopbackMcpHttpServer(server, executor); + } + + URI baseUri() { + return URI.create("http://127.0.0.1:" + this.server.getAddress().getPort()); + } + + @Override + public void close() { + this.server.stop(0); + this.executor.shutdownNow(); + } + + private static final class StreamableHandler implements HttpHandler { + + @Override + public void handle(HttpExchange exchange) throws IOException { + try (exchange) { + String method = exchange.getRequestMethod(); + if ("GET".equals(method)) { + Headers headers = exchange.getResponseHeaders(); + headers.add("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, STREAMABLE_PRIMER.length); + try (OutputStream outputStream = exchange.getResponseBody()) { + outputStream.write(STREAMABLE_PRIMER); + } + return; + } + if ("POST".equals(method)) { + exchange.getResponseHeaders().add("mcp-session-id", "loopback-session"); + exchange.sendResponseHeaders(202, -1); + return; + } + if ("DELETE".equals(method)) { + exchange.sendResponseHeaders(204, -1); + return; + } + exchange.sendResponseHeaders(405, EMPTY_BODY.length); + } + } + + } + + private static final class SseHandler implements HttpHandler { + + @Override + public void handle(HttpExchange exchange) throws IOException { + try (exchange) { + Headers headers = exchange.getResponseHeaders(); + headers.add("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, SSE_ENDPOINT_EVENT.length); + try (OutputStream outputStream = exchange.getResponseBody()) { + outputStream.write(SSE_ENDPOINT_EVENT); + } + } + } + + } + + private static final class MessageHandler implements HttpHandler { + + @Override + public void handle(HttpExchange exchange) throws IOException { + try (exchange) { + exchange.sendResponseHeaders(202, -1); + } + } + + } + + private static final class LoopbackThreadFactory implements ThreadFactory { + + @Override + public Thread newThread(Runnable runnable) { + Thread thread = new Thread(runnable); + thread.setDaemon(true); + thread.setName("loopback-mcp-http-server-" + thread.getId()); + return thread; + } + + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlersSendAsyncTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlersSendAsyncTests.java new file mode 100644 index 000000000..0b6289136 --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlersSendAsyncTests.java @@ -0,0 +1,78 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.IOException; +import java.net.InetSocketAddress; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; + +import com.sun.net.httpserver.HttpServer; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import reactor.core.Disposable; +import reactor.core.publisher.Hooks; + +import static org.assertj.core.api.Assertions.assertThat; + +class ResponseBodyHandlersSendAsyncTests { + + private final ExecutorService executor = Executors.newCachedThreadPool(); + + private final CountDownLatch releaseResponse = new CountDownLatch(1); + + private final List dropped = new CopyOnWriteArrayList<>(); + + private HttpServer server; + + @AfterEach + void tearDown() { + Hooks.resetOnErrorDropped(); + this.releaseResponse.countDown(); + if (this.server != null) { + this.server.stop(0); + } + this.executor.shutdownNow(); + } + + @Test + void cancellingBeforeTheResponseArrivesDropsNoError() throws Exception { + CountDownLatch requestReceived = new CountDownLatch(1); + this.server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + this.server.setExecutor(this.executor); + this.server.createContext("/", exchange -> { + // Never responds, so that the exchange is cancelled while awaiting headers. + requestReceived.countDown(); + try { + this.releaseResponse.await(); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + exchange.close(); + }); + this.server.start(); + Hooks.onErrorDropped(this.dropped::add); + + HttpRequest request = HttpRequest + .newBuilder(URI.create("http://127.0.0.1:" + this.server.getAddress().getPort() + "/")) + .build(); + Disposable exchange = ResponseBodyHandlers.sendAsync(HttpClient.newHttpClient(), request).subscribe(); + assertThat(requestReceived.await(5, TimeUnit.SECONDS)).isTrue(); + + // The HttpClient fails the aborted exchange within cancel() itself. + exchange.dispose(); + + assertThat(this.dropped).isEmpty(); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java new file mode 100644 index 000000000..459e42af0 --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java @@ -0,0 +1,185 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.util.Optional; + +import org.junit.jupiter.api.Test; + +import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEvent; +import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEventParser; + +import static org.assertj.core.api.Assertions.assertThat; + +class SseEventParserTests { + + @Test + void simpleDataEvent() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("data: hello")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().data()).isEqualTo("hello"); + assertThat(event.get().id()).isNull(); + assertThat(event.get().event()).isEqualTo("message"); + } + + @Test + void multiLineDataAccumulatesWithNewlineSeparatorAndTrims() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("data: first")).isEmpty(); + assertThat(p.feed("data: second")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().data()).isEqualTo("first\nsecond"); + } + + @Test + void idAndEventFieldsCaptured() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("id: 42")).isEmpty(); + assertThat(p.feed("event: message")).isEmpty(); + assertThat(p.feed("data: payload")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().id()).isEqualTo("42"); + assertThat(event.get().event()).isEqualTo("message"); + assertThat(event.get().data()).isEqualTo("payload"); + } + + @Test + void idPersistsAcrossEventsButEventTypeDoesNot() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("id: 1"); + p.feed("event: endpoint"); + p.feed("data: one"); + SseEvent first = p.feed("").orElseThrow(); + assertThat(first.id()).isEqualTo("1"); + assertThat(first.event()).isEqualTo("endpoint"); + + // An event that does not name its type is a message event, not another + // endpoint event + p.feed("data: two"); + SseEvent second = p.feed("").orElseThrow(); + assertThat(second.id()).isEqualTo("1"); + assertThat(second.event()).isEqualTo("message"); + assertThat(second.data()).isEqualTo("two"); + } + + @Test + void blankLineWithNoDataStillResetsEventType() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("event: endpoint"); + assertThat(p.feed("")).isEmpty(); + p.feed("data: payload"); + SseEvent event = p.feed("").orElseThrow(); + assertThat(event.event()).isEqualTo("message"); + } + + @Test + void emptyIdClearsLastEventId() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("id: 1"); + p.feed("data: one"); + assertThat(p.feed("").orElseThrow().id()).isEqualTo("1"); + + p.feed("id:"); + p.feed("data: two"); + assertThat(p.feed("").orElseThrow().id()).isNull(); + } + + @Test + void idContainingNullIsIgnored() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("id: 1"); + p.feed("id: 2\u0000" + "3"); + p.feed("data: payload"); + assertThat(p.feed("").orElseThrow().id()).isEqualTo("1"); + } + + @Test + void commentLineIgnored() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed(": this is a comment")).isEmpty(); + assertThat(p.feed("data: hello")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().data()).isEqualTo("hello"); + } + + @Test + void trailingIncompleteEventEmittedOnFlush() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("data: incomplete"); + Optional flushed = p.flush(); + assertThat(flushed).isPresent(); + assertThat(flushed.get().data()).isEqualTo("incomplete"); + } + + @Test + void flushWithNothingPendingIsEmpty() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.flush()).isEmpty(); + } + + @Test + void blankLineWithNoPendingDataIsEmpty() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("")).isEmpty(); + } + + @Test + void unknownFieldsAreIgnored() { + // The SSE spec mandates that unknown fields are ignored, so neither a standard + // field the parser does not act on nor a malformed line may fail the stream. + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("retry: 3000")).isEmpty(); + assertThat(p.feed("bogus line")).isEmpty(); + assertThat(p.feed("data: hello")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().data()).isEqualTo("hello"); + } + + @Test + void dataFieldWithEmptyValueStillDispatchesAnEvent() { + // Per the SSE spec a data field appends its value plus a separator, so a lone + // `data:` line leaves the buffer non-empty and the event is dispatched carrying + // empty data. Servers send exactly this to prime a stream, and dropping it leaves + // the request the stream answers hanging. + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("data:")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().data()).isEmpty(); + } + + @Test + void dataFieldWithOnlyASpaceIsEquivalentToNoValue() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("data: ")).isEmpty(); + Optional event = p.feed(""); + assertThat(event).isPresent(); + assertThat(event.get().data()).isEmpty(); + } + + @Test + void blankLineWithNoDataFieldDispatchesNothing() { + // `event:` alone leaves the data buffer empty, which per the spec is not an event + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("event: message")).isEmpty(); + assertThat(p.feed("")).isEmpty(); + } + + @Test + void valuelessDataFieldIsDispatchedOnFlush() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + assertThat(p.feed("data:")).isEmpty(); + Optional flushed = p.flush(); + assertThat(flushed).isPresent(); + assertThat(flushed.get().data()).isEmpty(); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderBoundTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderBoundTests.java new file mode 100644 index 000000000..63dd096c4 --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderBoundTests.java @@ -0,0 +1,131 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.util.List; + +import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.Utf8LineDecoder; +import io.modelcontextprotocol.spec.McpTransportException; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatCode; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +/** + * Covers the bound {@link Utf8LineDecoder} places on a single line. The decoder buffers + * characters until a line terminator arrives, so a peer that never sends one would + * otherwise force the transport to buffer its line in memory without limit. + * + * @author Daniel Garnier-Moiroux + */ +class Utf8LineDecoderBoundTests { + + private static final int MAX_SIZE = 16; + + private static List chunk(String text) { + return List.of(ByteBuffer.wrap(text.getBytes(StandardCharsets.UTF_8))); + } + + private static Utf8LineDecoder decoder() { + return new Utf8LineDecoder(MAX_SIZE); + } + + @Test + void acceptsUnterminatedLineOfExactlyMaxSize() { + Utf8LineDecoder dec = decoder(); + + assertThat(dec.decode(chunk("a".repeat(MAX_SIZE)))).isEmpty(); + assertThat(dec.flush()).containsExactly("a".repeat(MAX_SIZE)); + } + + @Test + void rejectsUnterminatedLineOneCharOverMaxSize() { + Utf8LineDecoder dec = decoder(); + + assertThatThrownBy(() -> dec.decode(chunk("a".repeat(MAX_SIZE + 1)))).isInstanceOf(McpTransportException.class) + .hasMessageContaining("Inbound line exceeds the maximum allowed size of " + MAX_SIZE + " bytes"); + } + + @Test + void accumulatesAcrossChunks() { + Utf8LineDecoder dec = decoder(); + + assertThat(dec.decode(chunk("a".repeat(10)))).isEmpty(); + assertThat(dec.decode(chunk("a".repeat(6)))).isEmpty(); + + assertThatThrownBy(() -> dec.decode(chunk("a"))).isInstanceOf(McpTransportException.class); + } + + @Test + void acceptsUnboundedTotalOfTerminatedLines() { + Utf8LineDecoder dec = decoder(); + + // Far more than MAX_SIZE in total, but no single line comes close to it. + assertThatCode(() -> { + for (int i = 0; i < 100; i++) { + dec.decode(chunk("a".repeat(MAX_SIZE / 2) + "\n")); + } + }).doesNotThrowAnyException(); + } + + @Test + void lineFeedRefillsTheBudget() { + Utf8LineDecoder dec = decoder(); + + assertThat(dec.decode(chunk("a".repeat(MAX_SIZE) + "\n"))).containsExactly("a".repeat(MAX_SIZE)); + // A fresh line, so the previous characters must not count towards it. + assertThat(dec.decode(chunk("a".repeat(MAX_SIZE)))).isEmpty(); + } + + @Test + void carriageReturnRefillsTheBudget() { + Utf8LineDecoder dec = decoder(); + + // The decoder terminates a line on a lone CR, so a CR empties its buffer and has + // to refill the budget too, or a peer framing short lines with CR alone would be + // rejected for exceeding a bound it never reached. + assertThatCode(() -> { + for (int i = 0; i < 100; i++) { + dec.decode(chunk("a".repeat(MAX_SIZE / 2) + "\r")); + } + }).doesNotThrowAnyException(); + } + + @Test + void crLfSplitAcrossChunksRefillsTheBudgetOnce() { + Utf8LineDecoder dec = decoder(); + + assertThat(dec.decode(chunk("a".repeat(12) + "\r"))).containsExactly("a".repeat(12)); + // The LF completes the terminator rather than ending a line of its own, so what + // follows it gets the whole budget. + assertThat(dec.decode(chunk("\n" + "a".repeat(MAX_SIZE)))).isEmpty(); + } + + @Test + void countsCharactersRatherThanBytes() { + Utf8LineDecoder dec = decoder(); + + // Each 'é' is two bytes but one character. Measuring characters is deliberately + // the more permissive of the two, so that a line is only rejected once it has + // genuinely exceeded the bound in bytes. + assertThat(dec.decode(chunk("é".repeat(MAX_SIZE)))).isEmpty(); + assertThat(dec.flush()).containsExactly("é".repeat(MAX_SIZE)); + } + + @Test + void reportsTheLinesDecodedBeforeTheBoundWasReached() { + Utf8LineDecoder dec = decoder(); + + // The offending run arrives in the same chunk as two good lines. Those are lost + // with the chunk, which is why the bound has to be generous enough that only a + // peer misbehaving can reach it. + assertThatThrownBy(() -> dec.decode(chunk("one\ntwo\n" + "a".repeat(MAX_SIZE + 1)))) + .isInstanceOf(McpTransportException.class); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderTests.java new file mode 100644 index 000000000..e97303f3d --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderTests.java @@ -0,0 +1,256 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.ByteArrayOutputStream; +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.util.List; + +import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.Utf8LineDecoder; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; + +class Utf8LineDecoderTests { + + /** + * A bound high enough that no test here can reach it: these tests cover decoding, and + * the bound has its own coverage in {@link Utf8LineDecoderBoundTests}. + */ + private static final int UNBOUNDED = Integer.MAX_VALUE; + + /** + * 0xFF cannot appear anywhere in well-formed UTF-8. One of these is what a peer + * mixing encodings, or a proxy corrupting a byte, puts on the wire. + */ + private static final byte[] INVALID_BYTE = { (byte) 0xFF }; + + /** + * The lead byte of the two-byte sequence for {@code 'é'} (U+00E9, 0xC3 0xA9). + */ + private static final byte[] TRUNCATED_LEAD_BYTE = { (byte) 0xC3 }; + + private static List chunk(String... parts) { + return List.of(toByteBuffers(parts)); + } + + /** + * A chunk whose bytes are passed through verbatim, so that bytes no encoder would + * produce reach the decoder as-is. + */ + private static List rawChunk(byte[]... parts) { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + for (byte[] part : parts) { + out.writeBytes(part); + } + return List.of(ByteBuffer.wrap(out.toByteArray())); + } + + private static byte[] utf8(String text) { + return text.getBytes(StandardCharsets.UTF_8); + } + + private static ByteBuffer[] toByteBuffers(String... parts) { + ByteBuffer[] bbs = new ByteBuffer[parts.length]; + for (int i = 0; i < parts.length; i++) { + bbs[i] = ByteBuffer.wrap(parts[i].getBytes(StandardCharsets.UTF_8)); + } + return bbs; + } + + @Test + void singleLineLf() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("hello\n"))).containsExactly("hello"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void singleLineCrLf() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("hello\r\n"))).containsExactly("hello"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void multipleLinesInOneChunk() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("one\ntwo\nthree\n"))).containsExactly("one", "two", "three"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void lineSplitAcrossChunks() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("hel"))).isEmpty(); + assertThat(dec.decode(chunk("lo\nworld"))).containsExactly("hello"); + assertThat(dec.flush()).containsExactly("world"); + } + + @Test + void lineSplitAcrossByteBuffersInSameChunk() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + // two byte-buffers, newline between them -- should still form one clean split + List chunk = List.of(ByteBuffer.wrap("part-one\npart-".getBytes(StandardCharsets.UTF_8)), + ByteBuffer.wrap("two\n".getBytes(StandardCharsets.UTF_8))); + assertThat(new Utf8LineDecoder(UNBOUNDED).decode(chunk)).containsExactly("part-one", "part-two"); + } + + @Test + void multiByteUtf8SplitAcrossChunks() { + // "€" is U+20AC → 0xE2 0x82 0xAC in UTF-8. Split between the first and second + // byte. + byte[] euro = "€".getBytes(StandardCharsets.UTF_8); + assertThat(euro).hasSize(3); + + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(List.of(ByteBuffer.wrap(new byte[] { euro[0] })))).isEmpty(); + assertThat(dec.decode(List.of(ByteBuffer.wrap(new byte[] { euro[1], euro[2], '\n' })))).containsExactly("€"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void consecutiveBlankLines() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("\n\n\n"))).containsExactly("", "", ""); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void trailingPartialLineEmittedOnFlush() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("incomplete"))).isEmpty(); + assertThat(dec.flush()).containsExactly("incomplete"); + } + + @Test + void trailingCrTerminatesTheLine() { + // A body whose last byte is a CR ends on a terminator, not part-way through a + // line, so there is nothing left to flush. + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("complete\r"))).containsExactly("complete"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void emptyInput() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(List.of())).isEmpty(); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void resumesSearchAfterTerminatorWhenLineWasSplitAcrossManyChunks() { + // The decoder remembers how far it has searched for a terminator, so the chunk + // that finally terminates a long line must not leave that mark behind and hide + // the lines that follow it. + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + for (int i = 0; i < 10; i++) { + assertThat(dec.decode(chunk("aaaa"))).isEmpty(); + } + assertThat(dec.decode(chunk("\nsecond\nthird\n"))).containsExactly("a".repeat(40), "second", "third"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void resumesSearchAcrossChunkWhenTerminatorFollowsUnterminatedPrefix() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("unterminated"))).isEmpty(); + assertThat(dec.decode(chunk("-still-going"))).isEmpty(); + assertThat(dec.decode(chunk("\n"))).containsExactly("unterminated-still-going"); + } + + @Test + void linesLongerThanInternalCharBuffer() { + // 4096 is the internal CharBuffer size; send a single line ~10k chars to force + // multiple overflow cycles. + StringBuilder big = new StringBuilder(); + for (int i = 0; i < 10_000; i++) { + big.append('a'); + } + big.append('\n'); + + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + List lines = dec.decode(chunk(big.toString())); + assertThat(lines).hasSize(1); + assertThat(lines.get(0)).hasSize(10_000); + } + + @Test + void loneCrTerminatesLine() { + // SSE takes its line endings from HTML, which terminates on CRLF, CR and LF + // alike, and HttpResponse.BodySubscribers#fromLineSubscriber -- the path this + // decoder replaces -- splits on all three. Splitting on LF alone leaves a + // CR-framed stream as one unterminated run: downstream a single unparseable line + // whose SSE fields are silently dropped, leaving the request the stream answers + // hanging, or past the decoder's bound a failed one. + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("one\rtwo\rthree\r"))).containsExactly("one", "two", "three"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void blankLinesFramedWithCr() { + // A CR ending a chunk terminates its line, so the CR opening the next one ends an + // empty line rather than completing a CRLF. Values match what + // HttpResponse.BodySubscribers#fromLineSubscriber produces for the same bytes. + assertThat(new Utf8LineDecoder(UNBOUNDED).decode(chunk("\r\r"))).containsExactly("", ""); + + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("one\r"))).containsExactly("one"); + assertThat(dec.decode(chunk("\r"))).containsExactly(""); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void sseFramedWithCrOnlyIsSplitIntoFieldLines() { + // The same stream as the SSE parser downstream has to receive it: one line per + // field, and the empty line that ends the event. + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("event: message\rdata: {\"a\":1}\r\r"))).containsExactly("event: message", + "data: {\"a\":1}", ""); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void crLfSplitAcrossChunks() { + // Splitting on a lone CR means emitting the line as soon as the CR arrives, so a + // LF opening the next chunk is the tail of a CRLF rather than an empty line of + // its own. The terminator also sits exactly at the point the previous search for + // one stopped. + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(chunk("hello\r"))).containsExactly("hello"); + assertThat(dec.decode(chunk("\nworld\r\n"))).containsExactly("world"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void malformedByteIsReplacedAndOtherLinesArePreserved() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(rawChunk(utf8("one\ncaf"), INVALID_BYTE, utf8("e\nthree\n")))).containsExactly("one", + "caf\uFFFDe", "three"); + assertThat(dec.flush()).isEmpty(); + } + + @Test + void trailingTruncatedCharacterIsReplacedOnFlush() { + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(rawChunk(utf8("one\ncaf"), TRUNCATED_LEAD_BYTE))).containsExactly("one"); + assertThat(dec.flush()).containsExactly("caf\uFFFD"); + } + + @Test + void incompleteMultiByteSequenceFollowedByValidDataIsReplaced() { + // character is cut short by a chunk boundary + // "€" is U+20AC → 0xE2 0x82 0xAC in UTF-8; only the first two bytes arrive. + byte[] euro = "€".getBytes(StandardCharsets.UTF_8); + + Utf8LineDecoder dec = new Utf8LineDecoder(UNBOUNDED); + assertThat(dec.decode(rawChunk(new byte[] { euro[0], euro[1] }))).isEmpty(); + assertThat(dec.decode(chunk("x\n"))).containsExactly("\uFFFDx"); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/spec/DefaultMcpTransportSessionTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/spec/DefaultMcpTransportSessionTests.java new file mode 100644 index 000000000..a8ffac41e --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/spec/DefaultMcpTransportSessionTests.java @@ -0,0 +1,26 @@ +package io.modelcontextprotocol.spec; + +import java.util.concurrent.atomic.AtomicBoolean; + +import org.junit.jupiter.api.Test; +import reactor.core.Disposable; +import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; + +import static org.assertj.core.api.Assertions.assertThat; + +class DefaultMcpTransportSessionTests { + + @Test + void closeGracefullyDisposesOpenConnectionsEvenWhenOnCloseFails() { + var disposed = new AtomicBoolean(); + Disposable disposable = () -> disposed.set(true); + var session = new DefaultMcpTransportSession(id -> Mono.error(new RuntimeException("boom"))); + session.addConnection(disposable); + + StepVerifier.create(session.closeGracefully()).expectErrorMessage("boom").verify(); + + assertThat(disposed.get()).isTrue(); + } + +} diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientBoundedReadTestSupport.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientBoundedReadTestSupport.java index a22dc1301..31b63936c 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientBoundedReadTestSupport.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientBoundedReadTestSupport.java @@ -8,14 +8,18 @@ import java.io.OutputStream; import java.net.InetSocketAddress; import java.nio.charset.StandardCharsets; +import java.time.Duration; +import java.util.Arrays; +import java.util.concurrent.CompletableFuture; import java.util.concurrent.Executors; -import com.sun.net.httpserver.HttpExchange; import com.sun.net.httpserver.HttpServer; import io.modelcontextprotocol.server.transport.TomcatTestUtil; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; +import static org.assertj.core.api.Assertions.assertThat; + /** * Shared fixture for the transport bounded-read tests: a bare {@link HttpServer} whose * response body is written by a per-test {@link Responder}. @@ -62,50 +66,88 @@ void stopServer() { } /** - * Registers a handler that responds with the given content type and body. + * Registers a handler that answers {@code method} requests to {@code path} with the + * given content type and a body written by {@code responder}. Any other method gets a + * 405, as from a server offering nothing else there, so that requests the test does + * not target, such as the GET stream the Streamable HTTP transport opens once + * initialized, neither reach the responder nor stand in for the targeted request. + * @return completes once the targeted response has been handled: with the + * {@link IOException} that cut the body short if the client hung up first, or with + * {@code null} if the body was written in full + */ + protected CompletableFuture respondWith(String method, String path, String contentType, + Responder responder) { + return respondWith(method, path, 200, contentType, responder); + } + + /** + * Like {@link #respondWith(String, String, String, Responder)}, but answering with + * {@code status} rather than 200. */ - protected void respondWith(String path, String contentType, Responder responder) { + protected CompletableFuture respondWith(String method, String path, int status, String contentType, + Responder responder) { + CompletableFuture response = new CompletableFuture<>(); this.server.createContext(path, exchange -> { - exchange.getResponseHeaders().set("Content-Type", contentType); - exchange.sendResponseHeaders(200, 0); - try (OutputStream body = exchange.getResponseBody()) { - responder.respond(body); - } - catch (IOException ignored) { - // The client aborts the response once the limit is exceeded, which closes - // the connection and makes further writes fail. That is the behaviour - // under test. + try { + if (!method.equals(exchange.getRequestMethod())) { + exchange.sendResponseHeaders(405, -1); + return; + } + exchange.getResponseHeaders().set("Content-Type", contentType); + exchange.sendResponseHeaders(status, 0); + try (OutputStream body = exchange.getResponseBody()) { + responder.respond(body); + response.complete(null); + } + catch (IOException ex) { + response.complete(ex); + } } finally { exchange.close(); } }); + return response; } /** - * A responder that writes {@code chunks} blocks of {@code 'a'} with no line - * terminator anywhere, so nothing downstream can ever flush a line. + * Asserts that the client hung up on {@code response}, which is how exceeding the + * bound must end: with the endless responders below, a client that read the body in + * full would never let it complete, and one that merely stopped reading would leave + * the server blocked writing into a stalled connection. */ - protected static Responder unterminatedLine(int chunks) { - return body -> { - byte[] chunk = new byte[MAX_SIZE]; - java.util.Arrays.fill(chunk, (byte) 'a'); - for (int i = 0; i < chunks; i++) { - body.write(chunk); - body.flush(); - } - }; + protected static void assertHungUp(CompletableFuture response) { + assertThat(response).succeedsWithin(Duration.ofSeconds(5)).isNotNull(); + } + + /** + * A responder that streams {@code 'a'} with no line terminator anywhere, so nothing + * downstream can ever flush a line. + */ + protected static Responder unterminatedLine() { + byte[] block = new byte[MAX_SIZE]; + Arrays.fill(block, (byte) 'a'); + return endlessly(block); } /** - * A responder that writes enough short, properly terminated lines to exceed the limit - * in aggregate. + * A responder that streams short, properly terminated lines, each small but exceeding + * the limit in aggregate. */ protected static Responder manyShortLines(String prefix) { + return endlessly((prefix + "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa\n").getBytes(StandardCharsets.UTF_8)); + } + + /** + * A responder that repeats {@code block} until the client hangs up, as a peer + * streaming without end would. There is no amount to tune: however much the socket + * buffers absorb, and whether the client closes with a FIN or a RST, the writes only + * stop once the connection is gone. + */ + private static Responder endlessly(byte[] block) { return body -> { - byte[] line = (prefix + "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa\n").getBytes(StandardCharsets.UTF_8); - for (int i = 0; i < (MAX_SIZE / line.length) + 64; i++) { - body.write(line); + while (true) { + body.write(block); body.flush(); } }; diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportBoundedReadTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportBoundedReadTests.java index df613265f..765ac5c1f 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportBoundedReadTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportBoundedReadTests.java @@ -48,38 +48,52 @@ void releaseStream() { void shouldRejectSingleLineExceedingMaxSize() { // A line that never terminates, so the line buffer underneath the SSE parser // would grow without limit before any event could be flushed. - respondWith(endpoint(), "text/event-stream", unterminatedLine(8)); + var response = respondWith("GET", endpoint(), "text/event-stream", unterminatedLine()); StepVerifier.create(connect()) .verifyErrorMatches(t -> messageContains(t, "Inbound line exceeds the maximum allowed size")); + assertHungUp(response); } @Test void shouldRejectEventExceedingMaxSize() { // Many short, terminated "data:" lines with no blank line to end the event. Each // line is small, but the accumulated event data would grow without limit. - respondWith(endpoint(), "text/event-stream", manyShortLines("data:")); + var response = respondWith("GET", endpoint(), "text/event-stream", manyShortLines("data:")); StepVerifier.create(connect()) .verifyErrorMatches(t -> messageContains(t, "Inbound SSE event exceeds the maximum allowed size")); + assertHungUp(response); } @Test - void shouldRejectPostResponseExceedingMaxSize() throws Exception { - // The response to a posted message is read into a string in full, so an oversized - // one must abort rather than accumulate. - respondWith(endpoint(), "text/event-stream", body -> { + void shouldRejectPostResponseExceedingMaxSize() { + // The response to a posted message is discarded on success, but a peer must still + // not be able to make the transport read an unbounded one. + respondWith("GET", endpoint(), "text/event-stream", body -> { body.write(("event:endpoint\ndata:" + MESSAGE_ENDPOINT + "\n\n").getBytes(StandardCharsets.UTF_8)); body.flush(); awaitTeardown(); }); - respondWith(MESSAGE_ENDPOINT, "application/json", unterminatedLine(64)); + var response = respondWith("POST", MESSAGE_ENDPOINT, "application/json", unterminatedLine()); HttpClientSseClientTransport transport = transport(); transport.connect(Function.identity()).block(Duration.ofSeconds(5)); StepVerifier.create(sendMessage(transport)) .verifyErrorMatches(t -> messageContains(t, "Inbound response body exceeds the maximum allowed size")); + assertHungUp(response); + } + + @Test + void shouldIncludeConnectErrorResponseBodyInError() { + // What the server says about a failure is the most useful part of it to report. + respondWith("GET", endpoint(), 500, "text/plain", + body -> body.write("upstream unavailable".getBytes(StandardCharsets.UTF_8))); + + StepVerifier.create(connect()) + .verifyErrorMatches(t -> messageContains(t, + "Failed to connect to SSE stream: 500, response body: upstream unavailable")); } private void awaitTeardown() { diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportBoundedReadTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportBoundedReadTests.java index 856dc88fb..59aaa8647 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportBoundedReadTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportBoundedReadTests.java @@ -34,48 +34,64 @@ protected String endpoint() { void shouldRejectSingleLineExceedingMaxSize() { // A line that never terminates, so the line buffer underneath the SSE parser // would grow without limit before any event could be flushed. - respondWith(endpoint(), "text/event-stream", unterminatedLine(8)); + var response = respondWith("POST", endpoint(), "text/event-stream", unterminatedLine()); StepVerifier.create(sendMessage()) .verifyErrorMatches(t -> messageContains(t, "Inbound line exceeds the maximum allowed size")); + assertHungUp(response); } @Test void shouldRejectEventExceedingMaxSize() { // Many short, terminated "data:" lines with no blank line to end the event. Each // line is small, but the accumulated event data would grow without limit. - respondWith(endpoint(), "text/event-stream", manyShortLines("data:")); + var response = respondWith("POST", endpoint(), "text/event-stream", manyShortLines("data:")); StepVerifier.create(sendMessage()) .verifyErrorMatches(t -> messageContains(t, "Inbound SSE event exceeds the maximum allowed size")); + assertHungUp(response); } @Test void shouldRejectJsonResponseExceedingMaxSize() { // A multi-line application/json response whose total size exceeds the limit. Each // line is small, but the aggregated body would grow without limit. - respondWith(endpoint(), "application/json", manyShortLines("")); + var response = respondWith("POST", endpoint(), "application/json", manyShortLines("")); StepVerifier.create(sendMessage()) .verifyErrorMatches(t -> messageContains(t, "Inbound response body exceeds the maximum allowed size")); + assertHungUp(response); } @Test - void shouldRejectDiscardedResponseExceedingMaxSize() { + void shouldRejectDiscardedResponseExceedingMaxSizeButReportProperError() { // A content type the transport neither parses as SSE nor as JSON, so the body is - // discarded. The line subscriber underneath still buffers each line, so an - // unterminated one must abort the response rather than accumulate. - respondWith(endpoint(), "text/plain", unterminatedLine(8)); + // discarded. Nothing accumulates, but a peer must still not be able to make the + // transport read an unbounded body only to throw it away. The error reported is + // the content type mismatch, so the bound only shows in the client hanging up. + var response = respondWith("POST", endpoint(), "text/plain", unterminatedLine()); StepVerifier.create(sendMessage()) - .verifyErrorMatches(t -> messageContains(t, "Inbound line exceeds the maximum allowed size")); + .verifyErrorMatches(t -> messageContains(t, "Unknown media type returned: text/plain")); + assertHungUp(response); + } + + @Test + void shouldIncludeErrorResponseBodyInError() { + // What the server says about a failure is the most useful part of it to report. + respondWith("POST", endpoint(), 404, "text/plain", + body -> body.write("no MCP server here".getBytes(StandardCharsets.UTF_8))); + + StepVerifier.create(sendMessage()) + .verifyErrorMatches( + t -> messageContains(t, "Server Not Found. Status code:404, response body: no MCP server here")); } @Test void shouldAcceptEventOfExactlyMaxSize() { // The bound is inclusive and the SSE framing around the payload is given its own // headroom, so a message of exactly maxResponseSize must still be delivered. - respondWith(endpoint(), "text/event-stream", body -> body + respondWith("POST", endpoint(), "text/event-stream", body -> body .write(("data:" + jsonRpcResponseOfExactly(MAX_SIZE) + "\n\n").getBytes(StandardCharsets.UTF_8))); StepVerifier.create(sendMessage()).verifyComplete(); @@ -84,7 +100,7 @@ void shouldAcceptEventOfExactlyMaxSize() { @Test void shouldAcceptJsonResponseOfExactlyMaxSize() { // Same inclusive bound on the aggregated body. - respondWith(endpoint(), "application/json", + respondWith("POST", endpoint(), "application/json", body -> body.write(jsonRpcResponseOfExactly(MAX_SIZE).getBytes(StandardCharsets.UTF_8))); StepVerifier.create(sendMessage()).verifyComplete(); diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyJsonResponseTest.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyJsonResponseTest.java deleted file mode 100644 index c2d19ef67..000000000 --- a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyJsonResponseTest.java +++ /dev/null @@ -1,94 +0,0 @@ -/* - * Copyright 2024-2025 the original author or authors. - */ - -package io.modelcontextprotocol.client.transport; - -import static org.mockito.ArgumentMatchers.any; -import static org.mockito.ArgumentMatchers.eq; -import static org.mockito.Mockito.atLeastOnce; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.verify; - -import java.io.IOException; -import java.net.InetSocketAddress; -import java.net.URI; -import java.net.URISyntaxException; - -import org.junit.jupiter.api.AfterAll; -import org.junit.jupiter.api.BeforeAll; -import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.Timeout; - -import com.sun.net.httpserver.HttpServer; - -import io.modelcontextprotocol.client.transport.customizer.McpSyncHttpClientRequestCustomizer; -import io.modelcontextprotocol.server.transport.TomcatTestUtil; -import io.modelcontextprotocol.spec.McpSchema; -import io.modelcontextprotocol.spec.ProtocolVersions; -import reactor.test.StepVerifier; - -/** - * Handles emplty application/json response with 200 OK status code. - * - * @author codezkk - */ -public class HttpClientStreamableHttpTransportEmptyJsonResponseTest { - - static int PORT = TomcatTestUtil.findAvailablePort(); - - static String host = "http://localhost:" + PORT; - - static HttpServer server; - - @BeforeAll - static void startContainer() throws IOException { - - server = HttpServer.create(new InetSocketAddress(PORT), 0); - - // Empty, 200 OK response for the /mcp endpoint - server.createContext("/mcp", exchange -> { - exchange.getResponseHeaders().set("Content-Type", "application/json"); - exchange.sendResponseHeaders(200, 0); - exchange.close(); - }); - - server.setExecutor(null); - server.start(); - } - - @AfterAll - static void stopContainer() { - server.stop(1); - } - - /** - * Regardless of the response (even if the response is null and the content-type is - * present), notify should handle it correctly. - */ - @Test - @Timeout(3) - void testNotificationInitialized() throws URISyntaxException { - - var uri = new URI(host + "/mcp"); - var mockRequestCustomizer = mock(McpSyncHttpClientRequestCustomizer.class); - var transport = HttpClientStreamableHttpTransport.builder(host) - .httpRequestCustomizer(mockRequestCustomizer) - .build(); - - var initializeRequest = McpSchema.InitializeRequest - .builder(ProtocolVersions.MCP_2025_03_26, McpSchema.ClientCapabilities.builder().roots(true).build(), - McpSchema.Implementation.builder("MCP Client", "0.3.1").build()) - .build(); - var testMessage = new McpSchema.JSONRPCRequest(McpSchema.METHOD_INITIALIZE, "test-id", initializeRequest); - - StepVerifier.create(transport.sendMessage(testMessage)).verifyComplete(); - - // Verify the customizer was called - verify(mockRequestCustomizer, atLeastOnce()).customize(any(), eq("POST"), eq(uri), eq( - "{\"jsonrpc\":\"2.0\",\"method\":\"initialize\",\"id\":\"test-id\",\"params\":{\"protocolVersion\":\"2025-03-26\",\"capabilities\":{\"roots\":{\"listChanged\":true}},\"clientInfo\":{\"name\":\"MCP Client\",\"version\":\"0.3.1\"}}}"), - any()); - - } - -} diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyResponseTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyResponseTests.java new file mode 100644 index 000000000..5faf3ae4e --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportEmptyResponseTests.java @@ -0,0 +1,147 @@ +/* + * Copyright 2024-2025 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.atLeastOnce; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; + +import java.io.IOException; +import java.io.OutputStream; +import java.net.InetSocketAddress; +import java.net.URI; +import java.net.URISyntaxException; +import java.nio.charset.StandardCharsets; +import java.util.Map; + +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; + +import com.sun.net.httpserver.HttpServer; + +import io.modelcontextprotocol.client.transport.customizer.McpSyncHttpClientRequestCustomizer; +import io.modelcontextprotocol.server.transport.TomcatTestUtil; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.test.StepVerifier; + +/** + * Handles 200 OK responses that carry no usable body, either as an empty application/json + * document or as a text/event-stream containing nothing but a stream primer. + * + * @author codezkk + */ +public class HttpClientStreamableHttpTransportEmptyResponseTests { + + static int PORT = TomcatTestUtil.findAvailablePort(); + + static String host = "http://localhost:" + PORT; + + static HttpServer server; + + /** + * An SSE event with an {@code event:} field but no data, which some servers send to + * open the response stream before any JSON-RPC payload is available. Note the + * valueless {@code data:} field: per the SSE spec this is identical to {@code data: } + * with a trailing space. + * @see SEP-1699 + */ + private static final byte[] SSE_PRIMER = """ + event: message + data: + + """.getBytes(StandardCharsets.UTF_8); + + @BeforeAll + static void startContainer() throws IOException { + + server = HttpServer.create(new InetSocketAddress(PORT), 0); + + // Empty, 200 OK response for the /mcp endpoint + server.createContext("/mcp", exchange -> { + exchange.getResponseHeaders().set("Content-Type", "application/json"); + exchange.sendResponseHeaders(200, 0); + exchange.close(); + }); + + // 200 OK text/event-stream carrying only a primer, for POSTs. The + // server-initiated GET stream is refused so that the transport falls back to + // request-response mode and the POST is the only thing under test. + server.createContext("/mcp-sse-primer", exchange -> { + try (exchange) { + if (!"POST".equals(exchange.getRequestMethod())) { + exchange.sendResponseHeaders(405, -1); + return; + } + exchange.getRequestBody().readAllBytes(); + exchange.getResponseHeaders().set("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, SSE_PRIMER.length); + try (OutputStream out = exchange.getResponseBody()) { + out.write(SSE_PRIMER); + } + } + }); + + server.setExecutor(null); + server.start(); + } + + @AfterAll + static void stopContainer() { + server.stop(1); + } + + /** + * Regardless of the response (even if the response is null and the content-type is + * present), notify should handle it correctly. + */ + @Test + @Timeout(3) + void testNotificationInitialized() throws URISyntaxException { + + var uri = new URI(host + "/mcp"); + var mockRequestCustomizer = mock(McpSyncHttpClientRequestCustomizer.class); + var transport = HttpClientStreamableHttpTransport.builder(host) + .httpRequestCustomizer(mockRequestCustomizer) + .build(); + + // Some servers answer a notification with an empty JSON body rather than 202. + var testMessage = new McpSchema.JSONRPCNotification(McpSchema.JSONRPC_VERSION, + McpSchema.METHOD_NOTIFICATION_INITIALIZED, null); + + StepVerifier.create(transport.sendMessage(testMessage)).verifyComplete(); + + // Verify the customizer was called + verify(mockRequestCustomizer, atLeastOnce()).customize(any(), eq("POST"), eq(uri), + eq("{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\"}"), any()); + + } + + /** + * A POST answered with {@code 200 text/event-stream} whose body holds only a stream + * primer must still complete, because the primer tells the client the stream is live + * and the message has been accepted. The primer's {@code data:} field carries no + * value, so this only holds as long as such a field still produces an event: a parser + * that drops it leaves no event to fire the transport's first-message callback, and + * {@code sendMessage} then never completes at all. + */ + @Test + @Timeout(5) + void testNotificationAnsweredWithSsePrimerOnly() { + + var transport = HttpClientStreamableHttpTransport.builder(host).endpoint("/mcp-sse-primer").build(); + + var testMessage = new McpSchema.JSONRPCNotification(McpSchema.JSONRPC_VERSION, "notifications/initialized", + Map.of()); + + StepVerifier.create(transport.sendMessage(testMessage)).verifyComplete(); + + } + +} diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportLargeResponseTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportLargeResponseTests.java new file mode 100644 index 000000000..51c62997b --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransportLargeResponseTests.java @@ -0,0 +1,295 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.io.IOException; +import java.io.OutputStream; +import java.net.InetSocketAddress; +import java.nio.charset.StandardCharsets; +import java.time.Duration; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.function.Consumer; + +import com.sun.net.httpserver.HttpServer; +import io.modelcontextprotocol.server.transport.TomcatTestUtil; +import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCMessage; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCNotification; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCRequest; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCResponse; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import reactor.core.publisher.Mono; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * End-to-end reproducer for + * #1042: a + * tool result of a few megabytes, delivered on the POST response's SSE stream as a single + * compact-JSON {@code data:} line, took the client ~5s to read where {@code curl} and + * {@link java.net.http.HttpResponse.BodyHandlers#ofString()} read the same bytes in + * ~0.4s. + * + *

+ * The reported setup is reproduced with a bare {@link HttpServer} in place of an MCP + * server, so that only the client's reading of the response is measured. What makes the + * payload expensive is that it arrives as one very long line, because compact JSON has no + * newline in it, so + * {@link #shouldReadOneLargeEventAboutAsFastAsTheSameBytesSplitOverManyEvents()} compares + * reading it against reading the same number of bytes split over many short events. That + * ratio is what the line length costs, with everything else (the wire, the JSON parsing, + * the machine) held constant. + * + *

+ * Measured here for a 4MiB payload, best of 12 reads each: ~26ms as one event against + * ~26ms split up, a ratio of ~1. Before the line decoder resumed its search for a line + * terminator where the previous search ended, instead of restarting it for every chunk + * that arrives, the same comparison measured ~180ms against ~30ms, and 2.0.0 measured + * ~100x. See {@code LargeSseEventDecodingTests} in {@code mcp-core} for the same + * comparison without a wire in between, and for how it scales with the payload. + * + *

+ * The per-read timings are logged. The issue also reported the first few reads of a large + * response taking an order of magnitude longer than the ones after them, which was that + * per-chunk cost paid while the hot loop was still being compiled: at this payload size + * the reads now go 404ms, 48ms, 39ms, then settle at ~26ms. The assertion is still made + * on the best read of many, so that it describes steady state rather than compilation. + * + * @author Daniel Garnier-Moiroux + */ +@Timeout(300) +class HttpClientStreamableHttpTransportLargeResponseTests { + + private static final Logger logger = LoggerFactory + .getLogger(HttpClientStreamableHttpTransportLargeResponseTests.class); + + private static final String ENDPOINT = "/mcp"; + + private static final String REQUEST_ID = "test-id"; + + /** + * The payload size in the report: ~4MiB of compact JSON, and therefore ~4MiB with no + * line terminator in it. + */ + private static final int PAYLOAD_SIZE = 4 * 1024 * 1024; + + /** + * How much of {@link #PAYLOAD_SIZE} each event carries when the same total is split + * over many events. + */ + private static final int SMALL_EVENT_SIZE = 64 * 1024; + + private static final int MEASURED_READS = 12; + + /** + * How much longer reading the payload as one event may take than reading the same + * bytes split over many events. A reader that goes over each line once is indifferent + * to how long the lines are, which measures ~1 here; the bound leaves headroom over + * that for a loaded machine, and is far below what rescanning the line measured + * (~6x). + */ + private static final double MAX_SINGLE_EVENT_PENALTY = 2.5; + + /** + * Writes the SSE body of the POST response, and may block until the client has + * reacted to what it has written so far. + */ + @FunctionalInterface + private interface SseResponder { + + void respond(OutputStream body) throws IOException, InterruptedException; + + } + + private HttpServer server; + + private String host; + + private volatile SseResponder responder; + + @BeforeEach + void startServer() throws IOException { + int port = TomcatTestUtil.findAvailablePort(); + this.host = "http://localhost:" + port; + this.server = HttpServer.create(new InetSocketAddress(port), 0); + this.server.setExecutor(Executors.newCachedThreadPool()); + this.server.createContext(ENDPOINT, exchange -> { + try (exchange) { + if (!"POST".equals(exchange.getRequestMethod())) { + // The transport opens a server-initiated stream after its first POST. + // 405 tells it there is none, which keeps this fixture to a single + // request-response exchange. + exchange.sendResponseHeaders(405, -1); + return; + } + exchange.getRequestBody().readAllBytes(); + exchange.getResponseHeaders().set("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, 0); + try (OutputStream body = exchange.getResponseBody()) { + this.responder.respond(body); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new IOException(e); + } + } + }); + this.server.start(); + } + + @AfterEach + void stopServer() { + if (this.server != null) { + this.server.stop(0); + } + } + + @Test + void shouldReceiveLargeSingleLineSseResponseIntact() throws Exception { + String payload = "a".repeat(PAYLOAD_SIZE); + this.responder = body -> writeEvent(body, jsonRpcResponse(payload)); + + List received = new CopyOnWriteArrayList<>(); + long elapsed = time(() -> readResponse(received::add)); + logger.info("received a {}KiB single-line SSE response in {}ms", PAYLOAD_SIZE / 1024, elapsed / 1_000_000); + + assertThat(received).hasSize(1); + JSONRPCResponse response = (JSONRPCResponse) received.get(0); + assertThat(response.id()).isEqualTo(REQUEST_ID); + assertThat(((Map) response.result()).get("content")).isEqualTo(payload); + } + + @Test + void shouldDeliverInterleavedNotificationBeforeTheLargeResultIsWritten() throws Exception { + // The response interleaves a progress notification before the result, on the same + // stream, which is why the issue rules out reading the whole body in one go: each + // event has to be delivered as its boundary arrives. This server refuses to write + // the result until the client has acknowledged the notification, so a reader that + // waits for the whole body deadlocks instead of quietly passing. + CountDownLatch notificationDelivered = new CountDownLatch(1); + AtomicBoolean deliveredBeforeResult = new AtomicBoolean(); + this.responder = body -> { + writeEvent(body, progressNotification("")); + deliveredBeforeResult.set(notificationDelivered.await(60, TimeUnit.SECONDS)); + writeEvent(body, jsonRpcResponse("a".repeat(PAYLOAD_SIZE))); + }; + + List received = new CopyOnWriteArrayList<>(); + readResponse(message -> { + received.add(message); + if (message instanceof JSONRPCNotification) { + notificationDelivered.countDown(); + } + }); + + assertThat(deliveredBeforeResult) + .as("the notification sent before the %dKiB result was not delivered until the whole response body had been read", + PAYLOAD_SIZE / 1024) + .isTrue(); + assertThat(received).hasSize(2); + assertThat(received.get(0)).isInstanceOf(JSONRPCNotification.class); + assertThat(received.get(1)).isInstanceOf(JSONRPCResponse.class); + } + + @Test + void shouldReadOneLargeEventWithinBudgetOfManySmallOnes() { + String payload = "a".repeat(PAYLOAD_SIZE); + String chunk = "a".repeat(SMALL_EVENT_SIZE); + SseResponder oneLargeEvent = body -> writeEvent(body, jsonRpcResponse(payload)); + SseResponder manySmallEvents = body -> { + for (int i = 0; i < PAYLOAD_SIZE / SMALL_EVENT_SIZE; i++) { + writeEvent(body, progressNotification(chunk)); + } + writeEvent(body, jsonRpcResponse("")); + }; + + // Interleaved, so that both shapes see the same machine and the same JIT state. + long oneEvent = Long.MAX_VALUE; + long manyEvents = Long.MAX_VALUE; + for (int i = 0; i < MEASURED_READS; i++) { + this.responder = oneLargeEvent; + long oneEventRead = time(() -> readResponse(message -> { + })); + this.responder = manySmallEvents; + long manyEventsRead = time(() -> readResponse(message -> { + })); + logger.info("read #{} of {}KiB: {}ms as one event, {}ms split over {}KiB events", i + 1, + PAYLOAD_SIZE / 1024, oneEventRead / 1_000_000, manyEventsRead / 1_000_000, SMALL_EVENT_SIZE / 1024); + oneEvent = Math.min(oneEvent, oneEventRead); + manyEvents = Math.min(manyEvents, manyEventsRead); + } + + double penalty = (double) oneEvent / Math.max(manyEvents, 1); + logger.info("best read: {}ms as one event, {}ms split up: ratio {}", oneEvent / 1_000_000, + manyEvents / 1_000_000, String.format("%.1f", penalty)); + + assertThat(penalty) + .as("reading %dKiB as a single SSE event took %.1fx as long as reading the same number of bytes split " + + "over %dKiB events, so the cost of an event grows with the length of its line", + PAYLOAD_SIZE / 1024, penalty, SMALL_EVENT_SIZE / 1024) + .isLessThan(MAX_SINGLE_EVENT_PENALTY); + } + + /** + * Sends one request and returns once the response has been delivered, handing every + * message received on the way to {@code onMessage}. + */ + private void readResponse(Consumer onMessage) { + HttpClientStreamableHttpTransport transport = HttpClientStreamableHttpTransport.builder(this.host) + .endpoint(ENDPOINT) + .build(); + CompletableFuture response = new CompletableFuture<>(); + JSONRPCRequest request = new JSONRPCRequest(McpSchema.JSONRPC_VERSION, "tools/call", REQUEST_ID, + Map.of("name", "large-response")); + try { + transport.connect(messages -> messages.doOnNext(message -> { + onMessage.accept(message); + if (message instanceof JSONRPCResponse) { + response.complete(message); + } + })).then(transport.sendMessage(request)).block(Duration.ofSeconds(120)); + response.get(120, TimeUnit.SECONDS); + } + catch (Exception e) { + throw new RuntimeException(e); + } + finally { + transport.closeGracefully().block(Duration.ofSeconds(10)); + } + } + + private static long time(Runnable read) { + long start = System.nanoTime(); + read.run(); + return System.nanoTime() - start; + } + + private static void writeEvent(OutputStream body, String data) throws IOException { + body.write(("event: message\ndata: " + data + "\n\n").getBytes(StandardCharsets.UTF_8)); + body.flush(); + } + + private static String jsonRpcResponse(String payload) { + return "{\"jsonrpc\":\"2.0\",\"id\":\"" + REQUEST_ID + "\",\"result\":{\"content\":\"" + payload + "\"}}"; + } + + private static String progressNotification(String payload) { + return "{\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\"params\":{\"progressToken\":\"" + + REQUEST_ID + "\",\"progress\":1,\"total\":2,\"message\":\"" + payload + "\"}}"; + } + +} From 6928451acb8566f8e8e2bbc5b78a5c0dde3c0b22 Mon Sep 17 00:00:00 2001 From: NewPeople-star <232190515+NewPeople-star@users.noreply.github.com> Date: Fri, 2 Oct 2026 22:00:16 +0800 Subject: [PATCH 3/3] Preserve whitespace in SSE field values --- .../transport/ResponseBodyHandlers.java | 19 ++++-- .../client/transport/SseEventParserTests.java | 65 +++++++++++++++++-- 2 files changed, 73 insertions(+), 11 deletions(-) diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java index 01e243da5..6e2ebf4c7 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java @@ -498,7 +498,7 @@ Optional feed(String line) { // valueless `data:` line still marks the event as carrying data and gets // dispatched with empty data. Servers send such an event to prime a // stream, and dropping it leaves the request it answers hanging. - String value = line.substring(5).trim(); + String value = fieldValue(line, 5); // Measured before appending, so that an event carrying exactly // maxSize of data is accepted: the trailing separator below is // stripped again before the event is emitted. @@ -509,7 +509,7 @@ Optional feed(String line) { data.append(value).append('\n'); } else if (line.startsWith("id:")) { - String value = line.substring(3).trim(); + String value = fieldValue(line, 3); // The spec ignores an id carrying a NULL, and an empty id resets the last // event ID, which leaves nothing to resume from. if (value.indexOf('\0') == -1) { @@ -517,7 +517,7 @@ else if (line.startsWith("id:")) { } } else if (line.startsWith("event:")) { - String value = line.substring(6).trim(); + String value = fieldValue(line, 6); event = value.isEmpty() ? null : value; } else if (line.startsWith(":")) { @@ -531,6 +531,15 @@ else if (line.startsWith(":")) { return Optional.empty(); } + private static String fieldValue(String line, int offset) { + // SSE removes only one optional U+0020 after the colon; all other + // leading and trailing whitespace belongs to the field value. + if (line.length() > offset && line.charAt(offset) == ' ') { + offset++; + } + return line.substring(offset); + } + /** * Emits the pending event, if a {@code data:} field was seen, and resets the * per-event state. The event type is reset even when nothing is dispatched, as @@ -543,7 +552,9 @@ Optional flush() { if (data.isEmpty()) { return Optional.empty(); } - SseEvent result = new SseEvent(id, type != null ? type : DEFAULT_EVENT_TYPE, data.toString().trim()); + // Remove only the final LF appended by the parser, preserving the payload. + SseEvent result = new SseEvent(id, type != null ? type : DEFAULT_EVENT_TYPE, + data.substring(0, data.length() - 1)); data.setLength(0); return Optional.of(result); } diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java index 459e42af0..3aa4b92a7 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java @@ -5,13 +5,20 @@ package io.modelcontextprotocol.client.transport; import java.util.Optional; +import java.util.stream.Stream; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import org.junit.jupiter.params.provider.ValueSource; import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEvent; import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEventParser; +import io.modelcontextprotocol.spec.McpTransportException; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; class SseEventParserTests { @@ -27,13 +34,56 @@ void simpleDataEvent() { } @Test - void multiLineDataAccumulatesWithNewlineSeparatorAndTrims() { + void multiLineDataAccumulatesWithNewlineSeparatorAndPreservesWhitespace() { SseEventParser p = new SseEventParser(Integer.MAX_VALUE); - assertThat(p.feed("data: first")).isEmpty(); - assertThat(p.feed("data: second")).isEmpty(); + assertThat(p.feed("data: first ")).isEmpty(); + assertThat(p.feed("data: \tsecond\t")).isEmpty(); Optional event = p.feed(""); assertThat(event).isPresent(); - assertThat(event.get().data()).isEqualTo("first\nsecond"); + assertThat(event.get().data()).isEqualTo(" first \n\tsecond\t"); + } + + @ParameterizedTest + @MethodSource("fieldValues") + void fieldValuesRemoveAtMostOneLeadingSpace(String input, String expected) { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("id:" + input); + p.feed("event:" + input); + p.feed("data:" + input); + SseEvent first = p.feed("").orElseThrow(); + assertThat(first.id()).isEqualTo(expected); + assertThat(first.event()).isEqualTo(expected); + assertThat(first.data()).isEqualTo(expected); + + p.feed("data: next"); + SseEvent second = p.feed("").orElseThrow(); + assertThat(second.id()).isEqualTo(expected); + assertThat(second.event()).isEqualTo("message"); + assertThat(second.data()).isEqualTo("next"); + } + + static Stream fieldValues() { + return Stream.of(Arguments.of("token", "token"), Arguments.of(" token", "token"), + Arguments.of(" token ", "token "), Arguments.of(" token ", " token "), + Arguments.of("\ttoken\t", "\ttoken\t"), Arguments.of(" \ttoken\t", "\ttoken\t"), + Arguments.of(" ", " "), Arguments.of("\t", "\t")); + } + + @Test + void emptyDataLinesPreserveLeadingAndTrailingNewlines() { + SseEventParser p = new SseEventParser(Integer.MAX_VALUE); + p.feed("data:"); + p.feed("data: payload"); + p.feed("data:"); + assertThat(p.feed("").orElseThrow().data()).isEqualTo("\npayload\n"); + } + + @Test + void preservedWhitespaceCountsTowardsSizeLimit() { + SseEventParser p = new SseEventParser(3); + p.feed("data: a "); + assertThat(p.feed("").orElseThrow().data()).isEqualTo(" a "); + assertThatThrownBy(() -> p.feed("data: a ")).isInstanceOf(McpTransportException.class); } @Test @@ -90,11 +140,12 @@ void emptyIdClearsLastEventId() { assertThat(p.feed("").orElseThrow().id()).isNull(); } - @Test - void idContainingNullIsIgnored() { + @ParameterizedTest + @ValueSource(strings = { "\0token", "token\0", "to\0ken" }) + void idContainingNullIsIgnored(String id) { SseEventParser p = new SseEventParser(Integer.MAX_VALUE); p.feed("id: 1"); - p.feed("id: 2\u0000" + "3"); + p.feed("id: " + id); p.feed("data: payload"); assertThat(p.feed("").orElseThrow().id()).isEqualTo("1"); }