Skip to content
Open
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 @@ -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.
* <p>
* 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"));
Expand Down Expand Up @@ -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"));
Expand Down
Original file line number Diff line number Diff line change
@@ -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();
Comment on lines +52 to +55
}

@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<FullHttpRequest>(false) {
@Override
protected void channelRead0(ChannelHandlerContext ctx, FullHttpRequest request) throws Exception {
handler.handle(ctx, request);
}
});
}
Comment on lines +85 to +92

private static <T> T readOutbound(EmbeddedChannel channel, Class<T> type) {
Object message = channel.readOutbound();
assertThat(message).isInstanceOf(type);
return type.cast(message);
}
}