Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -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;
}
}
}
}
Expand All @@ -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<String, Object> message = this.jsonMapper.readValue(line, new TypeRef<HashMap<String, Object>>() {
});
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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<JSONRPCMessage> 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();
}
}

}