From 1c7a13fd4a545f1b6c0bc19ee7cf7cfef7ef60cc Mon Sep 17 00:00:00 2001 From: Dongliang Xie Date: Fri, 15 May 2026 08:40:49 +0800 Subject: [PATCH 1/4] fix: propagate stdio process exit during initialization Signed-off-by: Dongliang Xie --- .../client/LifecycleInitializer.java | 20 +++++- .../McpStdioServerProcessExitException.java | 42 +++++++++++ .../transport/StdioClientTransport.java | 48 +++++++++++++ .../spec/McpClientSession.java | 19 ++++- .../client/LifecycleInitializerTests.java | 52 ++++++++++++++ .../spec/McpClientSessionTests.java | 15 ++++ .../client/FailingStdioServer.java | 17 +++++ ...ioMcpClientInitializationFailureTests.java | 72 +++++++++++++++++++ 8 files changed, 281 insertions(+), 4 deletions(-) create mode 100644 mcp-core/src/main/java/io/modelcontextprotocol/client/transport/McpStdioServerProcessExitException.java create mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/FailingStdioServer.java create mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java 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..52cf9c7e8 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java @@ -11,6 +11,7 @@ import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; +import io.modelcontextprotocol.client.transport.McpStdioServerProcessExitException; import io.modelcontextprotocol.spec.McpClientSession; import io.modelcontextprotocol.spec.McpError; import io.modelcontextprotocol.spec.McpSchema; @@ -225,6 +226,16 @@ private void close() { this.mcpSession().close(); } + private void close(Throwable cause) { + McpClientSession mcpClientSession = this.mcpSession(); + if (mcpClientSession != null) { + mcpClientSession.close(cause); + } + else { + this.error(cause); + } + } + private Mono closeGracefully() { return this.mcpSession().closeGracefully(); } @@ -259,6 +270,13 @@ public void handleException(Throwable t) { // the implicit initialization step. this.withInitialization("re-initializing", result -> Mono.empty()).subscribe(); } + else if (t instanceof McpStdioServerProcessExitException) { + DefaultInitialization previous = this.initializationRef.get(); + if (previous != null && previous.initializeResult() == null + && this.initializationRef.compareAndSet(previous, null)) { + previous.close(t); + } + } } /** @@ -355,4 +373,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..4a9da49a2 --- /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.McpTransportException; +import io.modelcontextprotocol.util.Assert; + +/** + * Thrown when an MCP stdio server process exits unexpectedly. + * + * @author Dongliang Xie + */ +public class McpStdioServerProcessExitException extends McpTransportException { + + 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..a4bf663db 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 @@ -14,6 +14,7 @@ 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; @@ -60,6 +61,10 @@ public class StdioClientTransport implements McpClientTransport { /** The server process being communicated with */ private Process process; + private final AtomicReference unexpectedExitException = new AtomicReference<>(); + + private final AtomicReference> exceptionHandler = new AtomicReference<>(); + private McpJsonMapper jsonMapper; /** Scheduler for handling inbound messages from the server process */ @@ -78,6 +83,8 @@ public class StdioClientTransport implements McpClientTransport { private volatile boolean isClosing = false; + private volatile boolean closeRequested = false; + // visible for tests private Consumer stdErrorHandler = error -> logger.info("STDERR Message received: {}", error); @@ -146,6 +153,7 @@ public Mono connect(Function, Mono> h startInboundProcessing(); startOutboundProcessing(); startErrorProcessing(); + startExitMonitoring(); logger.info("MCP server started"); }).subscribeOn(Schedulers.boundedElastic()); } @@ -172,6 +180,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 +252,14 @@ private void handleIncomingErrors() { @Override public Mono sendMessage(JSONRPCMessage message) { + McpStdioServerProcessExitException exitException = this.unexpectedExitException.get(); + if (exitException != null) { + return Mono.error(exitException); + } + if (!this.closeRequested && this.process != null && !this.process.isAlive()) { + exitException = signalUnexpectedProcessExit(this.process.exitValue()); + 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 @@ -252,6 +273,32 @@ public Mono sendMessage(JSONRPCMessage message) { } } + private void startExitMonitoring() { + this.process.onExit().thenAccept(process -> { + if (!closeRequested) { + signalUnexpectedProcessExit(process.exitValue()); + } + }); + } + + private McpStdioServerProcessExitException signalUnexpectedProcessExit(int exitCode) { + McpStdioServerProcessExitException exception = new McpStdioServerProcessExitException(exitCode, + this.params.getCommand()); + if (this.unexpectedExitException.compareAndSet(null, exception)) { + logger.warn(exception.getMessage()); + isClosing = true; + inboundSink.tryEmitComplete(); + outboundSink.tryEmitComplete(); + errorSink.tryEmitComplete(); + + Consumer handler = this.exceptionHandler.get(); + if (handler != null) { + handler.accept(exception); + } + } + return this.unexpectedExitException.get(); + } + /** * 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,6 +394,7 @@ protected void handleOutbound(Function, Flux closeGracefully() { return Mono.fromRunnable(() -> { + closeRequested = true; isClosing = true; logger.debug("Initiating graceful shutdown"); }).then(Mono.defer(() -> { 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..a005922de 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java @@ -123,14 +123,18 @@ 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")); + 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) { if (message instanceof McpSchema.JSONRPCResponse response) { logger.debug("Received response: {}", response); @@ -310,4 +314,13 @@ public void close() { dismissPendingResponses(); } + /** + * Closes the session immediately, failing pending operations with the given cause. + * @param cause the transport-level cause of the closure + */ + public void close(Throwable cause) { + Assert.notNull(cause, "The cause can not be null"); + dismissPendingResponses(cause); + } + } 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..64f3e37db 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java @@ -11,8 +11,10 @@ import java.util.function.Function; import io.modelcontextprotocol.client.LifecycleInitializer.Initialization; +import io.modelcontextprotocol.client.transport.McpStdioServerProcessExitException; import io.modelcontextprotocol.spec.McpClientSession; import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.McpTransportException; import io.modelcontextprotocol.spec.McpTransportSessionNotFoundException; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -302,6 +304,56 @@ void shouldHandleOtherExceptions() { verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); } + @Test + void shouldCloseInProgressInitializationOnStdioProcessExit() { + var cause = new McpStdioServerProcessExitException(127, "java"); + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenReturn(Mono.never()) + .thenReturn(Mono.just(MOCK_INIT_RESULT)); + + var subscription = initializer.withInitialization("test", init -> Mono.just(init.initializeResult())) + .subscribe(); + + initializer.handleException(cause); + subscription.dispose(); + + verify(mockClientSession).close(cause); + + StepVerifier.create(initializer.withInitialization("retry", init -> Mono.just(init.initializeResult()))) + .expectNext(MOCK_INIT_RESULT) + .verifyComplete(); + + verify(mockSessionSupplier, times(2)).apply(any(ContextView.class)); + } + + @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()).close(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()).close(cause); + 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..13cdff1af 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,21 @@ void testRequestTimeout() { session.close(); } + @Test + void testPendingRequestFailsWithCloseCause() { + 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.close(cause)).expectErrorSatisfies(error -> { + assertThat(error).isSameAs(cause); + }).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..c36540bb8 --- /dev/null +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java @@ -0,0 +1,72 @@ +/* + * Copyright 2024-2026 the original author or authors. + */ + +package io.modelcontextprotocol.client; + +import java.io.PrintWriter; +import java.io.StringWriter; +import java.nio.file.Path; +import java.time.Duration; +import java.util.concurrent.TimeUnit; + +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; + long elapsedMillis; + try { + long startNanos = System.nanoTime(); + failure = catchThrowable(client::initialize); + elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startNanos); + } + finally { + client.closeGracefully(); + } + + assertThat(failure).isNotNull(); + String stackTrace = stackTraceOf(failure); + assertThat(elapsedMillis).isLessThan(requestTimeout.toMillis()); + assertThat(stackTrace).contains("MCP server process exited", "with code 127") + .doesNotContain("TimeoutException"); + } + + 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 String stackTraceOf(Throwable failure) { + StringWriter writer = new StringWriter(); + failure.printStackTrace(new PrintWriter(writer)); + return writer.toString(); + } + +} From 71fdb1559db3cbf5eedad32d4099f8b6a9163b0a Mon Sep 17 00:00:00 2001 From: dragonfsky Date: Wed, 29 Jul 2026 13:57:52 +0800 Subject: [PATCH 2/4] fix: make stdio process exit terminal --- .../client/LifecycleInitializer.java | 44 ++++-- .../McpStdioServerProcessExitException.java | 4 +- .../transport/StdioClientTransport.java | 129 ++++++++++++++---- .../spec/McpClientSession.java | 50 +++++-- .../spec/McpTransportTerminatedException.java | 33 +++++ .../client/LifecycleInitializerTests.java | 60 ++++++-- .../spec/McpClientSessionTests.java | 45 +++++- ...ioMcpClientInitializationFailureTests.java | 53 +++++-- .../client/WaitingStdioServer.java | 16 +++ 9 files changed, 356 insertions(+), 78 deletions(-) create mode 100644 mcp-core/src/main/java/io/modelcontextprotocol/spec/McpTransportTerminatedException.java create mode 100644 mcp-test/src/test/java/io/modelcontextprotocol/client/WaitingStdioServer.java 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 52cf9c7e8..663ace7a2 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java @@ -11,11 +11,11 @@ import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; -import io.modelcontextprotocol.client.transport.McpStdioServerProcessExitException; import io.modelcontextprotocol.spec.McpClientSession; 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; @@ -94,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. */ @@ -226,13 +229,10 @@ private void close() { this.mcpSession().close(); } - private void close(Throwable cause) { + private void terminate(Throwable cause) { McpClientSession mcpClientSession = this.mcpSession(); if (mcpClientSession != null) { - mcpClientSession.close(cause); - } - else { - this.error(cause); + mcpClientSession.terminate(cause); } } @@ -243,10 +243,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; @@ -261,7 +264,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(); @@ -270,13 +281,6 @@ public void handleException(Throwable t) { // the implicit initialization step. this.withInitialization("re-initializing", result -> Mono.empty()).subscribe(); } - else if (t instanceof McpStdioServerProcessExitException) { - DefaultInitialization previous = this.initializationRef.get(); - if (previous != null && previous.initializeResult() == null - && this.initializationRef.compareAndSet(previous, null)) { - previous.close(t); - } - } } /** @@ -289,6 +293,11 @@ else if (t instanceof McpStdioServerProcessExitException) { */ 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); @@ -316,6 +325,11 @@ private Mono doInitialize(DefaultInitialization init initialization.setMcpClientSession(this.sessionSupplier.apply(ctx)); McpClientSession mcpClientSession = initialization.mcpSession(); + Throwable terminal = this.terminalFailure.get(); + if (terminal != null) { + mcpClientSession.terminate(terminal); + return Mono.error(terminal); + } String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1); 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 index 4a9da49a2..1a8b35c57 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/McpStdioServerProcessExitException.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/McpStdioServerProcessExitException.java @@ -4,7 +4,7 @@ package io.modelcontextprotocol.client.transport; -import io.modelcontextprotocol.spec.McpTransportException; +import io.modelcontextprotocol.spec.McpTransportTerminatedException; import io.modelcontextprotocol.util.Assert; /** @@ -12,7 +12,7 @@ * * @author Dongliang Xie */ -public class McpStdioServerProcessExitException extends McpTransportException { +public class McpStdioServerProcessExitException extends McpTransportTerminatedException { private static final long serialVersionUID = 1L; 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 a4bf663db..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,14 +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; @@ -43,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 @@ -59,9 +67,9 @@ 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 unexpectedExitException = new AtomicReference<>(); + private final AtomicReference lifecycle = new AtomicReference<>(new Lifecycle(State.NEW, null)); private final AtomicReference> exceptionHandler = new AtomicReference<>(); @@ -83,8 +91,6 @@ public class StdioClientTransport implements McpClientTransport { private volatile boolean isClosing = false; - private volatile boolean closeRequested = false; - // visible for tests private Consumer stdErrorHandler = error -> logger.info("STDERR Message received: {}", error); @@ -122,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<>(); @@ -140,20 +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()); } @@ -252,13 +288,16 @@ private void handleIncomingErrors() { @Override public Mono sendMessage(JSONRPCMessage message) { - McpStdioServerProcessExitException exitException = this.unexpectedExitException.get(); - if (exitException != null) { - return Mono.error(exitException); + Lifecycle current = this.lifecycle.get(); + if (current.terminalCause() != null) { + return Mono.error(current.terminalCause()); } - if (!this.closeRequested && this.process != null && !this.process.isAlive()) { - exitException = signalUnexpectedProcessExit(this.process.exitValue()); - return Mono.error(exitException); + 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 @@ -269,14 +308,29 @@ 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 -> { - if (!closeRequested) { + 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; } }); } @@ -284,19 +338,35 @@ private void startExitMonitoring() { private McpStdioServerProcessExitException signalUnexpectedProcessExit(int exitCode) { McpStdioServerProcessExitException exception = new McpStdioServerProcessExitException(exitCode, this.params.getCommand()); - if (this.unexpectedExitException.compareAndSet(null, exception)) { - logger.warn(exception.getMessage()); - isClosing = true; - inboundSink.tryEmitComplete(); - outboundSink.tryEmitComplete(); - errorSink.tryEmitComplete(); + 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; + } + } - Consumer handler = this.exceptionHandler.get(); - if (handler != null) { + 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 this.unexpectedExitException.get(); + return exception; } /** @@ -394,8 +464,9 @@ protected void handleOutbound(Function, Flux closeGracefully() { return Mono.fromRunnable(() -> { - closeRequested = true; - 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 @@ -436,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 a005922de..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<>(); @@ -125,10 +129,11 @@ public McpClientSession(Duration requestTimeout, McpClientTransport transport, private void dismissPendingResponses(Throwable cause) { this.pendingResponses.forEach((id, sink) -> { - logger.warn("Abruptly terminating exchange for request {}: {}", id, cause.toString()); - sink.error(cause); + 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() { @@ -260,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) { @@ -293,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); + }); } /** @@ -315,12 +340,15 @@ public void close() { } /** - * Closes the session immediately, failing pending operations with the given cause. - * @param cause the transport-level cause of the closure + * 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 close(Throwable cause) { + public void terminate(Throwable cause) { Assert.notNull(cause, "The cause can not be null"); - dismissPendingResponses(cause); + 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/LifecycleInitializerTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java index 64f3e37db..12989e2f2 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java @@ -11,11 +11,11 @@ import java.util.function.Function; import io.modelcontextprotocol.client.LifecycleInitializer.Initialization; -import io.modelcontextprotocol.client.transport.McpStdioServerProcessExitException; 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; @@ -305,24 +305,41 @@ void shouldHandleOtherExceptions() { } @Test - void shouldCloseInProgressInitializationOnStdioProcessExit() { - var cause = new McpStdioServerProcessExitException(127, "java"); - when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenReturn(Mono.never()) - .thenReturn(Mono.just(MOCK_INIT_RESULT)); + void shouldTerminateInProgressInitializationOnTransportTermination() { + var cause = new McpTransportTerminatedException("Transport terminated"); + 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).close(cause); + verify(mockClientSession).terminate(cause); + assertThat(initializer.isInitialized()).isFalse(); + assertThat(initializer.currentInitializationResult()).isNull(); StepVerifier.create(initializer.withInitialization("retry", init -> Mono.just(init.initializeResult()))) - .expectNext(MOCK_INIT_RESULT) - .verifyComplete(); + .expectErrorSatisfies(error -> assertThat(error).isSameAs(cause)) + .verify(); - verify(mockSessionSupplier, times(2)).apply(any(ContextView.class)); + verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); + subscription.dispose(); + } + + @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 @@ -336,7 +353,7 @@ void shouldIgnoreGenericTransportExceptionDuringInitialization() { initializer.handleException(cause); subscription.dispose(); - verify(mockClientSession, never()).close(cause); + verify(mockClientSession, never()).terminate(cause); } @Test @@ -350,7 +367,26 @@ void shouldKeepInitializedAfterTransportException() { initializer.handleException(cause); assertThat(initializer.isInitialized()).isTrue(); - verify(mockClientSession, never()).close(cause); + 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)); } 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 13cdff1af..64fd2363b 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java @@ -108,7 +108,7 @@ void testRequestTimeout() { } @Test - void testPendingRequestFailsWithCloseCause() { + 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))), @@ -117,11 +117,52 @@ void testPendingRequestFailsWithCloseCause() { Mono responseMono = session.sendRequest(TEST_METHOD, "test", responseType); - StepVerifier.create(responseMono).then(() -> session.close(cause)).expectErrorSatisfies(error -> { + 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/StdioMcpClientInitializationFailureTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java index c36540bb8..5658f18e2 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/StdioMcpClientInitializationFailureTests.java @@ -4,12 +4,13 @@ package io.modelcontextprotocol.client; -import java.io.PrintWriter; -import java.io.StringWriter; 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; @@ -41,21 +42,50 @@ void initializeShouldFailWithProcessExitInsteadOfRequestTimeout() { .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(); - String stackTrace = stackTraceOf(failure); assertThat(elapsedMillis).isLessThan(requestTimeout.toMillis()); - assertThat(stackTrace).contains("MCP server process exited", "with code 127") - .doesNotContain("TimeoutException"); + 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() { @@ -63,10 +93,15 @@ private String javaExecutable() { return Path.of(System.getProperty("java.home"), "bin", executable).toString(); } - private String stackTraceOf(Throwable failure) { - StringWriter writer = new StringWriter(); - failure.printStackTrace(new PrintWriter(writer)); - return writer.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); + } + +} From f4e19aada0d4da63e33eac5c066d2f8ebaee1caa Mon Sep 17 00:00:00 2001 From: dragonfsky Date: Wed, 29 Jul 2026 14:45:33 +0800 Subject: [PATCH 3/4] fix: share terminal initialization outcome --- .../client/LifecycleInitializer.java | 89 +++++++++++-------- ...nitializerPostInitializationHookTests.java | 36 ++++++++ .../client/LifecycleInitializerTests.java | 56 +++++++++++- 3 files changed, 140 insertions(+), 41 deletions(-) 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 663ace7a2..6ee442110 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java @@ -213,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) { @@ -222,7 +222,7 @@ 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() { @@ -230,6 +230,10 @@ private void 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); @@ -304,8 +308,20 @@ public Mono withInitialization(String actionName, Function 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(newInit.await(), initializationWork); + } + else { + initializationJob = previous.await(); + } return initializationJob.map(initializeResult -> this.initializationRef.get()) .timeout(this.initializationTimeout) @@ -322,46 +338,45 @@ public Mono withInitialization(String actionName, Function doInitialize(DefaultInitialization initialization, Function> postInitOperation, ContextView ctx) { - initialization.setMcpClientSession(this.sessionSupplier.apply(ctx)); + return Mono.defer(() -> { + initialization.setMcpClientSession(this.sessionSupplier.apply(ctx)); - McpClientSession mcpClientSession = initialization.mcpSession(); - Throwable terminal = this.terminalFailure.get(); - if (terminal != null) { - mcpClientSession.terminate(terminal); - return Mono.error(terminal); - } + McpClientSession mcpClientSession = initialization.mcpSession(); + Throwable terminal = this.terminalFailure.get(); + if (terminal != null) { + mcpClientSession.terminate(terminal); + return Mono.error(terminal); + } - String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1); + String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1); - McpSchema.InitializeRequest initializeRequest = McpSchema.InitializeRequest - .builder(latestVersion, this.clientCapabilities, this.clientInfo) - .build(); + 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); + 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 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()); - } + 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).onErrorResume(ex -> { - initialization.error(ex); - return Mono.error(ex); + 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); }); } 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 12989e2f2..28f68968c 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; @@ -305,15 +307,19 @@ void shouldHandleOtherExceptions() { } @Test - void shouldTerminateInProgressInitializationOnTransportTermination() { + void shouldTerminateInProgressInitializationOnTransportTermination() throws Exception { var cause = new McpTransportTerminatedException("Transport terminated"); when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenReturn(Mono.never()); - var subscription = initializer.withInitialization("test", init -> Mono.just(init.initializeResult())) - .subscribe(); + 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(); @@ -323,7 +329,6 @@ void shouldTerminateInProgressInitializationOnTransportTermination() { .verify(); verify(mockSessionSupplier, times(1)).apply(any(ContextView.class)); - subscription.dispose(); } @Test @@ -342,6 +347,49 @@ void shouldApplyTerminationThatArrivesBeforeSessionRegistration() { 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 shouldIgnoreGenericTransportExceptionDuringInitialization() { var cause = new McpTransportException("Transport closed"); From bc6b03e929920e9648c102c8e169b85ed2cd041b Mon Sep 17 00:00:00 2001 From: dragonfsky Date: Wed, 29 Jul 2026 20:06:12 +0800 Subject: [PATCH 4/4] fix: close terminal initialization races --- .../client/LifecycleInitializer.java | 16 ++++++-- .../client/LifecycleInitializerTests.java | 37 +++++++++++++++++++ 2 files changed, 49 insertions(+), 4 deletions(-) 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 6ee442110..f5c65faf2 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java @@ -306,6 +306,14 @@ public Mono withInitialization(String actionName, Function initializationJob; @@ -317,10 +325,10 @@ public Mono withInitialization(String actionName, Function this.initializationRef.get()) @@ -376,8 +384,8 @@ private Mono doInitialize(DefaultInitialization init }).flatMap(initializeResult -> { initialization.cacheResult(initializeResult); return postInitOperation.apply(initialization).thenReturn(initializeResult); - }).doOnNext(initialization::complete).doOnError(initialization::error); - }); + }); + }).doOnNext(initialization::complete).doOnError(initialization::error); } /** 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 28f68968c..407be06df 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java @@ -390,6 +390,43 @@ void shouldTerminateWinnerAndJoinerWhenTerminationArrivesBeforeSessionRegistrati 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");