diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java index 311390f3c..cacb30522 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java @@ -324,61 +324,44 @@ 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()); @@ -386,6 +369,34 @@ public void onStartAsync(jakarta.servlet.AsyncEvent event) throws IOException { } } + /** + * 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 diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpStreamableServerSession.java b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpStreamableServerSession.java index e7fac7b0d..7f892df09 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpStreamableServerSession.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpStreamableServerSession.java @@ -179,14 +179,20 @@ public Mono 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; } diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java index 8228ad3b1..0a918b6df 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java @@ -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; @@ -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; @@ -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 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 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 streamFuture, AtomicBoolean isOpen, Queue events) { + } + }