Skip to content

Commit f4e19aa

Browse files
committed
fix: share terminal initialization outcome
1 parent 71fdb15 commit f4e19aa

3 files changed

Lines changed: 140 additions & 41 deletions

File tree

mcp-core/src/main/java/io/modelcontextprotocol/client/LifecycleInitializer.java

Lines changed: 52 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -213,7 +213,7 @@ private Mono<McpSchema.InitializeResult> await() {
213213

214214
private void complete(McpSchema.InitializeResult initializeResult) {
215215
// inform all the subscribers waiting for the initialization
216-
this.initSink.emitValue(initializeResult, Sinks.EmitFailureHandler.FAIL_FAST);
216+
this.initSink.tryEmitValue(initializeResult);
217217
}
218218

219219
private void cacheResult(McpSchema.InitializeResult initializeResult) {
@@ -222,14 +222,18 @@ private void cacheResult(McpSchema.InitializeResult initializeResult) {
222222
}
223223

224224
private void error(Throwable t) {
225-
this.initSink.emitError(t, Sinks.EmitFailureHandler.FAIL_FAST);
225+
this.initSink.tryEmitError(t);
226226
}
227227

228228
private void close() {
229229
this.mcpSession().close();
230230
}
231231

232232
private void terminate(Throwable cause) {
233+
// Initialization has a single shared outcome. Publish the terminal failure
234+
// even when the session has not been installed yet so that both the owner
235+
// and all concurrent joiners observe it immediately.
236+
this.initSink.tryEmitError(cause);
233237
McpClientSession mcpClientSession = this.mcpSession();
234238
if (mcpClientSession != null) {
235239
mcpClientSession.terminate(cause);
@@ -304,8 +308,20 @@ public <T> Mono<T> withInitialization(String actionName, Function<Initialization
304308
boolean needsToInitialize = previous == null;
305309
logger.debug(needsToInitialize ? "Initialization process started" : "Joining previous initialization");
306310

307-
Mono<McpSchema.InitializeResult> initializationJob = needsToInitialize
308-
? this.doInitialize(newInit, this.postInitializationHook, ctx) : previous.await();
311+
Mono<McpSchema.InitializeResult> initializationJob;
312+
if (needsToInitialize) {
313+
// The work branch only publishes into the shared sink. Keeping it from
314+
// winning directly makes the owner and all joiners consume the same
315+
// first terminal signal.
316+
Mono<McpSchema.InitializeResult> initializationWork = this
317+
.doInitialize(newInit, this.postInitializationHook, ctx)
318+
.onErrorComplete()
319+
.then(Mono.never());
320+
initializationJob = Mono.firstWithSignal(newInit.await(), initializationWork);
321+
}
322+
else {
323+
initializationJob = previous.await();
324+
}
309325

310326
return initializationJob.map(initializeResult -> this.initializationRef.get())
311327
.timeout(this.initializationTimeout)
@@ -322,46 +338,45 @@ public <T> Mono<T> withInitialization(String actionName, Function<Initialization
322338
private Mono<McpSchema.InitializeResult> doInitialize(DefaultInitialization initialization,
323339
Function<Initialization, Mono<Void>> postInitOperation, ContextView ctx) {
324340

325-
initialization.setMcpClientSession(this.sessionSupplier.apply(ctx));
341+
return Mono.defer(() -> {
342+
initialization.setMcpClientSession(this.sessionSupplier.apply(ctx));
326343

327-
McpClientSession mcpClientSession = initialization.mcpSession();
328-
Throwable terminal = this.terminalFailure.get();
329-
if (terminal != null) {
330-
mcpClientSession.terminate(terminal);
331-
return Mono.error(terminal);
332-
}
344+
McpClientSession mcpClientSession = initialization.mcpSession();
345+
Throwable terminal = this.terminalFailure.get();
346+
if (terminal != null) {
347+
mcpClientSession.terminate(terminal);
348+
return Mono.error(terminal);
349+
}
333350

334-
String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1);
351+
String latestVersion = this.protocolVersions.get(this.protocolVersions.size() - 1);
335352

336-
McpSchema.InitializeRequest initializeRequest = McpSchema.InitializeRequest
337-
.builder(latestVersion, this.clientCapabilities, this.clientInfo)
338-
.build();
353+
McpSchema.InitializeRequest initializeRequest = McpSchema.InitializeRequest
354+
.builder(latestVersion, this.clientCapabilities, this.clientInfo)
355+
.build();
339356

340-
Mono<McpSchema.InitializeResult> result = mcpClientSession.sendRequest(McpSchema.METHOD_INITIALIZE,
341-
initializeRequest, McpAsyncClient.INITIALIZE_RESULT_TYPE_REF);
357+
Mono<McpSchema.InitializeResult> result = mcpClientSession.sendRequest(McpSchema.METHOD_INITIALIZE,
358+
initializeRequest, McpAsyncClient.INITIALIZE_RESULT_TYPE_REF);
342359

343-
return result.flatMap(initializeResult -> {
344-
logger.info("Server response with Protocol: {}, Capabilities: {}, Info: {} and Instructions {}",
345-
initializeResult.protocolVersion(), initializeResult.capabilities(), initializeResult.serverInfo(),
346-
initializeResult.instructions());
360+
return result.flatMap(initializeResult -> {
361+
logger.info("Server response with Protocol: {}, Capabilities: {}, Info: {} and Instructions {}",
362+
initializeResult.protocolVersion(), initializeResult.capabilities(),
363+
initializeResult.serverInfo(), initializeResult.instructions());
347364

348-
if (!this.protocolVersions.contains(initializeResult.protocolVersion())) {
349-
return Mono.error(McpError.builder(-32602)
350-
.message("Unsupported protocol version")
351-
.data("Unsupported protocol version from the server: " + initializeResult.protocolVersion())
352-
.build());
353-
}
365+
if (!this.protocolVersions.contains(initializeResult.protocolVersion())) {
366+
return Mono.error(McpError.builder(-32602)
367+
.message("Unsupported protocol version")
368+
.data("Unsupported protocol version from the server: " + initializeResult.protocolVersion())
369+
.build());
370+
}
354371

355-
return mcpClientSession.sendNotification(McpSchema.METHOD_NOTIFICATION_INITIALIZED, null)
356-
.contextWrite(
357-
c -> c.put(McpAsyncClient.NEGOTIATED_PROTOCOL_VERSION, initializeResult.protocolVersion()))
358-
.thenReturn(initializeResult);
359-
}).flatMap(initializeResult -> {
360-
initialization.cacheResult(initializeResult);
361-
return postInitOperation.apply(initialization).thenReturn(initializeResult);
362-
}).doOnNext(initialization::complete).onErrorResume(ex -> {
363-
initialization.error(ex);
364-
return Mono.error(ex);
372+
return mcpClientSession.sendNotification(McpSchema.METHOD_NOTIFICATION_INITIALIZED, null)
373+
.contextWrite(
374+
c -> c.put(McpAsyncClient.NEGOTIATED_PROTOCOL_VERSION, initializeResult.protocolVersion()))
375+
.thenReturn(initializeResult);
376+
}).flatMap(initializeResult -> {
377+
initialization.cacheResult(initializeResult);
378+
return postInitOperation.apply(initialization).thenReturn(initializeResult);
379+
}).doOnNext(initialization::complete).doOnError(initialization::error);
365380
});
366381
}
367382

mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerPostInitializationHookTests.java

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66

77
import java.time.Duration;
88
import java.util.List;
9+
import java.util.concurrent.CountDownLatch;
10+
import java.util.concurrent.TimeUnit;
911
import java.util.concurrent.atomic.AtomicInteger;
1012
import java.util.concurrent.atomic.AtomicReference;
1113
import java.util.function.Function;
@@ -14,11 +16,13 @@
1416
import io.modelcontextprotocol.spec.McpClientSession;
1517
import io.modelcontextprotocol.spec.McpSchema;
1618
import io.modelcontextprotocol.spec.McpTransportSessionNotFoundException;
19+
import io.modelcontextprotocol.spec.McpTransportTerminatedException;
1720
import org.junit.jupiter.api.BeforeEach;
1821
import org.junit.jupiter.api.Test;
1922
import org.mockito.Mock;
2023
import org.mockito.MockitoAnnotations;
2124
import reactor.core.publisher.Mono;
25+
import reactor.core.publisher.Sinks;
2226
import reactor.core.scheduler.Schedulers;
2327
import reactor.test.StepVerifier;
2428
import reactor.util.context.ContextView;
@@ -164,6 +168,38 @@ void shouldFailInitializationWhenPostInitializationHookFails() {
164168
verify(mockPostInitializationHook, times(1)).apply(any(Initialization.class));
165169
}
166170

171+
@Test
172+
void shouldFailInitializationWhenTransportTerminatesDuringPostInitializationHook() throws Exception {
173+
var cause = new McpTransportTerminatedException("Transport terminated during post-initialization hook");
174+
var hookEntered = new CountDownLatch(1);
175+
var hookGate = Sinks.<Void>one();
176+
177+
when(mockPostInitializationHook.apply(any(Initialization.class))).thenAnswer(invocation -> Mono.defer(() -> {
178+
hookEntered.countDown();
179+
return hookGate.asMono();
180+
}));
181+
182+
var initialization = initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))
183+
.materialize()
184+
.toFuture();
185+
assertThat(hookEntered.await(1, TimeUnit.SECONDS)).isTrue();
186+
187+
initializer.handleException(cause);
188+
189+
try {
190+
var signal = initialization.get(1, TimeUnit.SECONDS);
191+
assertThat(signal.isOnError()).isTrue();
192+
assertThat(signal.getThrowable()).hasCause(cause);
193+
}
194+
finally {
195+
hookGate.tryEmitEmpty();
196+
}
197+
198+
assertThat(initializer.isInitialized()).isFalse();
199+
assertThat(initializer.currentInitializationResult()).isNull();
200+
verify(mockClientSession).terminate(cause);
201+
}
202+
167203
@Test
168204
void shouldNotInvokePostInitializationHookWhenInitializationFails() {
169205
when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any()))

mcp-core/src/test/java/io/modelcontextprotocol/client/LifecycleInitializerTests.java

Lines changed: 52 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66

77
import java.time.Duration;
88
import java.util.List;
9+
import java.util.concurrent.CountDownLatch;
10+
import java.util.concurrent.TimeUnit;
911
import java.util.concurrent.atomic.AtomicInteger;
1012
import java.util.concurrent.atomic.AtomicReference;
1113
import java.util.function.Function;
@@ -305,15 +307,19 @@ void shouldHandleOtherExceptions() {
305307
}
306308

307309
@Test
308-
void shouldTerminateInProgressInitializationOnTransportTermination() {
310+
void shouldTerminateInProgressInitializationOnTransportTermination() throws Exception {
309311
var cause = new McpTransportTerminatedException("Transport terminated");
310312
when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenReturn(Mono.never());
311313

312-
var subscription = initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))
313-
.subscribe();
314+
var initialization = initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))
315+
.materialize()
316+
.toFuture();
314317

315318
initializer.handleException(cause);
316319

320+
var signal = initialization.get(1, TimeUnit.SECONDS);
321+
assertThat(signal.isOnError()).isTrue();
322+
assertThat(signal.getThrowable()).hasCause(cause);
317323
verify(mockClientSession).terminate(cause);
318324
assertThat(initializer.isInitialized()).isFalse();
319325
assertThat(initializer.currentInitializationResult()).isNull();
@@ -323,7 +329,6 @@ void shouldTerminateInProgressInitializationOnTransportTermination() {
323329
.verify();
324330

325331
verify(mockSessionSupplier, times(1)).apply(any(ContextView.class));
326-
subscription.dispose();
327332
}
328333

329334
@Test
@@ -342,6 +347,49 @@ void shouldApplyTerminationThatArrivesBeforeSessionRegistration() {
342347
verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any());
343348
}
344349

350+
@Test
351+
void shouldTerminateWinnerAndJoinerWhenTerminationArrivesBeforeSessionRegistration() throws Exception {
352+
var cause = new McpTransportTerminatedException("Transport terminated before session registration");
353+
var supplierEntered = new CountDownLatch(1);
354+
var releaseSupplier = new CountDownLatch(1);
355+
356+
when(mockSessionSupplier.apply(any(ContextView.class))).thenAnswer(invocation -> {
357+
supplierEntered.countDown();
358+
if (!releaseSupplier.await(5, TimeUnit.SECONDS)) {
359+
throw new IllegalStateException("Timed out waiting to release session supplier");
360+
}
361+
return mockClientSession;
362+
});
363+
364+
var winner = initializer.withInitialization("winner", init -> Mono.just(init.initializeResult()))
365+
.subscribeOn(Schedulers.boundedElastic())
366+
.materialize()
367+
.toFuture();
368+
assertThat(supplierEntered.await(1, TimeUnit.SECONDS)).isTrue();
369+
370+
var joiner = initializer.withInitialization("joiner", init -> Mono.just(init.initializeResult()))
371+
.materialize()
372+
.toFuture();
373+
374+
initializer.handleException(cause);
375+
376+
try {
377+
var winnerSignal = winner.get(1, TimeUnit.SECONDS);
378+
var joinerSignal = joiner.get(1, TimeUnit.SECONDS);
379+
380+
assertThat(winnerSignal.isOnError()).isTrue();
381+
assertThat(winnerSignal.getThrowable()).hasCause(cause);
382+
assertThat(joinerSignal.isOnError()).isTrue();
383+
assertThat(joinerSignal.getThrowable()).hasCause(cause);
384+
}
385+
finally {
386+
releaseSupplier.countDown();
387+
}
388+
389+
verify(mockSessionSupplier, times(1)).apply(any(ContextView.class));
390+
verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any());
391+
}
392+
345393
@Test
346394
void shouldIgnoreGenericTransportExceptionDuringInitialization() {
347395
var cause = new McpTransportException("Transport closed");

0 commit comments

Comments
 (0)