Skip to content

Commit 17de368

Browse files
committed
fix(core): keep streamable HTTP session alive when an SSE stream write fails
A failed write inside HttpServletStreamableMcpSessionTransport.sendMessage removed the whole McpStreamableServerSession from the provider. Any transient failure (slow or aborted client, momentary IO error) therefore made the server forget the session: the client's next POST got a 404 and client sessions surfaced 'MCP session with server terminated' (#952). A failed write now breaks only the affected SSE stream, mirroring close(): the transport is marked closed so further sends on the dead stream are no-ops, the async context is completed behind a guard, and session removal stays with explicit DELETE handling and lifecycle events. Clients can still reopen a stream via GET with Last-Event-ID. Added servlet-mock regression tests: a failed stream write no longer evicts the session (subsequent POST is accepted), while an explicit DELETE still removes it.
1 parent 4186ca1 commit 17de368

2 files changed

Lines changed: 210 additions & 2 deletions

File tree

‎mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java‎

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -768,8 +768,19 @@ public Mono<Void> sendMessage(McpSchema.JSONRPCMessage message, String messageId
768768
}
769769
catch (Exception e) {
770770
logger.error("Failed to send message to session {}: {}", this.sessionId, e.getMessage());
771-
HttpServletStreamableServerTransportProvider.this.sessions.remove(this.sessionId);
772-
this.asyncContext.complete();
771+
// A failed write breaks this SSE stream, not the session itself: the
772+
// client may reopen a stream (GET with Last-Event-ID) or keep
773+
// POSTing.
774+
// Session removal stays with explicit DELETE handling and lifecycle
775+
// events, mirroring close() below.
776+
this.closed = true;
777+
try {
778+
this.asyncContext.complete();
779+
}
780+
catch (Exception completionError) {
781+
logger.warn("Failed to complete async context for session {}: {}", this.sessionId,
782+
completionError.getMessage());
783+
}
773784
}
774785
finally {
775786
lock.unlock();
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,197 @@
1+
/*
2+
* Copyright 2025 - 2025 the original author or authors.
3+
*/
4+
5+
package io.modelcontextprotocol.server.transport;
6+
7+
import java.io.ByteArrayInputStream;
8+
import java.io.IOException;
9+
import java.io.PrintWriter;
10+
import java.io.StringWriter;
11+
import java.nio.charset.StandardCharsets;
12+
import java.time.Duration;
13+
import java.util.Map;
14+
15+
import com.google.gson.Gson;
16+
import com.google.gson.GsonBuilder;
17+
import com.google.gson.JsonObject;
18+
import com.google.gson.ToNumberPolicy;
19+
import com.google.gson.JsonSerializer;
20+
import io.modelcontextprotocol.json.McpJsonMapper;
21+
import io.modelcontextprotocol.spec.HttpHeaders;
22+
import io.modelcontextprotocol.spec.McpSchema;
23+
import io.modelcontextprotocol.spec.McpStreamableServerSession;
24+
import io.modelcontextprotocol.spec.json.gson.GsonMcpJsonMapper;
25+
import jakarta.servlet.AsyncContext;
26+
import jakarta.servlet.ReadListener;
27+
import jakarta.servlet.ServletInputStream;
28+
import jakarta.servlet.http.HttpServletRequest;
29+
import jakarta.servlet.http.HttpServletResponse;
30+
import org.junit.jupiter.api.BeforeEach;
31+
import org.junit.jupiter.api.Test;
32+
import org.mockito.ArgumentCaptor;
33+
import reactor.core.publisher.Mono;
34+
35+
import static org.mockito.ArgumentMatchers.eq;
36+
import static org.mockito.Mockito.mock;
37+
import static org.mockito.Mockito.verify;
38+
import static org.mockito.Mockito.when;
39+
40+
/**
41+
* Tests that a failed write on one SSE stream does not evict the whole streamable HTTP
42+
* session, while explicit DELETE requests still do.
43+
*/
44+
class HttpServletStreamableServerTransportProviderSessionRetentionTests {
45+
46+
private static final String INIT_BODY = """
47+
{"jsonrpc":"2.0","id":"1","method":"initialize","params":{"protocolVersion":"2025-06-18",
48+
"capabilities":{},"clientInfo":{"name":"it","version":"1.0"}}}""";
49+
50+
private static final String NOTIFICATION_BODY = """
51+
{"jsonrpc":"2.0","method":"notifications/test","params":{}}""";
52+
53+
private HttpServletStreamableServerTransportProvider provider;
54+
55+
private final StringWriter initResponseBody = new StringWriter();
56+
57+
@BeforeEach
58+
void setUp() {
59+
this.provider = HttpServletStreamableServerTransportProvider.builder()
60+
.jsonMapper(jsonMapper())
61+
.mcpEndpoint("/mcp")
62+
.build();
63+
this.provider.setSessionFactory(this::startSession);
64+
}
65+
66+
private McpJsonMapper jsonMapper() {
67+
// The test Gson needs a Throwable adapter: responseError serializes McpError
68+
// (a RuntimeException), and reflective serialization of Throwable fields is
69+
// not accessible on JDK 17+.
70+
Gson gson = new GsonBuilder().serializeNulls()
71+
.setObjectToNumberStrategy(ToNumberPolicy.LONG_OR_DOUBLE)
72+
.setNumberToNumberStrategy(ToNumberPolicy.LONG_OR_DOUBLE)
73+
.registerTypeHierarchyAdapter(Throwable.class, (JsonSerializer<Throwable>) (src, typeOfSrc, context) -> {
74+
JsonObject error = new JsonObject();
75+
error.addProperty("message", src.getMessage());
76+
return error;
77+
})
78+
.create();
79+
return new GsonMcpJsonMapper(gson);
80+
}
81+
82+
private McpStreamableServerSession.McpStreamableServerSessionInit startSession(
83+
McpSchema.InitializeRequest initializeRequest) {
84+
McpStreamableServerSession session = new McpStreamableServerSession("test-session-id",
85+
new McpSchema.ClientCapabilities(null, null, null, null),
86+
new McpSchema.Implementation("it", null, "1.0", null, null, null), Duration.ofSeconds(5), Map.of(),
87+
Map.of());
88+
McpSchema.InitializeResult initResult = new McpSchema.InitializeResult("2025-06-18",
89+
McpSchema.ServerCapabilities.builder().build(),
90+
new McpSchema.Implementation("test-server", null, "1.0", null, null, null), null, null);
91+
return new McpStreamableServerSession.McpStreamableServerSessionInit(session, Mono.just(initResult));
92+
}
93+
94+
@Test
95+
void sessionSurvivesFailedWriteOnSseStream() throws Exception {
96+
String sessionId = initializeSession();
97+
98+
openBrokenSseStream(sessionId);
99+
100+
// The write failure happens inside the notification delivery
101+
this.provider.notifyClient(sessionId, "notifications/test", Map.of()).block();
102+
103+
// The session must still be usable for subsequent client requests
104+
HttpServletRequest post = postRequest(NOTIFICATION_BODY, sessionId);
105+
HttpServletResponse response = mockResponse(new PrintWriter(new StringWriter()));
106+
this.provider.doPost(post, response);
107+
108+
verify(response).setStatus(HttpServletResponse.SC_ACCEPTED);
109+
}
110+
111+
@Test
112+
void explicitDeleteStillRemovesSession() throws Exception {
113+
String sessionId = initializeSession();
114+
115+
HttpServletRequest delete = mock(HttpServletRequest.class);
116+
when(delete.getRequestURI()).thenReturn("/mcp");
117+
when(delete.getHeader(HttpHeaders.MCP_SESSION_ID)).thenReturn(sessionId);
118+
HttpServletResponse deleteResponse = mockResponse(new PrintWriter(new StringWriter()));
119+
this.provider.doDelete(delete, deleteResponse);
120+
verify(deleteResponse).setStatus(HttpServletResponse.SC_OK);
121+
122+
HttpServletRequest post = postRequest(NOTIFICATION_BODY, sessionId);
123+
HttpServletResponse postResponse = mockResponse(new PrintWriter(new StringWriter()));
124+
this.provider.doPost(post, postResponse);
125+
126+
verify(postResponse).setStatus(HttpServletResponse.SC_NOT_FOUND);
127+
}
128+
129+
private String initializeSession() throws Exception {
130+
HttpServletRequest post = postRequest(INIT_BODY, null);
131+
HttpServletResponse response = mockResponse(new PrintWriter(this.initResponseBody));
132+
this.provider.doPost(post, response);
133+
134+
ArgumentCaptor<String> sessionId = ArgumentCaptor.forClass(String.class);
135+
verify(response).setHeader(eq(HttpHeaders.MCP_SESSION_ID), sessionId.capture());
136+
return sessionId.getValue();
137+
}
138+
139+
private void openBrokenSseStream(String sessionId) throws Exception {
140+
HttpServletRequest get = mock(HttpServletRequest.class);
141+
when(get.getRequestURI()).thenReturn("/mcp");
142+
when(get.getHeader(HttpHeaders.ACCEPT)).thenReturn("text/event-stream");
143+
when(get.getHeader(HttpHeaders.MCP_SESSION_ID)).thenReturn(sessionId);
144+
145+
PrintWriter brokenWriter = mock(PrintWriter.class);
146+
when(brokenWriter.checkError()).thenReturn(true);
147+
148+
AsyncContext asyncContext = mock(AsyncContext.class);
149+
when(get.startAsync()).thenReturn(asyncContext);
150+
151+
HttpServletResponse response = mockResponse(brokenWriter);
152+
this.provider.doGet(get, response);
153+
}
154+
155+
private HttpServletRequest postRequest(String body, String sessionId) throws IOException {
156+
HttpServletRequest request = mock(HttpServletRequest.class);
157+
when(request.getRequestURI()).thenReturn("/mcp");
158+
when(request.getHeader(HttpHeaders.ACCEPT)).thenReturn("application/json, text/event-stream");
159+
if (sessionId != null) {
160+
when(request.getHeader(HttpHeaders.MCP_SESSION_ID)).thenReturn(sessionId);
161+
}
162+
when(request.getContentLengthLong()).thenReturn((long) body.getBytes(StandardCharsets.UTF_8).length);
163+
when(request.getInputStream()).thenReturn(bodyStream(body));
164+
return request;
165+
}
166+
167+
private HttpServletResponse mockResponse(PrintWriter writer) throws IOException {
168+
HttpServletResponse response = mock(HttpServletResponse.class);
169+
when(response.getWriter()).thenReturn(writer);
170+
return response;
171+
}
172+
173+
private static ServletInputStream bodyStream(String body) {
174+
ByteArrayInputStream source = new ByteArrayInputStream(body.getBytes(StandardCharsets.UTF_8));
175+
return new ServletInputStream() {
176+
@Override
177+
public boolean isFinished() {
178+
return source.available() == 0;
179+
}
180+
181+
@Override
182+
public boolean isReady() {
183+
return true;
184+
}
185+
186+
@Override
187+
public void setReadListener(ReadListener listener) {
188+
}
189+
190+
@Override
191+
public int read() {
192+
return source.read();
193+
}
194+
};
195+
}
196+
197+
}

0 commit comments

Comments
 (0)