Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -324,68 +324,79 @@ protected void doGet(HttpServletRequest request, HttpServletResponse response)
HttpServletStreamableMcpSessionTransport sessionTransport = new HttpServletStreamableMcpSessionTransport(
sessionId, asyncContext, response.getWriter());

// Check if this is a replay request
if (request.getHeader(HttpHeaders.LAST_EVENT_ID) != null) {
String lastId = request.getHeader(HttpHeaders.LAST_EVENT_ID);
// Replay the messages the client missed while its stream was broken
String lastEventId = request.getHeader(HttpHeaders.LAST_EVENT_ID);
if (lastEventId != null
&& !this.tryReplayMissedMessages(session, lastEventId, sessionTransport, transportContext)) {
// The replay failed and already closed the transport
return;
}

try {
session.replay(lastId)
.contextWrite(ctx -> ctx.put(McpTransportContext.KEY, transportContext))
.toIterable()
.forEach(message -> {
try {
sessionTransport.sendMessage(message)
.contextWrite(ctx -> ctx.put(McpTransportContext.KEY, transportContext))
.block();
}
catch (Exception e) {
logger.error("Failed to replay message: {}", e.getMessage());
asyncContext.complete();
}
});
}
catch (Exception e) {
logger.error("Failed to replay messages: {}", e.getMessage());
asyncContext.complete();
// Establish the listening stream. Resumed streams are registered too, so
// that the session keeps delivering messages to the reconnected client and
// the async context is completed once the client goes away.
McpStreamableServerSession.McpStreamableServerSessionStream listeningStream = session
.listeningStream(sessionTransport);

asyncContext.addListener(new jakarta.servlet.AsyncListener() {
@Override
public void onComplete(jakarta.servlet.AsyncEvent event) throws IOException {
logger.debug("SSE connection completed for session: {}", sessionId);
listeningStream.close();
}
}
else {
// Establish new listening stream
McpStreamableServerSession.McpStreamableServerSessionStream listeningStream = session
.listeningStream(sessionTransport);

asyncContext.addListener(new jakarta.servlet.AsyncListener() {
@Override
public void onComplete(jakarta.servlet.AsyncEvent event) throws IOException {
logger.debug("SSE connection completed for session: {}", sessionId);
listeningStream.close();
}

@Override
public void onTimeout(jakarta.servlet.AsyncEvent event) throws IOException {
logger.debug("SSE connection timed out for session: {}", sessionId);
listeningStream.close();
}
@Override
public void onTimeout(jakarta.servlet.AsyncEvent event) throws IOException {
logger.debug("SSE connection timed out for session: {}", sessionId);
listeningStream.close();
}

@Override
public void onError(jakarta.servlet.AsyncEvent event) throws IOException {
logger.debug("SSE connection error for session: {}", sessionId);
listeningStream.close();
}
@Override
public void onError(jakarta.servlet.AsyncEvent event) throws IOException {
logger.debug("SSE connection error for session: {}", sessionId);
listeningStream.close();
}

@Override
public void onStartAsync(jakarta.servlet.AsyncEvent event) throws IOException {
// No action needed
}
});
}
@Override
public void onStartAsync(jakarta.servlet.AsyncEvent event) throws IOException {
// No action needed
}
});
}
catch (Exception e) {
logger.error("Failed to handle GET request for session {}: {}", sessionId, e.getMessage());
response.sendError(HttpServletResponse.SC_INTERNAL_SERVER_ERROR);
}
}

/**
* Replays the messages the client missed while its SSE stream was broken.
* @param session the session the client is resuming
* @param lastEventId the ID of the last event received by the client
* @param sessionTransport the transport of the resumed SSE stream
* @param transportContext the context extracted from the request
* @return {@code true} if the replay completed, {@code false} if it failed, in which
* case the transport has been closed
*/
private boolean tryReplayMissedMessages(McpStreamableServerSession session, String lastEventId,
McpStreamableServerTransport sessionTransport, McpTransportContext transportContext) {
try {
for (McpSchema.JSONRPCMessage message : session.replay(lastEventId)
.contextWrite(ctx -> ctx.put(McpTransportContext.KEY, transportContext))
.toIterable()) {
sessionTransport.sendMessage(message)
.contextWrite(ctx -> ctx.put(McpTransportContext.KEY, transportContext))
.block();
}
return true;
}
catch (Exception e) {
logger.error("Failed to replay messages for session {}: {}", session.getId(), e.getMessage());
sessionTransport.close();
return false;
}
}

