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;
     }

Reply via email to