Skip to content

Commit 9796c7c

Browse files
committed
Address connect hangs in legacy SSE connect and streamline code paths
Signed-off-by: Dariusz Jędrzejczyk <dariusz.jedrzejczyk@broadcom.com>
1 parent 46511cd commit 9796c7c

2 files changed

Lines changed: 150 additions & 27 deletions

File tree

‎mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransport.java‎

Lines changed: 37 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
import java.net.http.HttpResponse;
1212
import java.time.Duration;
1313
import java.util.List;
14-
import java.util.concurrent.atomic.AtomicBoolean;
14+
import java.util.Optional;
1515
import java.util.concurrent.atomic.AtomicReference;
1616
import java.util.function.Consumer;
1717
import java.util.function.Function;
@@ -389,14 +389,6 @@ public Mono<Void> connect(Function<Mono<JSONRPCMessage>, Mono<JSONRPCMessage>> h
389389
var transportContext = ctx.getOrDefault(McpTransportContext.KEY, McpTransportContext.EMPTY);
390390
return Mono.from(this.httpRequestCustomizer.customize(builder, "GET", uri, null, transportContext));
391391
}).flatMap(requestBuilder -> Mono.create(sink -> {
392-
// Once connect() has completed, a later failure can no longer be reported
393-
// through its sink: signalling it there would only have Reactor drop it.
394-
AtomicBoolean connected = new AtomicBoolean();
395-
Runnable markConnected = () -> {
396-
if (connected.compareAndSet(false, true)) {
397-
sink.success();
398-
}
399-
};
400392
Disposable connection = ResponseBodyHandlers.sendAsync(this.httpClient, requestBuilder.build())
401393
.flatMapMany(response -> {
402394
if (isClosing) {
@@ -417,7 +409,10 @@ public Mono<Void> connect(Function<Mono<JSONRPCMessage>, Mono<JSONRPCMessage>> h
417409
"Failed to connect to SSE stream: " + statusCode);
418410
}
419411
})
420-
.flatMap(sseEvent -> {
412+
// Every successfully processed event yields exactly one element, empty
413+
// when it carries no message, so that the first one can mark the
414+
// connection as established.
415+
.<Optional<JSONRPCMessage>>handle((sseEvent, events) -> {
421416
try {
422417
if (ENDPOINT_EVENT_TYPE.equals(sseEvent.event())) {
423418
String messageEndpointUri = sseEvent.data();
@@ -426,46 +421,61 @@ public Mono<Void> connect(Function<Mono<JSONRPCMessage>, Mono<JSONRPCMessage>> h
426421
}
427422
catch (InvalidSseMessageEndpointException e) {
428423
this.messageEndpointSink.tryEmitError(e);
429-
return Flux.error(e);
424+
events.error(e);
425+
return;
430426
}
431427
if (this.messageEndpointSink.tryEmitValue(messageEndpointUri).isSuccess()) {
432-
markConnected.run();
433-
return Flux.empty(); // No further processing needed
428+
events.next(Optional.empty());
429+
}
430+
else {
431+
events.error(new McpTransportException("Failed to handle SSE endpoint event"));
434432
}
435-
return Flux.error(new McpTransportException("Failed to handle SSE endpoint event"));
436433
}
437434
else if (MESSAGE_EVENT_TYPE.equals(sseEvent.event())) {
438435
String data = sseEvent.data();
439436
if (data == null || data.isBlank()) {
440437
logger.debug("Skipping SSE event with empty data (stream primer)");
441-
markConnected.run();
442-
return Flux.<McpSchema.JSONRPCMessage>empty();
438+
events.next(Optional.empty());
439+
}
440+
else {
441+
events.next(Optional.of(McpSchema.deserializeJsonRpcMessage(jsonMapper, data)));
443442
}
444-
JSONRPCMessage message = McpSchema.deserializeJsonRpcMessage(jsonMapper, data);
445-
markConnected.run();
446-
return Flux.just(message);
447443
}
448444
else {
449445
logger.debug("Received unrecognized SSE event type: {}", sseEvent);
450-
markConnected.run();
451-
return Flux.<McpSchema.JSONRPCMessage>empty();
446+
events.next(Optional.empty());
452447
}
453448
}
454449
catch (IOException e) {
455-
return Flux.<McpSchema.JSONRPCMessage>error(
456-
new McpTransportException("Error processing SSE event", e));
450+
events.error(new McpTransportException("Error processing SSE event", e));
451+
}
452+
})
453+
// connect() is resolved by the first signal only: any later failure is
454+
// merely logged below, as connect() has already completed by then.
455+
.switchOnFirst((first, events) -> {
456+
if (first.hasValue()) {
457+
sink.success();
458+
}
459+
else if (first.isOnError()) {
460+
sink.error(first.getThrowable());
457461
}
462+
else if (first.isOnComplete()) {
463+
sink.error(new McpTransportException("SSE stream closed before any event was received"));
464+
}
465+
return events;
458466
})
459-
.flatMap(jsonRpcMessage -> handler.apply(Mono.just(jsonRpcMessage)))
467+
.<JSONRPCMessage>handle((message, messages) -> message.ifPresent(messages::next))
468+
.flatMap(message -> handler.apply(Mono.just(message)))
460469
.onErrorComplete(t -> {
461470
if (!isClosing) {
462471
logger.warn("SSE stream observed an error", t);
463-
if (connected.compareAndSet(false, true)) {
464-
sink.error(t);
465-
}
466472
}
467473
return true;
468474
})
475+
// A closeGracefully() before the first signal cancels the stream:
476+
// complete
477+
// connect() instead of leaving it pending. A no-op once it has resolved.
478+
.doOnCancel(sink::success)
469479
.doFinally(s -> {
470480
Disposable ref = this.sseSubscription.getAndSet(null);
471481
if (ref != null && !ref.isDisposed()) {
Lines changed: 113 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,113 @@
1+
/*
2+
* Copyright 2026-2026 the original author or authors.
3+
*/
4+
5+
package io.modelcontextprotocol.client.transport;
6+
7+
import java.io.IOException;
8+
import java.net.InetSocketAddress;
9+
import java.time.Duration;
10+
import java.util.concurrent.CountDownLatch;
11+
import java.util.concurrent.ExecutorService;
12+
import java.util.concurrent.Executors;
13+
import java.util.concurrent.TimeUnit;
14+
import java.util.function.Function;
15+
16+
import com.sun.net.httpserver.HttpHandler;
17+
import com.sun.net.httpserver.HttpServer;
18+
import io.modelcontextprotocol.spec.McpTransportException;
19+
import io.modelcontextprotocol.spec.json.gson.GsonMcpJsonMapper;
20+
import org.junit.jupiter.api.AfterEach;
21+
import org.junit.jupiter.api.Test;
22+
import reactor.test.StepVerifier;
23+
24+
import static org.assertj.core.api.Assertions.assertThat;
25+
26+
/**
27+
* Verifies that {@link HttpClientSseClientTransport#connect} always resolves, even when
28+
* the SSE stream ends, fails or is closed before its first event.
29+
*/
30+
class HttpClientSseClientTransportConnectTests {
31+
32+
// Only bounds a regression: every test resolves without waiting on it.
33+
private static final Duration TIMEOUT = Duration.ofSeconds(5);
34+
35+
private final ExecutorService executor = Executors.newCachedThreadPool();
36+
37+
private final CountDownLatch releaseResponse = new CountDownLatch(1);
38+
39+
private HttpServer server;
40+
41+
@AfterEach
42+
void tearDown() {
43+
this.releaseResponse.countDown();
44+
if (this.server != null) {
45+
this.server.stop(0);
46+
}
47+
this.executor.shutdownNow();
48+
}
49+
50+
@Test
51+
void connectFailsWhenStreamEndsBeforeAnyEvent() throws IOException {
52+
HttpClientSseClientTransport transport = transport(exchange -> {
53+
exchange.getResponseHeaders().add("Content-Type", "text/event-stream");
54+
exchange.sendResponseHeaders(200, -1);
55+
exchange.close();
56+
});
57+
58+
StepVerifier.create(transport.connect(Function.identity()))
59+
.expectErrorSatisfies(e -> assertThat(e).isInstanceOf(McpTransportException.class)
60+
.hasMessageContaining("before any event"))
61+
.verify(TIMEOUT);
62+
}
63+
64+
@Test
65+
void connectFailsWhenStreamErrorsBeforeAnyEventWhileClosing() throws IOException {
66+
HttpClientSseClientTransport transport = transport(exchange -> {
67+
// The server drops the connection without responding.
68+
throw new IOException("dropped");
69+
});
70+
transport.closeGracefully().block(TIMEOUT);
71+
72+
StepVerifier.create(transport.connect(Function.identity())).expectError().verify(TIMEOUT);
73+
}
74+
75+
@Test
76+
void connectCompletesWhenClosedBeforeAnyEvent() throws IOException {
77+
CountDownLatch streamOpened = new CountDownLatch(1);
78+
HttpClientSseClientTransport transport = transport(exchange -> {
79+
exchange.getResponseHeaders().add("Content-Type", "text/event-stream");
80+
exchange.sendResponseHeaders(200, 0);
81+
exchange.getResponseBody().flush();
82+
streamOpened.countDown();
83+
try {
84+
this.releaseResponse.await();
85+
}
86+
catch (InterruptedException e) {
87+
Thread.currentThread().interrupt();
88+
}
89+
exchange.close();
90+
});
91+
92+
StepVerifier.create(transport.connect(Function.identity())).then(() -> {
93+
try {
94+
assertThat(streamOpened.await(TIMEOUT.toMillis(), TimeUnit.MILLISECONDS)).isTrue();
95+
}
96+
catch (InterruptedException e) {
97+
throw new IllegalStateException(e);
98+
}
99+
transport.closeGracefully().block(TIMEOUT);
100+
}).expectComplete().verify(TIMEOUT);
101+
}
102+
103+
private HttpClientSseClientTransport transport(HttpHandler sseHandler) throws IOException {
104+
this.server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0);
105+
this.server.setExecutor(this.executor);
106+
this.server.createContext("/sse", sseHandler);
107+
this.server.start();
108+
return HttpClientSseClientTransport.builder("http://127.0.0.1:" + this.server.getAddress().getPort())
109+
.jsonMapper(new GsonMcpJsonMapper())
110+
.build();
111+
}
112+
113+
}

0 commit comments

Comments
 (0)