diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java index f62cd7c71..f5c65faf2 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java @@ -15,6 +15,7 @@ import io.modelcontextprotocol.spec.McpError; import io.modelcontextprotocol.spec.McpSchema; import io.modelcontextprotocol.spec.McpTransportSessionNotFoundException; +import io.modelcontextprotocol.spec.McpTransportTerminatedException; import io.modelcontextprotocol.util.Assert; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -93,6 +94,9 @@ class LifecycleInitializer { private final AtomicReference initializationRef = new AtomicReference<>(); + /** Permanent transport failure that prevents any further initialization attempt. */ + private final AtomicReference terminalFailure = new AtomicReference<>(); + /** * The max timeout to await for the client-server connection to be initialized. */ @@ -209,7 +213,7 @@ private Mono await() { private void complete(McpSchema.InitializeResult initializeResult) { // inform all the subscribers waiting for the initialization - this.initSink.emitValue(initializeResult, Sinks.EmitFailureHandler.FAIL_FAST); + this.initSink.tryEmitValue(initializeResult); } private void cacheResult(McpSchema.InitializeResult initializeResult) { @@ -218,13 +222,24 @@ private void cacheResult(McpSchema.InitializeResult initializeResult) { } private void error(Throwable t) { - this.initSink.emitError(t, Sinks.EmitFailureHandler.FAIL_FAST); + this.initSink.tryEmitError(t); } private void close() { this.mcpSession().close(); } + private void terminate(Throwable cause) { + // Initialization has a single shared outcome. Publish the terminal failure + // even when the session has not been installed yet so that both the owner + // and all concurrent joiners observe it immediately. + this.initSink.tryEmitError(cause); + McpClientSession mcpClientSession = this.mcpSession(); + if (mcpClientSession != null) { + mcpClientSession.terminate(cause); + } + } + private Mono closeGracefully() { return this.mcpSession().closeGracefully(); } @@ -232,10 +247,13 @@ private Mono closeGracefully() { } public boolean isInitialized() { - return this.currentInitializationResult() != null; + return this.terminalFailure.get() == null && this.currentInitializationResult() != null; } public McpSchema.InitializeResult currentInitializationResult() { + if (this.terminalFailure.get() != null) { + return null; + } DefaultInitialization current = this.initializationRef.get(); McpSchema.InitializeResult initializeResult = current != null ? current.result.get() : null; return initializeResult; @@ -250,7 +268,15 @@ public McpSchema.InitializeResult currentInitializationResult() { * @param t The exception to handle */ public void handleException(Throwable t) { - if (t instanceof McpTransportSessionNotFoundException) { + if (t instanceof McpTransportTerminatedException) { + if (this.terminalFailure.compareAndSet(null, t)) { + DefaultInitialization current = this.initializationRef.get(); + if (current != null) { + current.terminate(t); + } + } + } + else if (t instanceof McpTransportSessionNotFoundException && this.terminalFailure.get() == null) { DefaultInitialization previous = this.initializationRef.getAndSet(null); if (previous != null) { previous.close(); @@ -271,14 +297,39 @@ public void handleException(Throwable t) { */ public Mono withInitialization(String actionName, Function> operation) { return Mono.deferContextual(ctx -> { + Throwable terminal = this.terminalFailure.get(); + if (terminal != null) { + return Mono.error(terminal); + } + DefaultInitialization newInit = new DefaultInitialization(); DefaultInitialization previous = this.initializationRef.compareAndExchange(null, newInit); boolean needsToInitialize = previous == null; + DefaultInitialization activeInitialization = needsToInitialize ? newInit : previous; + // Complete the handoff with handleException(): if termination won before + // registration, this check publishes it; if registration won first, the + // exception handler observes the active initialization. + Throwable terminalAfterRegistration = this.terminalFailure.get(); + if (terminalAfterRegistration != null) { + activeInitialization.terminate(terminalAfterRegistration); + } logger.debug(needsToInitialize ? "Initialization process started" : "Joining previous initialization"); - Mono initializationJob = needsToInitialize - ? this.doInitialize(newInit, this.postInitializationHook, ctx) : previous.await(); + Mono initializationJob; + if (needsToInitialize) { + // The work branch only publishes into the shared sink. Keeping it from + // winning directly makes the owner and all joiners consume the same + // first terminal signal. + Mono initializationWork = this + .doInitialize(newInit, this.postInitializationHook, ctx) + .onErrorComplete() + .then(Mono.never()); + initializationJob = Mono.firstWithSignal(activeInitialization.await(), initializationWork); + } + else { + initializationJob = activeInitialization.await(); + } return initializationJob.map(initializeResult -> this.initializationRef.get()) .timeout(this.initializationTimeout) @@ -295,42 +346,46 @@ public Mono withInitialization(String actionName, Function doInitialize(DefaultInitialization initialization, Function> postInitOperation, ContextView ctx) { - initialization.setMcpClientSession(this.sessionSupplier.apply(ctx)); - - McpClientSession mcpClientSession = initialization.mcpSession(); - - String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1); - - McpSchema.InitializeRequest initializeRequest = McpSchema.InitializeRequest - .builder(latestVersion, this.clientCapabilities, this.clientInfo) - .build(); - - Mono result = mcpClientSession.sendRequest(McpSchema.METHOD_INITIALIZE, - initializeRequest, McpAsyncClient.INITIALIZE_RESULT_TYPE_REF); - - return result.flatMap(initializeResult -> { - logger.info("Server response with Protocol: {}, Capabilities: {}, Info: {} and Instructions {}", - initializeResult.protocolVersion(), initializeResult.capabilities(), initializeResult.serverInfo(), - initializeResult.instructions()); + return Mono.defer(() -> { + initialization.setMcpClientSession(this.sessionSupplier.apply(ctx)); - if (!this.protocolVersions.contains(initializeResult.protocolVersion())) { - return Mono.error(McpError.builder(-32602) - .message("Unsupported protocol version") - .data("Unsupported protocol version from the server: " + initializeResult.protocolVersion()) - .build()); + McpClientSession mcpClientSession = initialization.mcpSession(); + Throwable terminal = this.terminalFailure.get(); + if (terminal != null) { + mcpClientSession.terminate(terminal); + return Mono.error(terminal); } - return mcpClientSession.sendNotification(McpSchema.METHOD_NOTIFICATION_INITIALIZED, null) - .contextWrite( - c -> c.put(McpAsyncClient.NEGOTIATED_PROTOCOL_VERSION, initializeResult.protocolVersion())) - .thenReturn(initializeResult); - }).flatMap(initializeResult -> { - initialization.cacheResult(initializeResult); - return postInitOperation.apply(initialization).thenReturn(initializeResult); - }).doOnNext(initialization::complete).onErrorResume(ex -> { - initialization.error(ex); - return Mono.error(ex); - }); + String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1); + + McpSchema.InitializeRequest initializeRequest = McpSchema.InitializeRequest + .builder(latestVersion, this.clientCapabilities, this.clientInfo) + .build(); + + Mono result = mcpClientSession.sendRequest(McpSchema.METHOD_INITIALIZE, + initializeRequest, McpAsyncClient.INITIALIZE_RESULT_TYPE_REF); + + return result.flatMap(initializeResult -> { + logger.info("Server response with Protocol: {}, Capabilities: {}, Info: {} and Instructions {}", + initializeResult.protocolVersion(), initializeResult.capabilities(), + initializeResult.serverInfo(), initializeResult.instructions()); + + if (!this.protocolVersions.contains(initializeResult.protocolVersion())) { + return Mono.error(McpError.builder(-32602) + .message("Unsupported protocol version") + .data("Unsupported protocol version from the server: " + initializeResult.protocolVersion()) + .build()); + } + + return mcpClientSession.sendNotification(McpSchema.METHOD_NOTIFICATION_INITIALIZED, null) + .contextWrite( + c -> c.put(McpAsyncClient.NEGOTIATED_PROTOCOL_VERSION, initializeResult.protocolVersion())) + .thenReturn(initializeResult); + }).flatMap(initializeResult -> { + initialization.cacheResult(initializeResult); + return postInitOperation.apply(initialization).thenReturn(initializeResult); + }); + }).doOnNext(initialization::complete).doOnError(initialization::error); } /** @@ -355,4 +410,4 @@ public Mono closeGracefully() { }); } -} \ No newline at end of file +} diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/McpStdioServerProcessExitException.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/McpStdioServerProcessExitException.java new file mode 100644 index 000000000..1a8b35c57 --- /dev/null +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/McpStdioServerProcessExitException.java @@ -0,0 +1,42 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import io.modelcontextprotocol.spec.McpTransportTerminatedException; +import io.modelcontextprotocol.util.Assert; + +/** + * Thrown when an MCP stdio server process exits unexpectedly. + * + * @author Dongliang Xie + */ +public class McpStdioServerProcessExitException extends McpTransportTerminatedException { + + private static final long serialVersionUID = 1L; + + private final int exitCode; + + private final String command; + + public McpStdioServerProcessExitException(int exitCode, String command) { + super(message(exitCode, command)); + this.exitCode = exitCode; + this.command = command; + } + + public int getExitCode() { + return this.exitCode; + } + + public String getCommand() { + return this.command; + } + + private static String message(int exitCode, String command) { + Assert.hasText(command, "The command can not be empty"); + return "MCP server process exited unexpectedly with code " + exitCode + " for command: " + command; + } + +} diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java index e73e43ef5..c3945b69c 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java @@ -10,13 +10,12 @@ import java.nio.charset.StandardCharsets; import java.time.Duration; import java.util.ArrayList; -import java.util.EnumSet; import java.util.List; import java.util.Set; import java.util.concurrent.Executors; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; import java.util.function.Function; -import java.util.stream.IntStream; import io.modelcontextprotocol.json.TypeRef; import io.modelcontextprotocol.json.McpJsonMapper; @@ -42,6 +41,16 @@ */ public class StdioClientTransport implements McpClientTransport { + private enum State { + + NEW, STARTING, RUNNING, CLOSING, TERMINATED + + } + + private record Lifecycle(State state, McpStdioServerProcessExitException terminalCause) { + + } + private static final Logger logger = LoggerFactory.getLogger(StdioClientTransport.class); // @formatter:off @@ -58,7 +67,11 @@ public class StdioClientTransport implements McpClientTransport { private final Sinks.Many outboundSink; /** The server process being communicated with */ - private Process process; + private volatile Process process; + + private final AtomicReference lifecycle = new AtomicReference<>(new Lifecycle(State.NEW, null)); + + private final AtomicReference> exceptionHandler = new AtomicReference<>(); private McpJsonMapper jsonMapper; @@ -115,9 +128,18 @@ public StdioClientTransport(ServerParameters params, McpJsonMapper jsonMapper) { @Override public Mono connect(Function, Mono> handler) { return Mono.fromRunnable(() -> { + Lifecycle current = this.lifecycle.get(); + Lifecycle starting = new Lifecycle(State.STARTING, null); + if (current.state() != State.NEW || !this.lifecycle.compareAndSet(current, starting)) { + current = this.lifecycle.get(); + if (current.terminalCause() != null) { + throw current.terminalCause(); + } + throw new IllegalStateException( + "Stdio client transport cannot be connected from state " + current.state()); + } + logger.info("MCP server starting."); - handleIncomingMessages(handler); - handleIncomingErrors(); // Prepare command and environment List fullCommand = new ArrayList<>(); @@ -133,19 +155,41 @@ public Mono connect(Function, Mono> h this.process = processBuilder.start(); } catch (IOException e) { + this.lifecycle.compareAndSet(starting, new Lifecycle(State.NEW, null)); throw new RuntimeException("Failed to start process with command: " + fullCommand, e); } // Validate process streams if (this.process.getInputStream() == null || process.getOutputStream() == null) { this.process.destroy(); + this.lifecycle.compareAndSet(starting, new Lifecycle(State.NEW, null)); throw new RuntimeException("Process input or output stream is null"); } + Lifecycle running = new Lifecycle(State.RUNNING, null); + if (!this.lifecycle.compareAndSet(starting, running)) { + this.process.destroy(); + current = this.lifecycle.get(); + if (current.terminalCause() != null) { + throw current.terminalCause(); + } + throw new IllegalStateException( + "Stdio client transport startup interrupted in state " + current.state()); + } + handleIncomingMessages(handler); + handleIncomingErrors(); + // Start threads startInboundProcessing(); startOutboundProcessing(); startErrorProcessing(); + startExitMonitoring(); + if (!this.process.isAlive()) { + McpStdioServerProcessExitException terminal = signalUnexpectedProcessExit(this.process.exitValue()); + if (terminal != null) { + throw terminal; + } + } logger.info("MCP server started"); }).subscribeOn(Schedulers.boundedElastic()); } @@ -172,6 +216,11 @@ public void setStdErrorHandler(Consumer errorHandler) { this.stdErrorHandler = errorHandler; } + @Override + public void setExceptionHandler(Consumer handler) { + this.exceptionHandler.set(handler); + } + /** * Waits for the server process to exit. * @throws RuntimeException if the process is interrupted while waiting @@ -239,6 +288,17 @@ private void handleIncomingErrors() { @Override public Mono sendMessage(JSONRPCMessage message) { + Lifecycle current = this.lifecycle.get(); + if (current.terminalCause() != null) { + return Mono.error(current.terminalCause()); + } + if ((current.state() == State.STARTING || current.state() == State.RUNNING) && this.process != null + && !this.process.isAlive()) { + McpStdioServerProcessExitException exitException = signalUnexpectedProcessExit(this.process.exitValue()); + if (exitException != null) { + return Mono.error(exitException); + } + } if (this.outboundSink.tryEmitNext(message).isSuccess()) { // TODO: essentially we could reschedule ourselves in some time and make // another attempt with the already read data but pause reading until @@ -248,10 +308,67 @@ public Mono sendMessage(JSONRPCMessage message) { return Mono.empty(); } else { + current = this.lifecycle.get(); + if (current.terminalCause() != null) { + return Mono.error(current.terminalCause()); + } return Mono.error(new RuntimeException("Failed to enqueue message")); } } + private void startExitMonitoring() { + this.process.onExit().thenAccept(process -> { + while (true) { + Lifecycle current = this.lifecycle.get(); + if (current.state() == State.TERMINATED) { + return; + } + if (current.state() == State.CLOSING) { + if (this.lifecycle.compareAndSet(current, new Lifecycle(State.TERMINATED, null))) { + return; + } + continue; + } + signalUnexpectedProcessExit(process.exitValue()); + return; + } + }); + } + + private McpStdioServerProcessExitException signalUnexpectedProcessExit(int exitCode) { + McpStdioServerProcessExitException exception = new McpStdioServerProcessExitException(exitCode, + this.params.getCommand()); + while (true) { + Lifecycle current = this.lifecycle.get(); + if (current.terminalCause() != null) { + return current.terminalCause(); + } + if (current.state() == State.CLOSING || current.state() == State.TERMINATED) { + return null; + } + if (this.lifecycle.compareAndSet(current, new Lifecycle(State.TERMINATED, exception))) { + break; + } + } + + logger.warn(exception.getMessage()); + this.isClosing = true; + this.inboundSink.tryEmitComplete(); + this.outboundSink.tryEmitComplete(); + this.errorSink.tryEmitComplete(); + + Consumer handler = this.exceptionHandler.get(); + if (handler != null) { + try { + handler.accept(exception); + } + catch (RuntimeException handlerFailure) { + logger.error("Transport exception handler failed", handlerFailure); + } + } + return exception; + } + /** * Starts the inbound processing thread that reads JSON-RPC messages from the * process's input stream. Messages are deserialized and emitted to the inbound sink. @@ -347,7 +464,9 @@ protected void handleOutbound(Function, Flux closeGracefully() { return Mono.fromRunnable(() -> { - isClosing = true; + this.lifecycle.updateAndGet( + current -> current.state() == State.TERMINATED ? current : new Lifecycle(State.CLOSING, null)); + this.isClosing = true; logger.debug("Initiating graceful shutdown"); }).then(Mono.defer(() -> { // First complete all sinks to stop accepting new messages @@ -388,7 +507,11 @@ public Mono closeGracefully() { catch (Exception e) { logger.error("Error during graceful shutdown", e); } - })).then().subscribeOn(Schedulers.boundedElastic()); + })) + .doFinally(signalType -> this.lifecycle.updateAndGet( + current -> current.state() == State.CLOSING ? new Lifecycle(State.TERMINATED, null) : current)) + .then() + .subscribeOn(Schedulers.boundedElastic()); } public Sinks.Many getErrorSink() { diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java index 3d7154278..64956d601 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java @@ -17,6 +17,7 @@ import java.util.UUID; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; /** @@ -50,6 +51,9 @@ public class McpClientSession implements McpSession { /** Map of pending responses keyed by request ID */ private final ConcurrentHashMap> pendingResponses = new ConcurrentHashMap<>(); + /** Permanent failure that prevents this session from sending further messages. */ + private final AtomicReference terminalCause = new AtomicReference<>(); + /** Map of request handlers keyed by method name */ private final ConcurrentHashMap> requestHandlers = new ConcurrentHashMap<>(); @@ -123,12 +127,17 @@ public McpClientSession(Duration requestTimeout, McpClientTransport transport, }, error -> logger.warn("Client failed during connect", error)); } - private void dismissPendingResponses() { + private void dismissPendingResponses(Throwable cause) { this.pendingResponses.forEach((id, sink) -> { - logger.info("Abruptly terminating exchange for request {}", id); - sink.error(new RuntimeException("MCP session with server terminated")); + if (this.pendingResponses.remove(id, sink)) { + logger.warn("Abruptly terminating exchange for request {}: {}", id, cause.toString()); + sink.error(cause); + } }); - this.pendingResponses.clear(); + } + + private void dismissPendingResponses() { + dismissPendingResponses(new RuntimeException("MCP session with server terminated")); } private void handle(McpSchema.JSONRPCMessage message) { @@ -256,13 +265,27 @@ public Mono sendRequest(String method, Object requestParams, TypeRef t String requestId = this.generateRequestId(); return Mono.deferContextual(ctx -> Mono.create(pendingResponseSink -> { + Throwable terminal = this.terminalCause.get(); + if (terminal != null) { + pendingResponseSink.error(terminal); + return; + } + logger.debug("Sending message for method {}", method); this.pendingResponses.put(requestId, pendingResponseSink); + + terminal = this.terminalCause.get(); + if (terminal != null && this.pendingResponses.remove(requestId, pendingResponseSink)) { + pendingResponseSink.error(terminal); + return; + } + McpSchema.JSONRPCRequest jsonrpcRequest = new McpSchema.JSONRPCRequest(method, requestId, requestParams); this.transport.sendMessage(jsonrpcRequest).contextWrite(ctx).subscribe(v -> { }, error -> { - this.pendingResponses.remove(requestId); - pendingResponseSink.error(error); + if (this.pendingResponses.remove(requestId, pendingResponseSink)) { + pendingResponseSink.error(error); + } }); })).timeout(this.requestTimeout).handle((jsonRpcResponse, deliveredResponseSink) -> { if (jsonRpcResponse.error() != null) { @@ -289,8 +312,14 @@ public Mono sendRequest(String method, Object requestParams, TypeRef t */ @Override public Mono sendNotification(String method, Object params) { - McpSchema.JSONRPCNotification jsonrpcNotification = new McpSchema.JSONRPCNotification(method, params); - return this.transport.sendMessage(jsonrpcNotification); + return Mono.defer(() -> { + Throwable terminal = this.terminalCause.get(); + if (terminal != null) { + return Mono.error(terminal); + } + McpSchema.JSONRPCNotification jsonrpcNotification = new McpSchema.JSONRPCNotification(method, params); + return this.transport.sendMessage(jsonrpcNotification); + }); } /** @@ -310,4 +339,16 @@ public void close() { dismissPendingResponses(); } + /** + * Permanently terminates the session because its transport can no longer be used. + * Pending and future operations fail with the first terminal cause observed. + * @param cause the permanent transport failure + */ + public void terminate(Throwable cause) { + Assert.notNull(cause, "The cause can not be null"); + if (this.terminalCause.compareAndSet(null, cause)) { + dismissPendingResponses(cause); + } + } + } diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpTransportTerminatedException.java b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpTransportTerminatedException.java new file mode 100644 index 000000000..aee906b29 --- /dev/null +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpTransportTerminatedException.java @@ -0,0 +1,33 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.spec; + +/** + * Exception raised when a transport instance becomes permanently unusable. + * + *

+ * Unlike recoverable transport errors such as a missing HTTP session, this exception + * indicates that the current transport instance cannot establish or resume an MCP session + * and must not be reused. + * + * @author Dongliang Xie + */ +public class McpTransportTerminatedException extends McpTransportException { + + private static final long serialVersionUID = 1L; + + public McpTransportTerminatedException(String message) { + super(message); + } + + public McpTransportTerminatedException(String message, Throwable cause) { + super(message, cause); + } + + public McpTransportTerminatedException(Throwable cause) { + super(cause); + } + +} diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerPostInitializationHookTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerPostInitializationHookTests.java index f9b5401fe..951d89ac0 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerPostInitializationHookTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerPostInitializationHookTests.java @@ -6,6 +6,8 @@ import java.time.Duration; import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; @@ -14,11 +16,13 @@ import io.modelcontextprotocol.spec.McpClientSession; import io.modelcontextprotocol.spec.McpSchema; import io.modelcontextprotocol.spec.McpTransportSessionNotFoundException; +import io.modelcontextprotocol.spec.McpTransportTerminatedException; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.mockito.Mock; import org.mockito.MockitoAnnotations; import reactor.core.publisher.Mono; +import reactor.core.publisher.Sinks; import reactor.core.scheduler.Schedulers; import reactor.test.StepVerifier; import reactor.util.context.ContextView; @@ -164,6 +168,38 @@ void shouldFailInitializationWhenPostInitializationHookFails() { verify(mockPostInitializationHook, times(1)).apply(any(Initialization.class)); } + @Test + void shouldFailInitializationWhenTransportTerminatesDuringPostInitializationHook() throws Exception { + var cause = new McpTransportTerminatedException("Transport terminated during post-initialization hook"); + var hookEntered = new CountDownLatch(1); + var hookGate = Sinks.one(); + + when(mockPostInitializationHook.apply(any(Initialization.class))).thenAnswer(invocation -> Mono.defer(() -> { + hookEntered.countDown(); + return hookGate.asMono(); + })); + + var initialization = initializer.withInitialization("test", init -> Mono.just(init.initializeResult())) + .materialize() + .toFuture(); + assertThat(hookEntered.await(1, TimeUnit.SECONDS)).isTrue(); + + initializer.handleException(cause); + + try { + var signal = initialization.get(1, TimeUnit.SECONDS); + assertThat(signal.isOnError()).isTrue(); + assertThat(signal.getThrowable()).hasCause(cause); + } + finally { + hookGate.tryEmitEmpty(); + } + + assertThat(initializer.isInitialized()).isFalse(); + assertThat(initializer.currentInitializationResult()).isNull(); + verify(mockClientSession).terminate(cause); + } + @Test void shouldNotInvokePostInitializationHookWhenInitializationFails() { when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())) diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java index be7b7dc53..407be06df 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java @@ -6,6 +6,8 @@ import java.time.Duration; import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; @@ -13,7 +15,9 @@ import io.modelcontextprotocol.client.LifecycleInitializer.Initialization; import io.modelcontextprotocol.spec.McpClientSession; import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.McpTransportException; import io.modelcontextprotocol.spec.McpTransportSessionNotFoundException; +import io.modelcontextprotocol.spec.McpTransportTerminatedException; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.mockito.Mock; @@ -302,6 +306,175 @@ void shouldHandleOtherExceptions() { verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); } + @Test + void shouldTerminateInProgressInitializationOnTransportTermination() throws Exception { + var cause = new McpTransportTerminatedException("Transport terminated"); + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenReturn(Mono.never()); + + var initialization = initializer.withInitialization("test", init -> Mono.just(init.initializeResult())) + .materialize() + .toFuture(); + + initializer.handleException(cause); + + var signal = initialization.get(1, TimeUnit.SECONDS); + assertThat(signal.isOnError()).isTrue(); + assertThat(signal.getThrowable()).hasCause(cause); + verify(mockClientSession).terminate(cause); + assertThat(initializer.isInitialized()).isFalse(); + assertThat(initializer.currentInitializationResult()).isNull(); + + StepVerifier.create(initializer.withInitialization("retry", init -> Mono.just(init.initializeResult()))) + .expectErrorSatisfies(error -> assertThat(error).isSameAs(cause)) + .verify(); + + verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); + } + + @Test + void shouldApplyTerminationThatArrivesBeforeSessionRegistration() { + var cause = new McpTransportTerminatedException("Transport terminated early"); + when(mockSessionSupplier.apply(any(ContextView.class))).thenAnswer(invocation -> { + initializer.handleException(cause); + return mockClientSession; + }); + + StepVerifier.create(initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))) + .expectErrorSatisfies(error -> assertThat(error).hasCause(cause)) + .verify(); + + verify(mockClientSession).terminate(cause); + verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any()); + } + + @Test + void shouldTerminateWinnerAndJoinerWhenTerminationArrivesBeforeSessionRegistration() throws Exception { + var cause = new McpTransportTerminatedException("Transport terminated before session registration"); + var supplierEntered = new CountDownLatch(1); + var releaseSupplier = new CountDownLatch(1); + + when(mockSessionSupplier.apply(any(ContextView.class))).thenAnswer(invocation -> { + supplierEntered.countDown(); + if (!releaseSupplier.await(5, TimeUnit.SECONDS)) { + throw new IllegalStateException("Timed out waiting to release session supplier"); + } + return mockClientSession; + }); + + var winner = initializer.withInitialization("winner", init -> Mono.just(init.initializeResult())) + .subscribeOn(Schedulers.boundedElastic()) + .materialize() + .toFuture(); + assertThat(supplierEntered.await(1, TimeUnit.SECONDS)).isTrue(); + + var joiner = initializer.withInitialization("joiner", init -> Mono.just(init.initializeResult())) + .materialize() + .toFuture(); + + initializer.handleException(cause); + + try { + var winnerSignal = winner.get(1, TimeUnit.SECONDS); + var joinerSignal = joiner.get(1, TimeUnit.SECONDS); + + assertThat(winnerSignal.isOnError()).isTrue(); + assertThat(winnerSignal.getThrowable()).hasCause(cause); + assertThat(joinerSignal.isOnError()).isTrue(); + assertThat(joinerSignal.getThrowable()).hasCause(cause); + } + finally { + releaseSupplier.countDown(); + } + + verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); + verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any()); + } + + @Test + void shouldShareSynchronousSessionSupplierFailureBetweenWinnerAndJoiner() throws Exception { + var cause = new IllegalStateException("Session supplier failed"); + var supplierEntered = new CountDownLatch(1); + var releaseSupplier = new CountDownLatch(1); + + when(mockSessionSupplier.apply(any(ContextView.class))).thenAnswer(invocation -> { + supplierEntered.countDown(); + if (!releaseSupplier.await(5, TimeUnit.SECONDS)) { + throw new IllegalStateException("Timed out waiting to release session supplier"); + } + throw cause; + }); + + var winner = initializer.withInitialization("winner", init -> Mono.just(init.initializeResult())) + .subscribeOn(Schedulers.boundedElastic()) + .materialize() + .toFuture(); + assertThat(supplierEntered.await(1, TimeUnit.SECONDS)).isTrue(); + + var joiner = initializer.withInitialization("joiner", init -> Mono.just(init.initializeResult())) + .materialize() + .toFuture(); + + releaseSupplier.countDown(); + + var winnerSignal = winner.get(1, TimeUnit.SECONDS); + var joinerSignal = joiner.get(1, TimeUnit.SECONDS); + + assertThat(winnerSignal.isOnError()).isTrue(); + assertThat(winnerSignal.getThrowable()).hasCause(cause); + assertThat(joinerSignal.isOnError()).isTrue(); + assertThat(joinerSignal.getThrowable()).hasCause(cause); + verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); + verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any()); + } + + @Test + void shouldIgnoreGenericTransportExceptionDuringInitialization() { + var cause = new McpTransportException("Transport closed"); + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenReturn(Mono.never()); + + var subscription = initializer.withInitialization("test", init -> Mono.just(init.initializeResult())) + .subscribe(); + + initializer.handleException(cause); + subscription.dispose(); + + verify(mockClientSession, never()).terminate(cause); + } + + @Test + void shouldKeepInitializedAfterTransportException() { + StepVerifier.create(initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))) + .expectNext(MOCK_INIT_RESULT) + .verifyComplete(); + + var cause = new McpTransportException("Transport closed"); + + initializer.handleException(cause); + + assertThat(initializer.isInitialized()).isTrue(); + verify(mockClientSession, never()).terminate(cause); + verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); + } + + @Test + void shouldBecomeTerminalAfterInitializationCompletes() { + StepVerifier.create(initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))) + .expectNext(MOCK_INIT_RESULT) + .verifyComplete(); + + var cause = new McpTransportTerminatedException("Transport terminated"); + initializer.handleException(cause); + + assertThat(initializer.isInitialized()).isFalse(); + assertThat(initializer.currentInitializationResult()).isNull(); + verify(mockClientSession).terminate(cause); + + StepVerifier.create(initializer.withInitialization("retry", init -> Mono.just(init.initializeResult()))) + .expectErrorSatisfies(error -> assertThat(error).isSameAs(cause)) + .verify(); + verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); + } + @Test void shouldCloseGracefully() { StepVerifier.create(initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))) diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java index ae5daf1f4..64fd2363b 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java @@ -107,6 +107,62 @@ void testRequestTimeout() { session.close(); } + @Test + void testPendingRequestFailsWithTerminalCause() { + var transport = new MockMcpClientTransport(); + var session = new McpClientSession(TIMEOUT, transport, Map.of(), + Map.of(TEST_NOTIFICATION, params -> Mono.fromRunnable(() -> logger.info("Status update: {}", params))), + Function.identity()); + var cause = new McpTransportException("Transport closed"); + + Mono responseMono = session.sendRequest(TEST_METHOD, "test", responseType); + + StepVerifier.create(responseMono).then(() -> session.terminate(cause)).expectErrorSatisfies(error -> { + assertThat(error).isSameAs(cause); + }).verify(); + } + + @Test + void testRequestFailsImmediatelyAfterTermination() { + var transport = new MockMcpClientTransport(); + var session = new McpClientSession(TIMEOUT, transport, Map.of(), Map.of(), Function.identity()); + var cause = new McpTransportTerminatedException("Transport terminated"); + + session.terminate(cause); + + StepVerifier.create(session.sendRequest(TEST_METHOD, "test", responseType)) + .expectErrorSatisfies(error -> assertThat(error).isSameAs(cause)) + .verify(); + } + + @Test + void testNotificationFailsImmediatelyAfterTermination() { + var transport = new MockMcpClientTransport(); + var session = new McpClientSession(TIMEOUT, transport, Map.of(), Map.of(), Function.identity()); + var cause = new McpTransportTerminatedException("Transport terminated"); + + session.terminate(cause); + + StepVerifier.create(session.sendNotification(TEST_NOTIFICATION, Map.of())) + .expectErrorSatisfies(error -> assertThat(error).isSameAs(cause)) + .verify(); + } + + @Test + void testFirstTerminalCauseWins() { + var transport = new MockMcpClientTransport(); + var session = new McpClientSession(TIMEOUT, transport, Map.of(), Map.of(), Function.identity()); + var firstCause = new McpTransportTerminatedException("First failure"); + var secondCause = new McpTransportTerminatedException("Second failure"); + + session.terminate(firstCause); + session.terminate(secondCause); + + StepVerifier.create(session.sendRequest(TEST_METHOD, "test", responseType)) + .expectErrorSatisfies(error -> assertThat(error).isSameAs(firstCause)) + .verify(); + } + @Test void testSendNotification() { var transport = new MockMcpClientTransport(); diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/FailingStdioServer.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/FailingStdioServer.java new file mode 100644 index 000000000..63fbaa2f9 --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/FailingStdioServer.java @@ -0,0 +1,17 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client; + +final class FailingStdioServer { + + private FailingStdioServer() { + } + + public static void main(String[] args) { + System.err.println("Exiting before MCP initialization with code 127"); + System.exit(127); + } + +} diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java new file mode 100644 index 000000000..5658f18e2 --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java @@ -0,0 +1,107 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client; + +import java.nio.file.Path; +import java.time.Duration; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Function; + +import io.modelcontextprotocol.client.transport.McpStdioServerProcessExitException; +import io.modelcontextprotocol.client.transport.ServerParameters; +import io.modelcontextprotocol.client.transport.StdioClientTransport; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; + +import static io.modelcontextprotocol.util.McpJsonMapperUtils.JSON_MAPPER; +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.catchThrowable; + +/** + * Tests for initialization failures reported by {@link StdioClientTransport}. + * + * @author Dongliang Xie + */ +@Timeout(10) +class StdioMcpClientInitializationFailureTests { + + @Test + void initializeShouldFailWithProcessExitInsteadOfRequestTimeout() { + Duration requestTimeout = Duration.ofSeconds(3); + String classpath = System.getProperty("java.class.path"); + ServerParameters stdioParams = ServerParameters.builder(javaExecutable()) + .args("-cp", classpath, FailingStdioServer.class.getName()) + .build(); + StdioClientTransport transport = new StdioClientTransport(stdioParams, JSON_MAPPER); + McpSyncClient client = McpClient.sync(transport) + .requestTimeout(requestTimeout) + .initializationTimeout(Duration.ofSeconds(5)) + .build(); + + Throwable failure; + Throwable retryFailure; + long elapsedMillis; + long retryElapsedMillis; + try { + long startNanos = System.nanoTime(); + failure = catchThrowable(client::initialize); + elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startNanos); + + long retryStartNanos = System.nanoTime(); + retryFailure = catchThrowable(client::initialize); + retryElapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - retryStartNanos); + } + finally { + client.closeGracefully(); + } + + assertThat(failure).isNotNull(); + assertThat(elapsedMillis).isLessThan(requestTimeout.toMillis()); + McpStdioServerProcessExitException processExit = findCause(failure, McpStdioServerProcessExitException.class); + assertThat(processExit).isNotNull(); + assertThat(processExit.getExitCode()).isEqualTo(127); + assertThat(processExit.getCommand()).isEqualTo(javaExecutable()); + + assertThat(retryFailure).isNotNull(); + assertThat(retryElapsedMillis).isLessThan(requestTimeout.toMillis()); + McpStdioServerProcessExitException retryProcessExit = findCause(retryFailure, + McpStdioServerProcessExitException.class); + assertThat(retryProcessExit).isSameAs(processExit); + } + + @Test + void gracefulCloseShouldNotReportUnexpectedProcessExit() { + String classpath = System.getProperty("java.class.path"); + ServerParameters stdioParams = ServerParameters.builder(javaExecutable()) + .args("-cp", classpath, WaitingStdioServer.class.getName()) + .build(); + StdioClientTransport transport = new StdioClientTransport(stdioParams, JSON_MAPPER); + AtomicReference transportFailure = new AtomicReference<>(); + transport.setExceptionHandler(transportFailure::set); + + transport.connect(Function.identity()).block(Duration.ofSeconds(3)); + transport.closeGracefully().block(Duration.ofSeconds(5)); + + assertThat(transportFailure).hasValue(null); + } + + private String javaExecutable() { + String executable = System.getProperty("os.name").toLowerCase().contains("win") ? "java.exe" : "java"; + return Path.of(System.getProperty("java.home"), "bin", executable).toString(); + } + + private T findCause(Throwable failure, Class type) { + Throwable current = failure; + while (current != null) { + if (type.isInstance(current)) { + return type.cast(current); + } + current = current.getCause(); + } + return null; + } + +} diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/WaitingStdioServer.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/WaitingStdioServer.java new file mode 100644 index 000000000..05946f694 --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/WaitingStdioServer.java @@ -0,0 +1,16 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client; + +final class WaitingStdioServer { + + private WaitingStdioServer() { + } + + public static void main(String[] args) throws InterruptedException { + Thread.sleep(30_000); + } + +}