diff --git a/arthas-mcp-server/src/main/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandler.java b/arthas-mcp-server/src/main/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandler.java index d55789adb5..4475b8a49a 100644 --- a/arthas-mcp-server/src/main/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandler.java +++ b/arthas-mcp-server/src/main/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandler.java @@ -197,13 +197,13 @@ protected void handle(ChannelHandlerContext ctx, FullHttpRequest request) throws } /** - * Handles GET requests to establish SSE connections and message replay. + * Handles GET requests to establish SSE listening streams. + *

+ * Resume via {@code last-event-id} is not supported; reject immediately with 404 so + * clients re-initialize (#3118). */ private void handleGetRequest(ChannelHandlerContext ctx, FullHttpRequest request) { - // TODO support last-event-id #3118 - // MCP 客户端在 SSE 断线重连时,可能会带上 last-event-id 尝试做消息回放。 - // Arthas MCP Server 不支持基于 last-event-id 的恢复逻辑:直接返回 404, - // 让客户端触发完整重置并重新走 Initialize 握手申请新的会话。 + // Unsupported resume: reject last-event-id immediately with 404. if (request.headers().get(HttpHeaders.LAST_EVENT_ID) != null) { sendError(ctx, HttpResponseStatus.NOT_FOUND, new McpError("Session not found, please re-initialize")); @@ -259,39 +259,15 @@ private void handleGetRequest(ChannelHandlerContext ctx, FullHttpRequest request NettyStreamableMcpSessionTransport sessionTransport = new NettyStreamableMcpSessionTransport( sessionId, ctx); - // Check if this is a replay request - String lastEventId = request.headers().get(HttpHeaders.LAST_EVENT_ID); - if (lastEventId != null) { - try { - // Replay messages from the last event ID - try { - session.replay(lastEventId).forEach(message -> { - try { - sessionTransport.sendMessage(message).join(); - } catch (Exception e) { - logger.error("Failed to replay message: {}", e.getMessage()); - ctx.close(); - } - }); - } catch (Exception e) { - logger.error("Failed to replay messages: {}", e.getMessage()); - ctx.close(); - } - } catch (Exception e) { - logger.error("Failed to replay messages: {}", e.getMessage()); - ctx.close(); - } - } else { - // Establish new listening stream - McpStreamableServerSession.McpStreamableServerSessionStream listeningStream = session - .listeningStream(sessionTransport); - - // Handle channel closure - ctx.channel().closeFuture().addListener(future -> { - logger.debug("SSE connection closed for session: {}", sessionId); - listeningStream.close(); - }); - } + // Establish new listening stream + McpStreamableServerSession.McpStreamableServerSessionStream listeningStream = session + .listeningStream(sessionTransport); + + // Handle channel closure + ctx.channel().closeFuture().addListener(future -> { + logger.debug("SSE connection closed for session: {}", sessionId); + listeningStream.close(); + }); } catch (Exception e) { logger.error("Failed to handle GET request for session {}: {}", sessionId, e.getMessage()); sendError(ctx, HttpResponseStatus.INTERNAL_SERVER_ERROR, new McpError("Internal server error")); diff --git a/arthas-mcp-server/src/test/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandlerTest.java b/arthas-mcp-server/src/test/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandlerTest.java new file mode 100644 index 0000000000..5199800e1b --- /dev/null +++ b/arthas-mcp-server/src/test/java/com/taobao/arthas/mcp/server/protocol/server/handler/McpStreamableHttpRequestHandlerTest.java @@ -0,0 +1,99 @@ +/* + * Copyright 2024-2024 the original author or authors. + */ + +package com.taobao.arthas.mcp.server.protocol.server.handler; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.taobao.arthas.mcp.server.protocol.spec.HttpHeaders; +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.SimpleChannelInboundHandler; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.http.DefaultFullHttpRequest; +import io.netty.handler.codec.http.FullHttpRequest; +import io.netty.handler.codec.http.FullHttpResponse; +import io.netty.handler.codec.http.HttpHeaderNames; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.util.CharsetUtil; +import io.netty.util.ReferenceCountUtil; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; + +class McpStreamableHttpRequestHandlerTest { + + private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper(); + private static final String MCP_ENDPOINT = "/mcp"; + private static final String TEXT_EVENT_STREAM = "text/event-stream"; + + @Test + void getWithLastEventIdShouldReturn404Immediately() { + McpStreamableHttpRequestHandler handler = newHandler(); + + EmbeddedChannel channel = newChannel(handler); + DefaultFullHttpRequest request = new DefaultFullHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.GET, MCP_ENDPOINT); + request.headers().set(HttpHeaderNames.ACCEPT, TEXT_EVENT_STREAM); + // Guard runs before session lookup; any session id is sufficient. + request.headers().set(HttpHeaders.MCP_SESSION_ID, "any-session"); + // Cherry Studio reconnects with last-event-id after the streamable response closes. + request.headers().set(HttpHeaders.LAST_EVENT_ID, "event-1"); + + channel.writeInbound(request); + + FullHttpResponse response = readOutbound(channel, FullHttpResponse.class); + assertThat(response.status()).isEqualTo(HttpResponseStatus.NOT_FOUND); + assertThat(response.content().toString(CharsetUtil.UTF_8)) + .contains("please re-initialize"); + ReferenceCountUtil.release(response); + + // Must not leave a hanging SSE stream open for the client to time out on. + assertThat((Object) channel.readOutbound()).isNull(); + assertThat(channel.isActive()).isFalse(); + channel.finishAndReleaseAll(); + } + + @Test + void getWithLastEventIdIsCaseInsensitive() { + McpStreamableHttpRequestHandler handler = newHandler(); + + EmbeddedChannel channel = newChannel(handler); + DefaultFullHttpRequest request = new DefaultFullHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.GET, MCP_ENDPOINT); + request.headers().set(HttpHeaderNames.ACCEPT, TEXT_EVENT_STREAM); + request.headers().set(HttpHeaders.MCP_SESSION_ID, "any-session"); + request.headers().set("Last-Event-ID", "event-1"); + + channel.writeInbound(request); + + FullHttpResponse response = readOutbound(channel, FullHttpResponse.class); + assertThat(response.status()).isEqualTo(HttpResponseStatus.NOT_FOUND); + ReferenceCountUtil.release(response); + assertThat(channel.isActive()).isFalse(); + channel.finishAndReleaseAll(); + } + + private static McpStreamableHttpRequestHandler newHandler() { + return McpStreamableHttpRequestHandler.builder() + .objectMapper(OBJECT_MAPPER) + .mcpEndpoint(MCP_ENDPOINT) + .build(); + } + + private static EmbeddedChannel newChannel(McpStreamableHttpRequestHandler handler) { + return new EmbeddedChannel(new SimpleChannelInboundHandler(false) { + @Override + protected void channelRead0(ChannelHandlerContext ctx, FullHttpRequest request) throws Exception { + handler.handle(ctx, request); + } + }); + } + + private static T readOutbound(EmbeddedChannel channel, Class type) { + Object message = channel.readOutbound(); + assertThat(message).isInstanceOf(type); + return type.cast(message); + } +}