/**
* Handles POST requests for incoming JSON-RPC messages from clients.
* @param request The HTTP servlet request containing the JSON-RPC message
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -179,14 +179,20 @@ public Mono<Void> delete() {
}

/**
* Create a listening stream (the generic HTTP GET request without Last-Event-ID
* header).
* Create a listening stream (the generic HTTP GET request, with or without a
* Last-Event-ID header). A session addresses a single listening stream at a time, so
* the stream being replaced, if any, is closed: no message would ever be sent to it
* again, and leaving it open would leak the underlying connection.
* @param transport The dedicated SSE transport stream
* @return a stream representation
*/
public McpStreamableServerSessionStream listeningStream(McpStreamableServerTransport transport) {
McpStreamableServerSessionStream listeningStream = new McpStreamableServerSessionStream(transport);
this.listeningStreamRef.set(listeningStream);
McpLoggableSession replaced = this.listeningStreamRef.getAndSet(listeningStream);
if (replaced instanceof McpStreamableServerSessionStream replacedStream) {
logger.debug("Closing the listening stream replaced in session {}", this.id);
replacedStream.close();
}
return listeningStream;
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,10 @@
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.Map;
import java.util.Queue;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentLinkedQueue;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.concurrent.atomic.AtomicReference;
import java.util.function.Function;
import java.util.stream.Stream;
Expand All @@ -24,6 +28,7 @@
import io.modelcontextprotocol.server.McpServer.SyncSpecification;
import io.modelcontextprotocol.server.transport.HttpServletStreamableServerTransportProvider;
import io.modelcontextprotocol.server.transport.TomcatTestUtil;
import io.modelcontextprotocol.spec.HttpHeaders;
import io.modelcontextprotocol.spec.McpSchema;
import jakarta.servlet.http.HttpServletRequest;
import jakarta.servlet.http.HttpServletResponse;
Expand Down Expand Up @@ -218,4 +223,100 @@ public void cancel() {
assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_REQUEST_ENTITY_TOO_LARGE);
}

@Test
void resumedStreamReceivesServerNotifications() throws Exception {
prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0").build();
var httpClient = HttpClient.newHttpClient();

var sessionId = initializeSession(httpClient);

// Resume the stream the way a client does once its SSE connection broke. The
// resumed stream must become the session listening stream, otherwise the
// reconnected client never receives anything again.
var stream = openListeningStream(httpClient, sessionId, sessionId + "_0");

awaitStreamOpen(stream);
awaitNotification(stream.events());
}

@Test
void replacedListeningStreamIsClosed() throws Exception {
prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0").build();
var httpClient = HttpClient.newHttpClient();

var sessionId = initializeSession(httpClient);

var firstStream = openListeningStream(httpClient, sessionId, null);
awaitStreamOpen(firstStream);
awaitNotification(firstStream.events());

// stream keeps receiving pings, so we just ensure we've removed the notification
firstStream.events().clear();
assertThat(firstStream.events()).noneMatch(line -> line.contains("notifications/resources/list_changed"));

// Resuming installs a new listening stream. The session can no longer
// address the first one, so it must not be left open.
var secondStream = openListeningStream(httpClient, sessionId, sessionId + "_0");
assertThat(firstStream.streamFuture()).succeedsWithin(Duration.ofSeconds(5));
awaitStreamOpen(secondStream);
await().atMost(Duration.ofSeconds(5)).untilAsserted(() -> {
mcpServerTransportProvider.notifyClients(McpSchema.METHOD_NOTIFICATION_RESOURCES_LIST_CHANGED, null)
.block();
assertThat(secondStream.events()).anyMatch(line -> line.contains("notifications/resources/list_changed"));
assertThat(firstStream.events()).noneMatch(line -> line.contains("notifications/resources/list_changed"));
});
}

private String initializeSession(HttpClient httpClient) throws Exception {
var initialize = HttpRequest.newBuilder()
.uri(URI.create("http://localhost:" + PORT + MESSAGE_ENDPOINT))
.header("Content-Type", "application/json")
.header("Accept", "text/event-stream, application/json")
.POST(HttpRequest.BodyPublishers.ofString("""
{"jsonrpc":"2.0","id":"init","method":"initialize","params":{
"protocolVersion":"2025-06-18","capabilities":{},
"clientInfo":{"name":"test-client","version":"1.0.0"}}}"""))
.build();

var response = httpClient.send(initialize, HttpResponse.BodyHandlers.ofString());
assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_OK);
return response.headers().firstValue(HttpHeaders.MCP_SESSION_ID).orElseThrow();
}

/**
* Opens an SSE listening stream with a GET request, collecting the received lines.
* @return a future completing once the server closes the stream
*/
private StreamResponse openListeningStream(HttpClient httpClient, String sessionId, String lastEventId) {
var get = HttpRequest.newBuilder()
.uri(URI.create("http://localhost:" + PORT + MESSAGE_ENDPOINT))
.header("Accept", "text/event-stream")
.header(HttpHeaders.MCP_SESSION_ID, sessionId);
if (lastEventId != null) {
get.header(HttpHeaders.LAST_EVENT_ID, lastEventId);
}
Queue<String> events = new ConcurrentLinkedQueue<>();
var eventsReceived = new AtomicBoolean(false);
var clientFuture = httpClient.sendAsync(get.GET().build(), HttpResponse.BodyHandlers.ofLines())
.thenAccept(response -> {
eventsReceived.set(true);
response.body().forEach(events::add);
});
return new StreamResponse(clientFuture, eventsReceived, events);
}

private void awaitNotification(Queue<String> events) {
mcpServerTransportProvider.notifyClients(McpSchema.METHOD_NOTIFICATION_RESOURCES_LIST_CHANGED, null).block();
await().atMost(Duration.ofSeconds(5)).pollDelay(Duration.ofMillis(100)).untilAsserted(() -> {
assertThat(events).anyMatch(line -> line.contains("notifications/resources/list_changed"));
});
}

private static void awaitStreamOpen(StreamResponse stream) {
await().atMost(Duration.ofSeconds(5)).untilAsserted(() -> assertThat(stream.isOpen()).isTrue());
}

record StreamResponse(CompletableFuture<Void> streamFuture, AtomicBoolean isOpen, Queue<String> events) {
}

}
Loading