Skip to content

Commit bc6b03e

Browse files
committed
fix: close terminal initialization races
1 parent f4e19aa commit bc6b03e

2 files changed

Lines changed: 49 additions & 4 deletions

File tree

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

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -306,6 +306,14 @@ public <T> Mono<T> withInitialization(String actionName, Function<Initialization
306306
DefaultInitialization previous = this.initializationRef.compareAndExchange(null, newInit);
307307

308308
boolean needsToInitialize = previous == null;
309+
DefaultInitialization activeInitialization = needsToInitialize ? newInit : previous;
310+
// Complete the handoff with handleException(): if termination won before
311+
// registration, this check publishes it; if registration won first, the
312+
// exception handler observes the active initialization.
313+
Throwable terminalAfterRegistration = this.terminalFailure.get();
314+
if (terminalAfterRegistration != null) {
315+
activeInitialization.terminate(terminalAfterRegistration);
316+
}
309317
logger.debug(needsToInitialize ? "Initialization process started" : "Joining previous initialization");
310318

311319
Mono<McpSchema.InitializeResult> initializationJob;
@@ -317,10 +325,10 @@ public <T> Mono<T> withInitialization(String actionName, Function<Initialization
317325
.doInitialize(newInit, this.postInitializationHook, ctx)
318326
.onErrorComplete()
319327
.then(Mono.never());
320-
initializationJob = Mono.firstWithSignal(newInit.await(), initializationWork);
328+
initializationJob = Mono.firstWithSignal(activeInitialization.await(), initializationWork);
321329
}
322330
else {
323-
initializationJob = previous.await();
331+
initializationJob = activeInitialization.await();
324332
}
325333

326334
return initializationJob.map(initializeResult -> this.initializationRef.get())
@@ -376,8 +384,8 @@ private Mono<McpSchema.InitializeResult> doInitialize(DefaultInitialization init
376384
}).flatMap(initializeResult -> {
377385
initialization.cacheResult(initializeResult);
378386
return postInitOperation.apply(initialization).thenReturn(initializeResult);
379-
}).doOnNext(initialization::complete).doOnError(initialization::error);
380-
});
387+
});
388+
}).doOnNext(initialization::complete).doOnError(initialization::error);
381389
}
382390

383391
/**

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

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -390,6 +390,43 @@ void shouldTerminateWinnerAndJoinerWhenTerminationArrivesBeforeSessionRegistrati
390390
verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any());
391391
}
392392

393+
@Test
394+
void shouldShareSynchronousSessionSupplierFailureBetweenWinnerAndJoiner() throws Exception {
395+
var cause = new IllegalStateException("Session supplier failed");
396+
var supplierEntered = new CountDownLatch(1);
397+
var releaseSupplier = new CountDownLatch(1);
398+
399+
when(mockSessionSupplier.apply(any(ContextView.class))).thenAnswer(invocation -> {
400+
supplierEntered.countDown();
401+
if (!releaseSupplier.await(5, TimeUnit.SECONDS)) {
402+
throw new IllegalStateException("Timed out waiting to release session supplier");
403+
}
404+
throw cause;
405+
});
406+
407+
var winner = initializer.withInitialization("winner", init -> Mono.just(init.initializeResult()))
408+
.subscribeOn(Schedulers.boundedElastic())
409+
.materialize()
410+
.toFuture();
411+
assertThat(supplierEntered.await(1, TimeUnit.SECONDS)).isTrue();
412+
413+
var joiner = initializer.withInitialization("joiner", init -> Mono.just(init.initializeResult()))
414+
.materialize()
415+
.toFuture();
416+
417+
releaseSupplier.countDown();
418+
419+
var winnerSignal = winner.get(1, TimeUnit.SECONDS);
420+
var joinerSignal = joiner.get(1, TimeUnit.SECONDS);
421+
422+
assertThat(winnerSignal.isOnError()).isTrue();
423+
assertThat(winnerSignal.getThrowable()).hasCause(cause);
424+
assertThat(joinerSignal.isOnError()).isTrue();
425+
assertThat(joinerSignal.getThrowable()).hasCause(cause);
426+
verify(mockSessionSupplier, times(1)).apply(any(ContextView.class));
427+
verify(mockClientSession, never()).sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any());
428+
}
429+
393430
@Test
394431
void shouldIgnoreGenericTransportExceptionDuringInitialization() {
395432
var cause = new McpTransportException("Transport closed");

0 commit comments

Comments
 (0)