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 + "\"}}"; + } + +}