This is an automated email from the ASF dual-hosted git repository.
Aias00 pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/shenyu.git
The following commit(s) were added to refs/heads/master by this push:
new c63fc2a98f fix(mcp): isolate legacy session requests (#7337)
c63fc2a98f is described below
commit c63fc2a98fad989127aaf89964c5963ff9957d61
Author: BobSong <[email protected]>
AuthorDate: Thu Oct 1 06:27:18 2026 +0800
fix(mcp): isolate legacy session requests (#7337)
---
.../mcp/server/callback/ShenyuToolCallback.java | 12 +-
.../mcp/server/session/McpSessionHelper.java | 3 +
...henyuStreamableHttpServerTransportProvider.java | 76 ++++++-----
.../server/callback/ShenyuToolCallbackTest.java | 21 ++++
...uStreamableHttpServerTransportProviderTest.java | 139 ++++++++++++++++++++-
5 files changed, 212 insertions(+), 39 deletions(-)
diff --git
a/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallback.java
b/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallback.java
index 721c46556d..524da0a726 100644
---
a/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallback.java
+++
b/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallback.java
@@ -150,7 +150,7 @@ public class ShenyuToolCallback implements ToolCallback {
final String configStr = extractRequestConfig(shenyuTool);
// Get pre-stored exchange and plugin chain
- final ServerWebExchange originExchange =
getOriginExchange(sessionId);
+ final ServerWebExchange originExchange =
getOriginExchange(mcpExchange, sessionId);
final ShenyuPluginChain chain = getPluginChain(originExchange);
// Execute the tool call through the plugin chain
@@ -813,13 +813,19 @@ public class ShenyuToolCallback implements ToolCallback {
}
/**
- * Gets the origin ServerWebExchange for the given session ID.
+ * Gets the request-local exchange, falling back to the legacy session
holder.
*
+ * @param mcpExchange the current MCP request exchange
* @param sessionId the session ID
* @return the origin ServerWebExchange
* @throws IllegalStateException if exchange cannot be retrieved
*/
- private ServerWebExchange getOriginExchange(final String sessionId) {
+ private ServerWebExchange getOriginExchange(final McpSyncServerExchange
mcpExchange, final String sessionId) {
+ final Object contextualExchange =
Objects.isNull(mcpExchange.transportContext()) ? null
+ :
mcpExchange.transportContext().get(McpSessionHelper.SHENYU_EXCHANGE_CONTEXT_KEY);
+ if (contextualExchange instanceof ServerWebExchange exchange) {
+ return exchange;
+ }
final ServerWebExchange exchange =
ShenyuMcpExchangeHolder.get(sessionId);
if (Objects.nonNull(exchange)) {
LOG.debug("Found existing exchange for session: {}", sessionId);
diff --git
a/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/session/McpSessionHelper.java
b/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/session/McpSessionHelper.java
index fe3c114f4c..13a6bcba33 100644
---
a/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/session/McpSessionHelper.java
+++
b/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/session/McpSessionHelper.java
@@ -45,6 +45,9 @@ import java.util.Objects;
*/
public class McpSessionHelper {
+ /** Reactor transport-context key carrying the request-local ShenYu
exchange. */
+ public static final String SHENYU_EXCHANGE_CONTEXT_KEY =
"shenyu.serverWebExchange";
+
private static final Logger LOG =
LoggerFactory.getLogger(McpSessionHelper.class);
/**
diff --git
a/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProvider.java
b/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProvider.java
index e827be0158..344257f912 100644
---
a/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProvider.java
+++
b/shenyu-plugin/shenyu-plugin-mcp-server/src/main/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProvider.java
@@ -26,8 +26,10 @@ import io.modelcontextprotocol.spec.McpSchema;
import io.modelcontextprotocol.spec.McpServerSession;
import io.modelcontextprotocol.spec.McpServerTransport;
import io.modelcontextprotocol.spec.McpServerTransportProvider;
+import io.modelcontextprotocol.common.McpTransportContext;
import io.modelcontextprotocol.util.Assert;
import org.apache.shenyu.plugin.mcp.server.holder.ShenyuMcpExchangeHolder;
+import org.apache.shenyu.plugin.mcp.server.session.McpSessionHelper;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.http.HttpHeaders;
@@ -395,7 +397,7 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
} else {
LOGGER.debug("Exchange mapping already exists for session {},
using existing exchange", requestedSessionId);
}
- return processWithExistingSession(existingSession, requestedSessionId,
message, messageId)
+ return processWithExistingSession(exchange, existingSession,
requestedSessionId, message, messageId)
.map(result -> {
if (!requestedSessionId.equals(result.getSessionId())) {
LOGGER.info("Returning actual session ID {} instead of
requested ID {}", result.getSessionId(), requestedSessionId);
@@ -436,7 +438,7 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
LOGGER.debug("Bound exchange to temporary session: {} before
handshake", actualSessionId);
initializeSessionDirectly(tempSession, actualSessionId);
tempTransport.resetCapturedMessage();
- return processWithExistingSession(tempSession, actualSessionId,
message, messageId)
+ return processWithExistingSession(exchange, tempSession,
actualSessionId, message, messageId)
.doFinally(signalType -> {
LOGGER.debug("Cleaning up temporary session: {}
(signal: {})", actualSessionId, signalType);
removeSession(actualSessionId);
@@ -483,7 +485,7 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
LOGGER.debug("Bound exchange to restored session: {} before
handshake", actualSessionId);
initializeSessionDirectly(newSession, actualSessionId);
newTransport.resetCapturedMessage();
- return processWithExistingSession(newSession, actualSessionId,
message, messageId)
+ return processWithExistingSession(exchange, newSession,
actualSessionId, message, messageId)
.doFinally(signalType -> {
LOGGER.debug("Cleaning up restored session: {}
(signal: {})", actualSessionId, signalType);
removeSession(actualSessionId);
@@ -509,24 +511,20 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
* This method contains the core message processing logic that is shared
* between regular requests, temporary sessions, and restored sessions.
*
+ * @param exchange the current request exchange
* @param session the MCP server session
* @param sessionId the session identifier
* @param message the JSON-RPC message
* @param messageId the message ID for correlation
* @return a Mono containing the processing result
*/
- private Mono<MessageHandlingResult> processWithExistingSession(final
McpServerSession session,
+ private Mono<MessageHandlingResult> processWithExistingSession(final
ServerWebExchange exchange,
+ final
McpServerSession session,
final
String sessionId,
final
McpSchema.JSONRPCMessage message,
final
Object messageId) {
// Verify exchange is available before processing
- final ServerWebExchange verifyExchange =
ShenyuMcpExchangeHolder.get(sessionId);
- if (Objects.isNull(verifyExchange)) {
- LOGGER.error("CRITICAL: No exchange found in
ShenyuMcpExchangeHolder for session {} when processing business request. This
will cause ToolCallback to fail.", sessionId);
- } else {
- LOGGER.debug("Exchange verification passed for session {}
(exchange ID: {})",
- sessionId, System.identityHashCode(verifyExchange));
- }
+ LOGGER.debug("Processing request exchange {} for session {}",
System.identityHashCode(exchange), sessionId);
final StreamableHttpSessionTransport transport =
getSessionTransport(sessionId);
// JSON-RPC notifications (messages without an id, e.g.
notifications/initialized,
// notifications/cancelled) must be acknowledged with HTTP 202 and an
empty body
@@ -535,7 +533,10 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
// must not wait for a captured transport response here. Otherwise a
stale response
// from a previous request could be replayed with the wrong id.
if (message instanceof McpSchema.JSONRPCNotification) {
- return session.handle(message)
+ final McpTransportContext transportContext =
McpTransportContext.create(
+ Map.of(McpSessionHelper.SHENYU_EXCHANGE_CONTEXT_KEY,
exchange));
+ return Mono.defer(() -> session.handle(message))
+ .contextWrite(context ->
context.put(McpTransportContext.KEY, transportContext))
.cast(Object.class)
.doOnSuccess(result -> LOGGER.debug("Successfully
processed notification for session: {}", sessionId))
.thenReturn(createMessageHandlingResult(HttpStatus.ACCEPTED.value(), null,
sessionId))
@@ -547,14 +548,16 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
});
}
// Let MCP framework handle the message - framework will send response
through transport
- return session.handle(message)
+ final McpTransportContext transportContext =
McpTransportContext.create(
+ Map.of(McpSessionHelper.SHENYU_EXCHANGE_CONTEXT_KEY,
exchange));
+ return Mono.defer(() -> session.handle(message))
+ // The exchange belongs to this JSON-RPC request, not to the
session.
+ .contextWrite(context -> context.put(McpTransportContext.KEY,
transportContext))
.cast(Object.class)
.doOnSuccess(result -> LOGGER.debug("Successfully processed
message for session: {}", sessionId))
.then(waitForTransportResponse(transport, sessionId,
messageId))
- .doOnNext(result -> {
- // Clear the response captured for this specific message
id after it has
- // been delivered, so that a subsequent message on this
session cannot
- // observe a stale response from a previous request.
+ .doFinally(signal -> {
+ // Release only this request's response on completion,
error or cancellation.
if (Objects.nonNull(transport)) {
transport.resetCapturedMessage(messageId);
}
@@ -804,6 +807,10 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
public void removeSession(final String sessionId) {
final McpServerSession removedSession = sessions.remove(sessionId);
final StreamableHttpSessionTransport removedTransport =
sessionTransports.remove(sessionId);
+ if (Objects.nonNull(removedTransport)) {
+ removedTransport.closed = true;
+ removedTransport.resetCapturedMessage();
+ }
ShenyuMcpExchangeHolder.remove(sessionId);
if (Objects.nonNull(removedSession) ||
Objects.nonNull(removedTransport)) {
LOGGER.debug("Removed session and transport: {}", sessionId);
@@ -834,13 +841,8 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
final String
sessionId,
final Object
messageId) {
return Mono.fromCallable(() -> {
- final McpSchema.JSONRPCMessage correlatedResponse =
Objects.nonNull(transport)
- ? transport.getLastSentMessage(messageId) : null;
- if (Objects.nonNull(messageId) &&
Objects.nonNull(correlatedResponse)) {
- LOGGER.debug("Retrieved correlated response for message id {}
on session: {}", messageId, sessionId);
- return createMessageHandlingResult(200, correlatedResponse,
sessionId);
- } else if (Objects.nonNull(transport) &&
transport.isResponseReady() && Objects.nonNull(transport.getLastSentMessage()))
{
- final McpSchema.JSONRPCMessage sentMessage =
transport.getLastSentMessage();
+ if (Objects.nonNull(transport) &&
transport.isResponseReady(messageId)) {
+ final McpSchema.JSONRPCMessage sentMessage =
transport.getResponse(messageId);
LOGGER.debug("Retrieved captured response from transport for
session: {}", sessionId);
return createMessageHandlingResult(200, sentMessage,
sessionId);
} else {
@@ -1058,9 +1060,9 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
private volatile McpSchema.JSONRPCMessage lastSentMessage;
- private volatile boolean responseReady;
+ private final ConcurrentHashMap<Object, McpSchema.JSONRPCMessage>
responses = new ConcurrentHashMap<>();
- private final Map<String, McpSchema.JSONRPCMessage> messageResponses =
new ConcurrentHashMap<>();
+ private volatile boolean responseReady;
/**
* Creates a new session transport with auto-generated session ID.
@@ -1100,7 +1102,7 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
*/
public McpSchema.JSONRPCMessage getLastSentMessage(final Object
messageId) {
if (Objects.nonNull(messageId)) {
- return messageResponses.get(String.valueOf(messageId));
+ return responses.get(messageId);
}
return lastSentMessage;
}
@@ -1114,17 +1116,22 @@ public class
ShenyuStreamableHttpServerTransportProvider implements McpServerTra
return responseReady;
}
+ public boolean isResponseReady(final Object messageId) {
+ return Objects.nonNull(messageId) &&
responses.containsKey(messageId);
+ }
+
+ public McpSchema.JSONRPCMessage getResponse(final Object messageId) {
+ return responses.get(messageId);
+ }
+
@Override
public Mono<Void> sendMessage(final McpSchema.JSONRPCMessage message) {
if (!closed) {
this.lastSentMessage = message;
- this.responseReady = true;
- if (message instanceof McpSchema.JSONRPCResponse) {
- final Object responseId = ((McpSchema.JSONRPCResponse)
message).id();
- if (Objects.nonNull(responseId)) {
- this.messageResponses.put(String.valueOf(responseId),
message);
- }
+ if (message instanceof McpSchema.JSONRPCResponse response &&
Objects.nonNull(response.id())) {
+ responses.put(response.id(), message);
}
+ this.responseReady = true;
LOGGER.debug("Captured response message for session: {}",
sessionId);
}
return Mono.empty();
@@ -1161,6 +1168,7 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
* is executed during short-reconnect or temporary session creation.
*/
public void resetCapturedMessage() {
+ this.responses.clear();
this.lastSentMessage = null;
this.responseReady = false;
}
@@ -1173,7 +1181,7 @@ public class ShenyuStreamableHttpServerTransportProvider
implements McpServerTra
*/
public void resetCapturedMessage(final Object messageId) {
if (Objects.nonNull(messageId)) {
- this.messageResponses.remove(String.valueOf(messageId));
+ this.responses.remove(messageId);
}
}
}
diff --git
a/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallbackTest.java
b/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallbackTest.java
index eb2302558a..4a4b079582 100644
---
a/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallbackTest.java
+++
b/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/callback/ShenyuToolCallbackTest.java
@@ -18,6 +18,10 @@
package org.apache.shenyu.plugin.mcp.server.callback;
import io.modelcontextprotocol.server.McpSyncServerExchange;
+import io.modelcontextprotocol.common.McpTransportContext;
+import org.springframework.test.util.ReflectionTestUtils;
+import java.util.Map;
+import static org.junit.jupiter.api.Assertions.assertSame;
import org.apache.shenyu.common.constant.Constants;
import org.apache.shenyu.plugin.api.ShenyuPluginChain;
import org.apache.shenyu.plugin.api.context.ShenyuContext;
@@ -283,6 +287,23 @@ class ShenyuToolCallbackTest {
});
}
+ @Test
+ void testRequestContextTakesPrecedenceOverSessionHolder() {
+ shenyuToolCallback = new ShenyuToolCallback(toolDefinition);
+ when(mcpSyncServerExchange.transportContext()).thenReturn(
+
McpTransportContext.create(Map.of(McpSessionHelper.SHENYU_EXCHANGE_CONTEXT_KEY,
exchange)));
+ assertSame(exchange,
ReflectionTestUtils.invokeMethod(shenyuToolCallback, "getOriginExchange",
mcpSyncServerExchange, "shared-session"));
+ exchangeHolderMock.verifyNoInteractions();
+ }
+
+ @Test
+ void testLegacySessionHolderFallback() {
+ shenyuToolCallback = new ShenyuToolCallback(toolDefinition);
+ exchangeHolderMock.when(() ->
ShenyuMcpExchangeHolder.get("legacy-session")).thenReturn(exchange);
+ assertSame(exchange,
ReflectionTestUtils.invokeMethod(shenyuToolCallback, "getOriginExchange",
mcpSyncServerExchange, "legacy-session"));
+ exchangeHolderMock.verify(() ->
ShenyuMcpExchangeHolder.get("legacy-session"));
+ }
+
@Test
void testConstructorWithValidToolDefinition() {
ShenyuToolCallback callback = new ShenyuToolCallback(toolDefinition);
diff --git
a/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProviderTest.java
b/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProviderTest.java
index 2e17944c90..f2b160e5ac 100644
---
a/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProviderTest.java
+++
b/shenyu-plugin/shenyu-plugin-mcp-server/src/test/java/org/apache/shenyu/plugin/mcp/server/transport/ShenyuStreamableHttpServerTransportProviderTest.java
@@ -21,6 +21,11 @@ import com.fasterxml.jackson.databind.ObjectMapper;
import io.modelcontextprotocol.spec.McpSchema;
import io.modelcontextprotocol.spec.McpServerSession;
import io.modelcontextprotocol.server.McpRequestHandler;
+import org.apache.shenyu.plugin.mcp.server.session.McpSessionHelper;
+import org.springframework.web.server.ServerWebExchange;
+import reactor.core.publisher.Flux;
+import reactor.core.Disposable;
+import java.util.concurrent.atomic.AtomicBoolean;
import org.apache.shenyu.plugin.mcp.server.holder.ShenyuMcpExchangeHolder;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
@@ -37,6 +42,7 @@ import reactor.core.publisher.Mono;
import reactor.test.StepVerifier;
import java.lang.reflect.Field;
+import java.lang.reflect.Method;
import java.time.Duration;
import java.util.Collections;
import java.util.List;
@@ -204,6 +210,25 @@ class ShenyuStreamableHttpServerTransportProviderTest {
assertEquals(0, readMap(provider, "sessionTransports").size());
}
+ @Test
+ void testCompatibilityLookupSharesRequestResponseCleanup() throws
Exception {
+ final ShenyuStreamableHttpServerTransportProvider provider =
providerWithRealSessions();
+ final String sessionId = initialize(provider);
+ performRequest(provider, postRequest(TOOLS_LIST_REQUEST_BODY,
sessionId));
+
+ final Object transport = readMap(provider,
"sessionTransports").get(sessionId);
+ assertNotNull(transport);
+ final Method lookup =
transport.getClass().getDeclaredMethod("getLastSentMessage", Object.class);
+ lookup.setAccessible(true);
+ assertNull(lookup.invoke(transport, "req-1"));
+
+ final Method reset =
transport.getClass().getDeclaredMethod("resetCapturedMessage", Object.class);
+ reset.setAccessible(true);
+ reset.invoke(transport, new Object[] {null});
+ assertNull(lookup.invoke(transport, "req-1"));
+ provider.removeSession(sessionId);
+ }
+
@SuppressWarnings("unchecked")
private Map<String, ?> readMap(final
ShenyuStreamableHttpServerTransportProvider provider, final String fieldName)
throws Exception {
@@ -213,6 +238,117 @@ class ShenyuStreamableHttpServerTransportProviderTest {
}
private ShenyuStreamableHttpServerTransportProvider
providerWithRealSessions() {
+ return providerWithHandler((exchange, params) ->
Mono.just(Map.<String, Object>of("tools", List.of())));
+ }
+
+ @Test
+ void testConcurrentRequestsKeepExchangeAndResponseCorrelation() throws
Exception {
+ final ShenyuStreamableHttpServerTransportProvider provider =
providerWithHandler((exchange, params) -> {
+ final ServerWebExchange requestExchange = (ServerWebExchange)
exchange.transportContext()
+ .get(McpSessionHelper.SHENYU_EXCHANGE_CONTEXT_KEY);
+ final String marker =
requestExchange.getRequest().getHeaders().getFirst("X-Test-Request");
+ return Mono.delay(Duration.ofMillis(80 - Integer.parseInt(marker)
* 5L))
+ .thenReturn(Map.<String, Object>of("marker", marker));
+ });
+ final String sessionId = initialize(provider);
+ StepVerifier.create(Flux.range(0, 8).flatMap(index -> {
+ final String id = Integer.toString(index);
+ final MockServerHttpRequest request =
MockServerHttpRequest.post("/mcp/streamablehttp")
+ .header("Content-Type",
"application/json").header(SESSION_ID_HEADER, sessionId)
+ .header("X-Test-Request",
id).body(TOOLS_LIST_REQUEST_BODY.replace("req-1", id));
+ final MockServerWebExchange exchange =
MockServerWebExchange.from(request);
+ return
provider.handleUnifiedEndpoint(ServerRequest.create(exchange,
HandlerStrategies.withDefaults().messageReaders()))
+ .flatMap(response -> response.writeTo(exchange,
RESPONSE_CONTEXT))
+ .then(Mono.defer(() ->
exchange.getResponse().getBodyAsString()))
+ .doOnNext(body -> {
+ assertTrue(body.contains("\"id\":\"" + id + "\""),
body);
+ assertTrue(body.contains("\"marker\":\"" + id + "\""),
body);
+ });
+ })).expectNextCount(8).verifyComplete();
+ final Object transport = readMap(provider,
"sessionTransports").get(sessionId);
+ final Field responses =
transport.getClass().getDeclaredField("responses");
+ responses.setAccessible(true);
+ assertTrue(((Map<?, ?>) responses.get(transport)).isEmpty(),
"Completed requests must release captured responses");
+ provider.removeSession(sessionId);
+ assertNull(ShenyuMcpExchangeHolder.get(sessionId));
+ }
+
+ @Test
+ void testCancelReachesSessionHandler() {
+ final AtomicBoolean cancelled = new AtomicBoolean();
+ final ShenyuStreamableHttpServerTransportProvider provider =
providerWithHandler((exchange, params) ->
+ Mono.<Map<String, Object>>never().doOnCancel(() ->
cancelled.set(true)));
+ final String sessionId = initialize(provider);
+
StepVerifier.create(provider.handleUnifiedEndpoint(createRequest(postRequest(TOOLS_LIST_REQUEST_BODY,
sessionId))))
+
.thenAwait(Duration.ofMillis(100)).thenCancel().verify(Duration.ofSeconds(3));
+ assertTrue(cancelled.get());
+ provider.removeSession(sessionId);
+ }
+
+ private String initialize(final
ShenyuStreamableHttpServerTransportProvider provider) {
+ final String sessionId = performRequest(provider,
postRequest(INITIALIZE_REQUEST_BODY, null))
+ .getHeaders().getFirst(SESSION_ID_HEADER);
+ assertNotNull(sessionId);
+ performRequest(provider, postRequest(INITIALIZED_NOTIFICATION_BODY,
sessionId));
+ return sessionId;
+ }
+
+ @Test
+ void testFailedRequestDoesNotPoisonNextRequest() {
+ final AtomicBoolean fail = new AtomicBoolean(true);
+ final ShenyuStreamableHttpServerTransportProvider provider =
providerWithHandler((exchange, params) -> {
+ if (fail.getAndSet(false)) {
+ return Mono.error(new IllegalStateException("fixture
failure"));
+ }
+ return Mono.just(Map.<String, Object>of("tools", List.of()));
+ });
+ final String sessionId = initialize(provider);
+ final String failed = performRequest(provider,
postRequest(TOOLS_LIST_REQUEST_BODY, sessionId)).getBodyAsString().block();
+ assertNotNull(failed);
+ assertTrue(failed.contains("\"error\""), failed);
+ final String succeeded = performRequest(provider,
postRequest(TOOLS_LIST_REQUEST_BODY.replace("req-1", "req-2"), sessionId))
+ .getBodyAsString().block();
+ assertNotNull(succeeded);
+ assertTrue(succeeded.contains("\"id\":\"req-2\""), succeeded);
+ assertTrue(succeeded.contains("\"tools\""), succeeded);
+ provider.removeSession(sessionId);
+ }
+
+ @Test
+ void testCancellingOneRequestDoesNotCancelAnotherInSameSession() {
+ final AtomicBoolean cancelled = new AtomicBoolean();
+ final ShenyuStreamableHttpServerTransportProvider provider =
providerWithHandler((exchange, params) -> {
+ final ServerWebExchange requestExchange = (ServerWebExchange)
exchange.transportContext()
+ .get(McpSessionHelper.SHENYU_EXCHANGE_CONTEXT_KEY);
+ if
("cancel".equals(requestExchange.getRequest().getHeaders().getFirst("X-Test-Request")))
{
+ return Mono.<Map<String, Object>>never().doOnCancel(() ->
cancelled.set(true));
+ }
+ return Mono.just(Map.<String, Object>of("tools", List.of()));
+ });
+ final String sessionId = initialize(provider);
+ final MockServerHttpRequest request =
MockServerHttpRequest.post("/mcp/streamablehttp")
+ .header("Content-Type",
"application/json").header(SESSION_ID_HEADER, sessionId)
+ .header("X-Test-Request",
"cancel").body(TOOLS_LIST_REQUEST_BODY);
+ final Disposable pending =
provider.handleUnifiedEndpoint(createRequest(request)).subscribe();
+ try {
+ final String other = performRequest(provider,
postRequest(TOOLS_LIST_REQUEST_BODY.replace("req-1", "req-2"), sessionId))
+ .getBodyAsString().block(Duration.ofSeconds(3));
+ assertNotNull(other);
+ assertTrue(other.contains("\"id\":\"req-2\""), other);
+ assertTrue(other.contains("\"tools\""), other);
+ pending.dispose();
+ assertTrue(cancelled.get());
+ final String next = performRequest(provider,
postRequest(TOOLS_LIST_REQUEST_BODY.replace("req-1", "req-3"), sessionId))
+ .getBodyAsString().block(Duration.ofSeconds(3));
+ assertNotNull(next);
+ assertTrue(next.contains("\"id\":\"req-3\""), next);
+ } finally {
+ pending.dispose();
+ provider.removeSession(sessionId);
+ }
+ }
+
+ private ShenyuStreamableHttpServerTransportProvider
providerWithHandler(final McpRequestHandler<Map<String, Object>> handler) {
ShenyuStreamableHttpServerTransportProvider provider =
new ShenyuStreamableHttpServerTransportProvider(new
ObjectMapper(), "/mcp/streamablehttp");
provider.setSessionFactory(transport -> new McpServerSession(
@@ -224,8 +360,7 @@ class ShenyuStreamableHttpServerTransportProviderTest {
McpSchema.ServerCapabilities.builder().build(),
new McpSchema.Implementation("ShenyuMcpServer",
"1.0.0"),
"test")),
- Map.<String, McpRequestHandler<?>>of("tools/list",
- (McpRequestHandler<Map<String, Object>>) (exchange,
params) -> Mono.just(Map.<String, Object>of("tools", List.of()))),
+ Map.<String, McpRequestHandler<?>>of("tools/list", handler),
Map.of()));
return provider;
}