diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtils.java b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtils.java index 5901c384c..ddbced4eb 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtils.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtils.java @@ -18,9 +18,32 @@ */ final class HttpServletRequestUtils { + private static final String APPLICATION_JSON = "application/json"; + private HttpServletRequestUtils() { } + /** + * Checks whether a {@code Content-Type} header value denotes + * {@code application/json}. Only the media type is compared, case-insensitively; + * parameters such as {@code charset} are ignored. This is not a substring search, so + * a value like {@code text/plain; a=application/json} is rejected. + *
+ * Requiring {@code application/json} prevents browsers from sending cross-origin + * JSON-RPC messages as CORS "simple requests" (e.g. with {@code text/plain}), which + * would otherwise reach the server without a preflight. + * @param contentType The {@code Content-Type} header value, may be {@code null} + * @return {@code true} if the media type is {@code application/json} + */ + static boolean isJsonContentType(String contentType) { + if (contentType == null) { + return false; + } + int parametersStart = contentType.indexOf(';'); + String mediaType = parametersStart == -1 ? contentType : contentType.substring(0, parametersStart); + return APPLICATION_JSON.equalsIgnoreCase(mediaType.trim()); + } + /** * Reads the request body, decoded using the request's character encoding (or UTF-8 if * not specified), while bounding the number of bytes read. diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletSseServerTransportProvider.java b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletSseServerTransportProvider.java index a14271afc..18ac4147c 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletSseServerTransportProvider.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletSseServerTransportProvider.java @@ -377,6 +377,14 @@ protected void doPost(HttpServletRequest request, HttpServletResponse response) return; } + if (!HttpServletRequestUtils.isJsonContentType(request.getContentType())) { + this.responseError(response, HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE, + McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + return; + } + // Get the session ID from the request parameter String sessionId = request.getParameter("sessionId"); if (sessionId == null) { @@ -481,6 +489,16 @@ private void sendEvent(PrintWriter writer, String eventType, String data) throws } } + private void responseError(HttpServletResponse response, int httpCode, McpError mcpError) throws IOException { + response.setContentType(APPLICATION_JSON); + response.setCharacterEncoding(UTF_8); + response.setStatus(httpCode); + String jsonError = jsonMapper.writeValueAsString(mcpError); + PrintWriter writer = response.getWriter(); + writer.write(jsonError); + writer.flush(); + } + /** * Cleans up resources when the servlet is being destroyed. *
diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStatelessServerTransport.java b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStatelessServerTransport.java
index 9cd0d04e1..4d346b732 100644
--- a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStatelessServerTransport.java
+++ b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStatelessServerTransport.java
@@ -168,7 +168,13 @@ protected void doPost(HttpServletRequest request, HttpServletResponse response)
return;
}
- McpTransportContext transportContext = this.contextExtractor.extract(request);
+ if (!HttpServletRequestUtils.isJsonContentType(request.getContentType())) {
+ this.responseError(response, HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE,
+ McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST)
+ .message("Unsupported Media Type: Content-Type must be application/json")
+ .build());
+ return;
+ }
String accept = request.getHeader(ACCEPT);
if (accept == null || !(accept.contains(APPLICATION_JSON) && accept.contains(TEXT_EVENT_STREAM))) {
@@ -179,6 +185,8 @@ protected void doPost(HttpServletRequest request, HttpServletResponse response)
return;
}
+ McpTransportContext transportContext = this.contextExtractor.extract(request);
+
try {
String body = HttpServletRequestUtils.readBody(request, this.requestMaxSize);
diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java
index 5f883cb98..63e149ac9 100644
--- a/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java
+++ b/mcp-core/src/main/java/io/modelcontextprotocol/server/transport/HttpServletStreamableServerTransportProvider.java
@@ -503,6 +503,14 @@ protected void doPost(HttpServletRequest request, HttpServletResponse response)
badRequestErrors.add("application/json required in Accept header");
}
+ if (!HttpServletRequestUtils.isJsonContentType(request.getContentType())) {
+ this.responseError(response, HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE,
+ McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST)
+ .message("Unsupported Media Type: Content-Type must be application/json")
+ .build());
+ return;
+ }
+
McpTransportContext transportContext = this.contextExtractor.extract(request);
try {
diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtilsTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtilsTests.java
index 48c997c04..022b0a7a5 100644
--- a/mcp-core/src/test/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtilsTests.java
+++ b/mcp-core/src/test/java/io/modelcontextprotocol/server/transport/HttpServletRequestUtilsTests.java
@@ -12,6 +12,9 @@
import jakarta.servlet.ServletInputStream;
import jakarta.servlet.http.HttpServletRequest;
import org.junit.jupiter.api.Test;
+import org.junit.jupiter.params.ParameterizedTest;
+import org.junit.jupiter.params.provider.NullAndEmptySource;
+import org.junit.jupiter.params.provider.ValueSource;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
@@ -100,6 +103,22 @@ void honorsRequestCharacterEncoding() throws Exception {
assertThat(body).isEqualTo("café");
}
+ @ParameterizedTest
+ @ValueSource(strings = { "application/json", "application/json; charset=utf-8", "application/json;charset=UTF-8",
+ "Application/JSON", " application/json ; charset=utf-8" })
+ void acceptsJsonContentType(String contentType) {
+ assertThat(HttpServletRequestUtils.isJsonContentType(contentType)).isTrue();
+ }
+
+ @ParameterizedTest
+ @NullAndEmptySource
+ @ValueSource(strings = { "text/plain", "text/plain;charset=UTF-8", "text/plain; a=application/json",
+ "application/x-www-form-urlencoded", "multipart/form-data", "application/json-seq", "application/jsonp",
+ "application/json, text/plain", "text/event-stream" })
+ void rejectsNonJsonContentType(String contentType) {
+ assertThat(HttpServletRequestUtils.isJsonContentType(contentType)).isFalse();
+ }
+
private static HttpServletRequest requestWithBody(String body, String characterEncoding) throws IOException {
HttpServletRequest request = mock(HttpServletRequest.class);
when(request.getInputStream()).thenReturn(servletInputStream(body.getBytes(StandardCharsets.UTF_8)));
diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletSseIntegrationTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletSseIntegrationTests.java
index c4e334f4e..72c54bccd 100644
--- a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletSseIntegrationTests.java
+++ b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletSseIntegrationTests.java
@@ -4,6 +4,9 @@
package io.modelcontextprotocol.server;
+import java.io.BufferedReader;
+import java.io.InputStream;
+import java.io.InputStreamReader;
import java.net.URI;
import java.net.http.HttpClient;
import java.net.http.HttpRequest;
@@ -12,6 +15,9 @@
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.Map;
+import java.util.concurrent.CompletableFuture;
+import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicBoolean;
import java.util.stream.Stream;
import io.modelcontextprotocol.AbstractMcpClientServerIntegrationTests;
@@ -22,6 +28,7 @@
import io.modelcontextprotocol.server.McpServer.SyncSpecification;
import io.modelcontextprotocol.server.transport.HttpServletSseServerTransportProvider;
import io.modelcontextprotocol.server.transport.TomcatTestUtil;
+import io.modelcontextprotocol.spec.McpSchema;
import jakarta.servlet.http.HttpServletRequest;
import jakarta.servlet.http.HttpServletResponse;
import org.apache.catalina.LifecycleException;
@@ -33,8 +40,12 @@
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.Timeout;
+import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
+import org.junit.jupiter.params.provider.ValueSource;
+import reactor.core.publisher.Mono;
+import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA;
import static org.assertj.core.api.Assertions.assertThat;
@Timeout(15)
@@ -194,6 +205,55 @@ public void cancel() {
assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_REQUEST_ENTITY_TOO_LARGE);
}
+ @ParameterizedTest
+ @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" })
+ void rejectsNonJsonContentType(String contentType) throws Exception {
+ var httpClient = HttpClient.newHttpClient();
+ var toolCalled = new AtomicBoolean();
+ prepareAsyncServerBuilder().capabilities(McpSchema.ServerCapabilities.builder().tools(true).build())
+ .tools(McpServerFeatures.AsyncToolSpecification.builder()
+ .tool(McpSchema.Tool.builder("tool1", EMPTY_JSON_SCHEMA).build())
+ .callHandler((exchange, request) -> {
+ toolCalled.set(true);
+ return Mono.just(McpSchema.CallToolResult.builder().build());
+ })
+ .build())
+ .build();
+
+ // Establish an SSE session to obtain a valid session ID
+ var sseRequest = HttpRequest.newBuilder()
+ .uri(URI.create("http://localhost:" + PORT + CUSTOM_SSE_ENDPOINT))
+ .header("Accept", "text/event-stream")
+ .GET()
+ .build();
+ HttpResponse