diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java index 60a8d7c13..35b2954f5 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/StdioClientTransport.java @@ -11,7 +11,9 @@ import java.time.Duration; import java.util.ArrayList; import java.util.EnumSet; +import java.util.HashMap; import java.util.List; +import java.util.Map; import java.util.Set; import java.util.concurrent.Executors; import java.util.function.Consumer; @@ -23,6 +25,7 @@ import io.modelcontextprotocol.spec.McpClientTransport; import io.modelcontextprotocol.spec.McpSchema; import io.modelcontextprotocol.spec.McpSchema.JSONRPCMessage; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCResponse; import io.modelcontextprotocol.util.Assert; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -290,7 +293,15 @@ private void startInboundProcessing() { if (!isClosing) { logger.error("Error processing inbound message for line: {}", line, e); } - break; + // One malformed message must not end the session. Fail the + // request it answers, if any, and keep reading. + JSONRPCResponse failure = failureForMalformedResponse(line, e); + if (failure != null && !this.inboundSink.tryEmitNext(failure).isSuccess()) { + if (!isClosing) { + logger.error("Failed to enqueue inbound message: {}", failure); + } + break; + } } } } @@ -311,6 +322,32 @@ private void startInboundProcessing() { }); } + /** + * Builds an error response for a response that could not be deserialized, so that the + * request it answers fails right away instead of waiting for its timeout. + * @param line the raw message + * @param cause the deserialization failure + * @return the error response, or {@code null} if the line is not a response with a + * usable id + */ + private JSONRPCResponse failureForMalformedResponse(String line, Exception cause) { + try { + Map message = this.jsonMapper.readValue(line, new TypeRef>() { + }); + Object id = message.get("id"); + boolean isResponse = !message.containsKey("method") + && (message.containsKey("result") || message.containsKey("error")); + if (isResponse && (id instanceof String || id instanceof Integer || id instanceof Long)) { + return JSONRPCResponse.error(id, new JSONRPCResponse.JSONRPCError(McpSchema.ErrorCodes.INTERNAL_ERROR, + "Received a malformed JSON-RPC response", cause.getMessage())); + } + } + catch (Exception ignored) { + // not JSON, so there is no request to fail + } + return null; + } + /** * Reads a single line, mirroring {@link BufferedReader#readLine()}, but aborting once * more than {@code maxSize} characters have been read without encountering a line diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/StdioClientTransportTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/StdioClientTransportTests.java index 0aad3934a..9ce69fc80 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/StdioClientTransportTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/client/transport/StdioClientTransportTests.java @@ -6,12 +6,20 @@ import java.io.ByteArrayOutputStream; import java.io.PrintStream; +import java.nio.file.Files; +import java.nio.file.Path; import java.time.Duration; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import io.modelcontextprotocol.spec.McpSchema; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCMessage; +import io.modelcontextprotocol.spec.McpSchema.JSONRPCResponse; import org.awaitility.Awaitility; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; import reactor.test.StepVerifier; import static io.modelcontextprotocol.util.McpJsonMapperUtils.JSON_MAPPER; @@ -67,4 +75,50 @@ void shouldRejectInboundMessageExceedingMaxSize() throws Exception { } } + @Test + void shouldFailMalformedResponsesAndKeepProcessing(@TempDir Path tempDir) throws Exception { + // A server process that answers with two malformed responses and a line that + // is not JSON, followed by a valid response. It then stays alive until its + // stdin is closed. + Path serverOutput = tempDir.resolve("server-output.jsonl"); + Files.write(serverOutput, + List.of("{\"id\":\"missing-jsonrpc\",\"result\":{}}", + "{\"jsonrpc\":\"2.0\",\"id\":2,\"result\":{},\"error\":{\"code\":-32000,\"message\":\"boom\"}}", + "this is not json", "{\"jsonrpc\":\"2.0\",\"id\":\"valid\",\"result\":{}}")); + ServerParameters params = ServerParameters.builder("sh") + .args("-c", "cat '" + serverOutput.toString().replace('\\', '/') + "'; cat > /dev/null") + .build(); + + List received = new CopyOnWriteArrayList<>(); + StdioClientTransport transport = new StdioClientTransport(params, JSON_MAPPER); + try { + StepVerifier.create(transport.connect(msg -> msg.doOnNext(received::add))).verifyComplete(); + + Awaitility.await() + .atMost(Duration.ofSeconds(5)) + .pollInterval(Duration.ofMillis(100)) + .untilAsserted(() -> assertThat(received).hasSize(3)); + + // each malformed response fails the request it answers + assertThat(received.get(0)).isInstanceOfSatisfying(JSONRPCResponse.class, response -> { + assertThat(response.id()).isEqualTo("missing-jsonrpc"); + assertThat(response.result()).isNull(); + assertThat(response.error().code()).isEqualTo(McpSchema.ErrorCodes.INTERNAL_ERROR); + }); + assertThat(received.get(1)).isInstanceOfSatisfying(JSONRPCResponse.class, response -> { + assertThat(response.id()).isEqualTo(2); + assertThat(response.result()).isNull(); + assertThat(response.error().code()).isEqualTo(McpSchema.ErrorCodes.INTERNAL_ERROR); + }); + // the transport is still reading, so the valid response gets through + assertThat(received.get(2)).isInstanceOfSatisfying(JSONRPCResponse.class, response -> { + assertThat(response.id()).isEqualTo("valid"); + assertThat(response.error()).isNull(); + }); + } + finally { + StepVerifier.create(transport.closeGracefully()).verifyComplete(); + } + } + }