diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java index 3d7154278..808b1940b 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpClientSession.java @@ -264,21 +264,25 @@ public Mono sendRequest(String method, Object requestParams, TypeRef t this.pendingResponses.remove(requestId); pendingResponseSink.error(error); }); - })).timeout(this.requestTimeout).handle((jsonRpcResponse, deliveredResponseSink) -> { - if (jsonRpcResponse.error() != null) { - logger.info("Server returned a JSON-RPC error when calling method {}: {}", method, - jsonRpcResponse.error()); - deliveredResponseSink.error(new McpError(jsonRpcResponse.error())); - } - else { - if (typeRef.getType().equals(Void.class)) { - deliveredResponseSink.complete(); + })) + .timeout(this.requestTimeout) + .doOnError(e -> this.pendingResponses.remove(requestId)) + .doOnCancel(() -> this.pendingResponses.remove(requestId)) + .handle((jsonRpcResponse, deliveredResponseSink) -> { + if (jsonRpcResponse.error() != null) { + logger.info("Server returned a JSON-RPC error when calling method {}: {}", method, + jsonRpcResponse.error()); + deliveredResponseSink.error(new McpError(jsonRpcResponse.error())); } else { - deliveredResponseSink.next(this.transport.unmarshalFrom(jsonRpcResponse.result(), typeRef)); + if (typeRef.getType().equals(Void.class)) { + deliveredResponseSink.complete(); + } + else { + deliveredResponseSink.next(this.transport.unmarshalFrom(jsonRpcResponse.result(), typeRef)); + } } - } - }); + }); } /** diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpServerSession.java b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpServerSession.java index 8f86138f0..587000c30 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpServerSession.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/spec/McpServerSession.java @@ -184,19 +184,23 @@ public Mono sendRequest(String method, Object requestParams, TypeRef t this.pendingResponses.remove(requestId); sink.error(error); }); - }).timeout(requestTimeout).handle((jsonRpcResponse, sink) -> { - if (jsonRpcResponse.error() != null) { - sink.error(new McpError(jsonRpcResponse.error())); - } - else { - if (typeRef.getType().equals(Void.class)) { - sink.complete(); + }) + .timeout(requestTimeout) + .doOnError(e -> this.pendingResponses.remove(requestId)) + .doOnCancel(() -> this.pendingResponses.remove(requestId)) + .handle((jsonRpcResponse, sink) -> { + if (jsonRpcResponse.error() != null) { + sink.error(new McpError(jsonRpcResponse.error())); } else { - sink.next(this.transport.unmarshalFrom(jsonRpcResponse.result(), typeRef)); + if (typeRef.getType().equals(Void.class)) { + sink.complete(); + } + else { + sink.next(this.transport.unmarshalFrom(jsonRpcResponse.result(), typeRef)); + } } - } - }); + }); } @Override diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java index ae5daf1f4..a2ab5e8c5 100644 --- a/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java +++ b/mcp-core/src/test/java/io/modelcontextprotocol/spec/McpClientSessionTests.java @@ -303,4 +303,22 @@ void testGracefulShutdown() { StepVerifier.create(session.closeGracefully()).verifyComplete(); } + @Test + void testRequestTimeoutRemovesPendingResponse() throws Exception { + var transport = new MockMcpClientTransport(); + var session = new McpClientSession(Duration.ofMillis(50), transport, Map.of(), Map.of(), Function.identity()); + + Mono responseMono = session.sendRequest(TEST_METHOD, "test", responseType); + + StepVerifier.create(responseMono).expectError(java.util.concurrent.TimeoutException.class).verify(); + + var field = McpClientSession.class.getDeclaredField("pendingResponses"); + field.setAccessible(true); + @SuppressWarnings("unchecked") + var pendingResponses = (java.util.Map) field.get(session); + assertThat(pendingResponses).isEmpty(); + + session.close(); + } + }