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 721bdf195..82bcc68d6 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 @@ -23,6 +23,8 @@ */ final class HttpServletRequestUtils { + private static final String APPLICATION_JSON = "application/json"; + private HttpServletRequestUtils() { } @@ -41,6 +43,27 @@ static Map> extractHeaders(HttpServletRequest request) { return headers; } + /** + * 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 c2ea6dd4e..2a3909c7c 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 @@ -354,6 +354,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) { @@ -452,6 +460,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 e61375bff..291a15b93 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 @@ -170,7 +170,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 059acda21..c9dbb2439 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 @@ -427,6 +427,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-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxServerRequestUtils.java b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxServerRequestUtils.java new file mode 100644 index 000000000..74e333153 --- /dev/null +++ b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxServerRequestUtils.java @@ -0,0 +1,42 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.server.transport; + +import org.springframework.http.InvalidMediaTypeException; +import org.springframework.http.MediaType; +import org.springframework.web.reactive.function.server.ServerRequest; + +/** + * Utility methods for working with {@link ServerRequest}. For internal use only. + * + * @author Daniel Garnier-Moiroux + */ +final class WebFluxServerRequestUtils { + + private WebFluxServerRequestUtils() { + } + + /** + * Checks whether the request's {@code Content-Type} header denotes + * {@code application/json}. Only the media type is compared, case-insensitively; + * parameters such as {@code charset} are ignored. A missing or malformed header 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 request The incoming server request + * @return {@code true} if the media type is {@code application/json} + */ + static boolean isJsonContentType(ServerRequest request) { + try { + return request.headers().contentType().map(MediaType.APPLICATION_JSON::equalsTypeAndSubtype).orElse(false); + } + catch (InvalidMediaTypeException ex) { + return false; + } + } + +} diff --git a/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxSseServerTransportProvider.java b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxSseServerTransportProvider.java index e950417d4..ad52e43d2 100644 --- a/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxSseServerTransportProvider.java +++ b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxSseServerTransportProvider.java @@ -359,6 +359,13 @@ private Mono handleMessage(ServerRequest request) { return ServerResponse.status(e.getStatusCode()).bodyValue(e.getMessage()); } + if (!WebFluxServerRequestUtils.isJsonContentType(request)) { + return ServerResponse.status(HttpStatus.UNSUPPORTED_MEDIA_TYPE) + .bodyValue(McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + } + if (request.queryParam("sessionId").isEmpty()) { return ServerResponse.badRequest().bodyValue(new McpError("Session ID missing in message endpoint")); } diff --git a/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStatelessServerTransport.java b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStatelessServerTransport.java index bbb0493e4..a64bcdc12 100644 --- a/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStatelessServerTransport.java +++ b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStatelessServerTransport.java @@ -114,6 +114,13 @@ private Mono handlePost(ServerRequest request) { return ServerResponse.status(e.getStatusCode()).bodyValue(e.getMessage()); } + if (!WebFluxServerRequestUtils.isJsonContentType(request)) { + return ServerResponse.status(HttpStatus.UNSUPPORTED_MEDIA_TYPE) + .bodyValue(McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + } + McpTransportContext transportContext = this.contextExtractor.extract(request); List acceptHeaders = request.headers().asHttpHeaders().getAccept(); diff --git a/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStreamableServerTransportProvider.java b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStreamableServerTransportProvider.java index 223c2f009..85e355968 100644 --- a/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStreamableServerTransportProvider.java +++ b/mcp-spring/mcp-spring-webflux/src/main/java/io/modelcontextprotocol/server/transport/WebFluxStreamableServerTransportProvider.java @@ -246,6 +246,13 @@ private Mono handlePost(ServerRequest request) { return ServerResponse.status(e.getStatusCode()).bodyValue(e.getMessage()); } + if (!WebFluxServerRequestUtils.isJsonContentType(request)) { + return ServerResponse.status(HttpStatus.UNSUPPORTED_MEDIA_TYPE) + .bodyValue(McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + } + McpTransportContext transportContext = this.contextExtractor.extract(request); List acceptHeaders = request.headers().asHttpHeaders().getAccept(); diff --git a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxSseIntegrationTests.java b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxSseIntegrationTests.java index eb8abb90c..1b72db7a6 100644 --- a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxSseIntegrationTests.java +++ b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxSseIntegrationTests.java @@ -5,16 +5,25 @@ package io.modelcontextprotocol; import java.time.Duration; +import java.util.List; import java.util.Map; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; 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 org.springframework.core.ParameterizedTypeReference; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.codec.ServerSentEvent; import org.springframework.http.server.reactive.HttpHandler; import org.springframework.http.server.reactive.ReactorHttpHandlerAdapter; +import org.springframework.web.reactive.function.client.ClientResponse; import org.springframework.web.reactive.function.client.WebClient; import org.springframework.web.reactive.function.server.RouterFunctions; import org.springframework.web.reactive.function.server.ServerRequest; @@ -26,12 +35,20 @@ import io.modelcontextprotocol.server.McpServer; import io.modelcontextprotocol.server.McpServer.AsyncSpecification; import io.modelcontextprotocol.server.McpServer.SingleSessionSyncSpecification; +import io.modelcontextprotocol.server.McpServerFeatures; import io.modelcontextprotocol.server.McpTransportContextExtractor; import io.modelcontextprotocol.server.TestUtil; import io.modelcontextprotocol.server.transport.WebFluxSseServerTransportProvider; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.core.Disposable; +import reactor.core.publisher.Mono; +import reactor.core.publisher.Sinks; import reactor.netty.DisposableServer; import reactor.netty.http.server.HttpServer; +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; +import static org.assertj.core.api.Assertions.assertThat; + @Timeout(15) class WebFluxSseIntegrationTests extends AbstractMcpClientServerIntegrationTests { @@ -103,4 +120,69 @@ public void after() { } } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsNonJsonContentType(String contentType) { + var webClient = WebClient.create("http://localhost:" + PORT); + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(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 the session-scoped message endpoint. The + // SSE stream must stay open for the session to remain active. + var messageEndpoint = Sinks.one(); + Disposable sseSubscription = webClient.get() + .uri(CUSTOM_SSE_ENDPOINT) + .accept(MediaType.TEXT_EVENT_STREAM) + .retrieve() + .bodyToFlux(new ParameterizedTypeReference>() { + }) + .filter(event -> WebFluxSseServerTransportProvider.ENDPOINT_EVENT_TYPE.equals(event.event())) + .subscribe(event -> messageEndpoint.tryEmitValue(event.data())); + try { + String endpoint = messageEndpoint.asMono().block(Duration.ofSeconds(5)); + + // Initialize the session, so that the tool call below would be handled if it + // were accepted + for (String message : List.of(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2024-11-05","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""", """ + {"jsonrpc":"2.0","method":"notifications/initialized"}""")) { + var initResponse = webClient.post() + .uri(endpoint) + .contentType(MediaType.APPLICATION_JSON) + .bodyValue(message) + .exchangeToMono(ClientResponse::toBodilessEntity) + .block(); + assertThat(initResponse.getStatusCode()).isEqualTo(HttpStatus.OK); + } + + // CORS-safelisted content types can be sent cross-origin by a browser without + // a preflight, so they must be rejected before the message is handled + var response = webClient.post() + .uri(endpoint) + .contentType(MediaType.parseMediaType(contentType)) + .bodyValue( + """ + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + finally { + sseSubscription.dispose(); + } + } + } diff --git a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStatelessIntegrationTests.java b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStatelessIntegrationTests.java index 96a786a9e..ebb8f6c3f 100644 --- a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStatelessIntegrationTests.java +++ b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStatelessIntegrationTests.java @@ -5,13 +5,18 @@ package io.modelcontextprotocol; import java.time.Duration; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; 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 org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; import org.springframework.http.server.reactive.HttpHandler; import org.springframework.http.server.reactive.ReactorHttpHandlerAdapter; import org.springframework.web.reactive.function.client.WebClient; @@ -22,11 +27,17 @@ import io.modelcontextprotocol.server.McpServer; import io.modelcontextprotocol.server.McpServer.StatelessAsyncSpecification; import io.modelcontextprotocol.server.McpServer.StatelessSyncSpecification; +import io.modelcontextprotocol.server.McpStatelessServerFeatures; import io.modelcontextprotocol.server.TestUtil; import io.modelcontextprotocol.server.transport.WebFluxStatelessServerTransport; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.core.publisher.Mono; import reactor.netty.DisposableServer; import reactor.netty.http.server.HttpServer; +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; +import static org.assertj.core.api.Assertions.assertThat; + @Timeout(15) class WebFluxStatelessIntegrationTests extends AbstractStatelessIntegrationTests { @@ -88,4 +99,35 @@ public void after() { } } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsNonJsonContentType(String contentType) { + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpStatelessServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(EMPTY_JSON_SCHEMA).build()) + .callHandler((transportContext, request) -> { + toolCalled.set(true); + return Mono.just(McpSchema.CallToolResult.builder().build()); + }) + .build()) + .build(); + + // CORS-safelisted content types can be sent cross-origin by a browser without a + // preflight, so they must be rejected before the message is handled + var response = WebClient.create("http://localhost:" + PORT) + .post() + .uri(CUSTOM_MESSAGE_ENDPOINT) + .contentType(MediaType.parseMediaType(contentType)) + .accept(MediaType.APPLICATION_JSON, MediaType.TEXT_EVENT_STREAM) + .bodyValue(""" + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + } diff --git a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStreamableIntegrationTests.java b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStreamableIntegrationTests.java index 5ab651931..18d81671b 100644 --- a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStreamableIntegrationTests.java +++ b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/WebFluxStreamableIntegrationTests.java @@ -6,13 +6,19 @@ import java.time.Duration; import java.util.Map; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import org.junit.jupiter.api.AfterEach; 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 org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; import org.springframework.http.server.reactive.HttpHandler; import org.springframework.http.server.reactive.ReactorHttpHandlerAdapter; import org.springframework.web.reactive.function.client.WebClient; @@ -26,12 +32,19 @@ import io.modelcontextprotocol.server.McpServer; import io.modelcontextprotocol.server.McpServer.AsyncSpecification; import io.modelcontextprotocol.server.McpServer.SyncSpecification; +import io.modelcontextprotocol.server.McpServerFeatures; import io.modelcontextprotocol.server.McpTransportContextExtractor; import io.modelcontextprotocol.server.TestUtil; import io.modelcontextprotocol.server.transport.WebFluxStreamableServerTransportProvider; +import io.modelcontextprotocol.spec.HttpHeaders; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.core.publisher.Mono; import reactor.netty.DisposableServer; import reactor.netty.http.server.HttpServer; +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; +import static org.assertj.core.api.Assertions.assertThat; + @Timeout(15) class WebFluxStreamableIntegrationTests extends AbstractMcpClientServerIntegrationTests { @@ -100,4 +113,76 @@ public void after() { } } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsInitializeWithNonJsonContentType(String contentType) { + var webClient = WebClient.create("http://localhost:" + PORT); + prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0").build(); + + // CORS-safelisted content types can be sent cross-origin by a browser without a + // preflight, so they must be rejected before a session is created + var response = webClient.post() + .uri(CUSTOM_MESSAGE_ENDPOINT) + .contentType(MediaType.parseMediaType(contentType)) + .accept(MediaType.TEXT_EVENT_STREAM, MediaType.APPLICATION_JSON) + .bodyValue(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2025-06-18","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(response.getHeaders().containsKey(HttpHeaders.MCP_SESSION_ID)).isFalse(); + } + + @Test + void rejectsToolCallWithNonJsonContentType() { + var webClient = WebClient.create("http://localhost:" + PORT); + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0") + .capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(EMPTY_JSON_SCHEMA).build()) + .callHandler((exchange, request) -> { + toolCalled.set(true); + return Mono.just(McpSchema.CallToolResult.builder().build()); + }) + .build()) + .build(); + var sessionId = initializeSession(webClient); + + var response = webClient.post() + .uri(CUSTOM_MESSAGE_ENDPOINT) + .contentType(MediaType.parseMediaType("text/plain;charset=UTF-8")) + .accept(MediaType.TEXT_EVENT_STREAM, MediaType.APPLICATION_JSON) + .header(HttpHeaders.MCP_SESSION_ID, sessionId) + .bodyValue(""" + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + + private String initializeSession(WebClient webClient) { + var response = webClient.post() + .uri(CUSTOM_MESSAGE_ENDPOINT) + .contentType(MediaType.APPLICATION_JSON) + .accept(MediaType.TEXT_EVENT_STREAM, MediaType.APPLICATION_JSON) + .bodyValue(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2025-06-18","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + var sessionId = response.getHeaders().getFirst(HttpHeaders.MCP_SESSION_ID); + assertThat(sessionId).isNotNull(); + return sessionId; + } + } diff --git a/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/server/transport/WebFluxServerRequestUtilsTests.java b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/server/transport/WebFluxServerRequestUtilsTests.java new file mode 100644 index 000000000..1b383b7db --- /dev/null +++ b/mcp-spring/mcp-spring-webflux/src/test/java/io/modelcontextprotocol/server/transport/WebFluxServerRequestUtilsTests.java @@ -0,0 +1,42 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.server.transport; + +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.NullAndEmptySource; +import org.junit.jupiter.params.provider.ValueSource; + +import org.springframework.mock.web.reactive.function.server.MockServerRequest; +import org.springframework.web.reactive.function.server.ServerRequest; + +import static org.assertj.core.api.Assertions.assertThat; + +class WebFluxServerRequestUtilsTests { + + @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(WebFluxServerRequestUtils.isJsonContentType(requestWithContentType(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", "not a media type" }) + void rejectsNonJsonContentType(String contentType) { + assertThat(WebFluxServerRequestUtils.isJsonContentType(requestWithContentType(contentType))).isFalse(); + } + + private static ServerRequest requestWithContentType(String contentType) { + MockServerRequest.Builder builder = MockServerRequest.builder(); + if (contentType != null) { + builder.header("Content-Type", contentType); + } + return builder.build(); + } + +} diff --git a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcServerRequestUtils.java b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcServerRequestUtils.java new file mode 100644 index 000000000..1661ba3fa --- /dev/null +++ b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcServerRequestUtils.java @@ -0,0 +1,42 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.server.transport; + +import org.springframework.http.InvalidMediaTypeException; +import org.springframework.http.MediaType; +import org.springframework.web.servlet.function.ServerRequest; + +/** + * Utility methods for working with {@link ServerRequest}. For internal use only. + * + * @author Daniel Garnier-Moiroux + */ +final class WebMvcServerRequestUtils { + + private WebMvcServerRequestUtils() { + } + + /** + * Checks whether the request's {@code Content-Type} header denotes + * {@code application/json}. Only the media type is compared, case-insensitively; + * parameters such as {@code charset} are ignored. A missing or malformed header 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 request The incoming server request + * @return {@code true} if the media type is {@code application/json} + */ + static boolean isJsonContentType(ServerRequest request) { + try { + return request.headers().contentType().map(MediaType.APPLICATION_JSON::equalsTypeAndSubtype).orElse(false); + } + catch (InvalidMediaTypeException ex) { + return false; + } + } + +} diff --git a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcSseServerTransportProvider.java b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcSseServerTransportProvider.java index e1eb67311..64724e655 100644 --- a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcSseServerTransportProvider.java +++ b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcSseServerTransportProvider.java @@ -340,6 +340,13 @@ private ServerResponse handleMessage(ServerRequest request) { return ServerResponse.status(e.getStatusCode()).body(e.getMessage()); } + if (!WebMvcServerRequestUtils.isJsonContentType(request)) { + return ServerResponse.status(HttpStatus.UNSUPPORTED_MEDIA_TYPE) + .body(McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + } + if (request.param(SESSION_ID).isEmpty()) { return ServerResponse.badRequest().body(new McpError("Session ID missing in message endpoint")); } diff --git a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStatelessServerTransport.java b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStatelessServerTransport.java index 2c379192c..ef0209029 100644 --- a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStatelessServerTransport.java +++ b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStatelessServerTransport.java @@ -118,6 +118,13 @@ private ServerResponse handlePost(ServerRequest request) { return ServerResponse.status(e.getStatusCode()).body(e.getMessage()); } + if (!WebMvcServerRequestUtils.isJsonContentType(request)) { + return ServerResponse.status(HttpStatus.UNSUPPORTED_MEDIA_TYPE) + .body(McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + } + McpTransportContext transportContext = this.contextExtractor.extract(request); List acceptHeaders = request.headers().asHttpHeaders().getAccept(); diff --git a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStreamableServerTransportProvider.java b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStreamableServerTransportProvider.java index 4f701a9db..f282d54b6 100644 --- a/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStreamableServerTransportProvider.java +++ b/mcp-spring/mcp-spring-webmvc/src/main/java/io/modelcontextprotocol/server/transport/WebMvcStreamableServerTransportProvider.java @@ -343,6 +343,13 @@ private ServerResponse handlePost(ServerRequest request) { return ServerResponse.status(e.getStatusCode()).body(e.getMessage()); } + if (!WebMvcServerRequestUtils.isJsonContentType(request)) { + return ServerResponse.status(HttpStatus.UNSUPPORTED_MEDIA_TYPE) + .body(McpError.builder(McpSchema.ErrorCodes.INVALID_REQUEST) + .message("Unsupported Media Type: Content-Type must be application/json") + .build()); + } + List acceptHeaders = request.headers().asHttpHeaders().getAccept(); if (!acceptHeaders.contains(MediaType.TEXT_EVENT_STREAM) || !acceptHeaders.contains(MediaType.APPLICATION_JSON)) { diff --git a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcSseIntegrationTests.java b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcSseIntegrationTests.java index 045f9b3dd..ef0be7707 100644 --- a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcSseIntegrationTests.java +++ b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcSseIntegrationTests.java @@ -3,10 +3,13 @@ */ package io.modelcontextprotocol.server; +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; import static org.assertj.core.api.Assertions.assertThat; import java.time.Duration; +import java.util.List; import java.util.Map; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import org.apache.catalina.LifecycleException; @@ -14,10 +17,17 @@ import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; 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 org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.core.ParameterizedTypeReference; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.codec.ServerSentEvent; +import org.springframework.web.reactive.function.client.ClientResponse; import org.springframework.web.reactive.function.client.WebClient; import org.springframework.web.servlet.config.annotation.EnableWebMvc; import org.springframework.web.servlet.function.RouterFunction; @@ -32,6 +42,10 @@ import io.modelcontextprotocol.server.McpServer.AsyncSpecification; import io.modelcontextprotocol.server.McpServer.SingleSessionSyncSpecification; import io.modelcontextprotocol.server.transport.WebMvcSseServerTransportProvider; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.core.Disposable; +import reactor.core.publisher.Mono; +import reactor.core.publisher.Sinks; import reactor.core.scheduler.Schedulers; @Timeout(15) @@ -134,4 +148,69 @@ protected SingleSessionSyncSpecification prepareSyncServerBuilder() { return McpServer.sync(mcpServerTransportProvider); } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsNonJsonContentType(String contentType) { + var webClient = WebClient.create("http://localhost:" + PORT); + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(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 the session-scoped message endpoint. The + // SSE stream must stay open for the session to remain active. + var messageEndpoint = Sinks.one(); + Disposable sseSubscription = webClient.get() + .uri("/sse") + .accept(MediaType.TEXT_EVENT_STREAM) + .retrieve() + .bodyToFlux(new ParameterizedTypeReference>() { + }) + .filter(event -> WebMvcSseServerTransportProvider.ENDPOINT_EVENT_TYPE.equals(event.event())) + .subscribe(event -> messageEndpoint.tryEmitValue(event.data())); + try { + String endpoint = messageEndpoint.asMono().block(Duration.ofSeconds(5)); + + // Initialize the session, so that the tool call below would be handled if it + // were accepted + for (String message : List.of(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2024-11-05","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""", """ + {"jsonrpc":"2.0","method":"notifications/initialized"}""")) { + var initResponse = webClient.post() + .uri(endpoint) + .contentType(MediaType.APPLICATION_JSON) + .bodyValue(message) + .exchangeToMono(ClientResponse::toBodilessEntity) + .block(); + assertThat(initResponse.getStatusCode()).isEqualTo(HttpStatus.OK); + } + + // CORS-safelisted content types can be sent cross-origin by a browser without + // a preflight, so they must be rejected before the message is handled + var response = webClient.post() + .uri(endpoint) + .contentType(MediaType.parseMediaType(contentType)) + .bodyValue( + """ + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + finally { + sseSubscription.dispose(); + } + } + } diff --git a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStatelessIntegrationTests.java b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStatelessIntegrationTests.java index 8c7b0a85e..faf2b5e63 100644 --- a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStatelessIntegrationTests.java +++ b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStatelessIntegrationTests.java @@ -3,9 +3,11 @@ */ package io.modelcontextprotocol.server; +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; import static org.assertj.core.api.Assertions.assertThat; import java.time.Duration; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import org.apache.catalina.LifecycleException; @@ -13,10 +15,14 @@ import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; 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 org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; import org.springframework.web.reactive.function.client.WebClient; import org.springframework.web.servlet.config.annotation.EnableWebMvc; import org.springframework.web.servlet.function.RouterFunction; @@ -29,6 +35,8 @@ import io.modelcontextprotocol.server.McpServer.StatelessAsyncSpecification; import io.modelcontextprotocol.server.McpServer.StatelessSyncSpecification; import io.modelcontextprotocol.server.transport.WebMvcStatelessServerTransport; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.core.publisher.Mono; import reactor.core.scheduler.Schedulers; @Timeout(15) @@ -131,4 +139,35 @@ public void after() { } } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsNonJsonContentType(String contentType) { + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpStatelessServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(EMPTY_JSON_SCHEMA).build()) + .callHandler((transportContext, request) -> { + toolCalled.set(true); + return Mono.just(McpSchema.CallToolResult.builder().build()); + }) + .build()) + .build(); + + // CORS-safelisted content types can be sent cross-origin by a browser without a + // preflight, so they must be rejected before the message is handled + var response = WebClient.create("http://localhost:" + PORT) + .post() + .uri(MESSAGE_ENDPOINT) + .contentType(MediaType.parseMediaType(contentType)) + .accept(MediaType.APPLICATION_JSON, MediaType.TEXT_EVENT_STREAM) + .bodyValue(""" + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + } diff --git a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStreamableIntegrationTests.java b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStreamableIntegrationTests.java index cb7b4a2a0..2b1807122 100644 --- a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStreamableIntegrationTests.java +++ b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/WebMvcStreamableIntegrationTests.java @@ -3,21 +3,28 @@ */ package io.modelcontextprotocol.server; +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; import static org.assertj.core.api.Assertions.assertThat; import java.time.Duration; import java.util.Map; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import org.apache.catalina.LifecycleException; import org.apache.catalina.LifecycleState; import org.junit.jupiter.api.AfterEach; 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 org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; import org.springframework.web.reactive.function.client.WebClient; import org.springframework.web.servlet.function.ServerRequest; import org.springframework.web.servlet.config.annotation.EnableWebMvc; @@ -32,6 +39,9 @@ import io.modelcontextprotocol.server.McpServer.AsyncSpecification; import io.modelcontextprotocol.server.McpServer.SyncSpecification; import io.modelcontextprotocol.server.transport.WebMvcStreamableServerTransportProvider; +import io.modelcontextprotocol.spec.HttpHeaders; +import io.modelcontextprotocol.spec.McpSchema; +import reactor.core.publisher.Mono; import reactor.core.scheduler.Schedulers; @Timeout(15) @@ -150,4 +160,76 @@ protected void prepareClients(int port, String mcpEndpoint) { .requestTimeout(Duration.ofHours(10))); } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsInitializeWithNonJsonContentType(String contentType) { + var webClient = WebClient.create("http://localhost:" + PORT); + prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0").build(); + + // CORS-safelisted content types can be sent cross-origin by a browser without a + // preflight, so they must be rejected before a session is created + var response = webClient.post() + .uri(MESSAGE_ENDPOINT) + .contentType(MediaType.parseMediaType(contentType)) + .accept(MediaType.TEXT_EVENT_STREAM, MediaType.APPLICATION_JSON) + .bodyValue(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2025-06-18","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(response.getHeaders().containsKey(HttpHeaders.MCP_SESSION_ID)).isFalse(); + } + + @Test + void rejectsToolCallWithNonJsonContentType() { + var webClient = WebClient.create("http://localhost:" + PORT); + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0") + .capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(EMPTY_JSON_SCHEMA).build()) + .callHandler((exchange, request) -> { + toolCalled.set(true); + return Mono.just(McpSchema.CallToolResult.builder().build()); + }) + .build()) + .build(); + var sessionId = initializeSession(webClient); + + var response = webClient.post() + .uri(MESSAGE_ENDPOINT) + .contentType(MediaType.parseMediaType("text/plain;charset=UTF-8")) + .accept(MediaType.TEXT_EVENT_STREAM, MediaType.APPLICATION_JSON) + .header(HttpHeaders.MCP_SESSION_ID, sessionId) + .bodyValue(""" + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.UNSUPPORTED_MEDIA_TYPE); + assertThat(response.getBody()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + + private String initializeSession(WebClient webClient) { + var response = webClient.post() + .uri(MESSAGE_ENDPOINT) + .contentType(MediaType.APPLICATION_JSON) + .accept(MediaType.TEXT_EVENT_STREAM, MediaType.APPLICATION_JSON) + .bodyValue(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2025-06-18","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""") + .exchangeToMono(clientResponse -> clientResponse.toEntity(String.class)) + .block(); + assertThat(response.getStatusCode()).isEqualTo(HttpStatus.OK); + var sessionId = response.getHeaders().getFirst(HttpHeaders.MCP_SESSION_ID); + assertThat(sessionId).isNotNull(); + return sessionId; + } + } diff --git a/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/transport/WebMvcServerRequestUtilsTests.java b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/transport/WebMvcServerRequestUtilsTests.java new file mode 100644 index 000000000..7d474f1ce --- /dev/null +++ b/mcp-spring/mcp-spring-webmvc/src/test/java/io/modelcontextprotocol/server/transport/WebMvcServerRequestUtilsTests.java @@ -0,0 +1,45 @@ +/* + * Copyright 2026-2026 the original author or authors. + */ + +package io.modelcontextprotocol.server.transport; + +import java.util.List; + +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.NullAndEmptySource; +import org.junit.jupiter.params.provider.ValueSource; + +import org.springframework.http.converter.StringHttpMessageConverter; +import org.springframework.mock.web.MockHttpServletRequest; +import org.springframework.web.servlet.function.ServerRequest; + +import static org.assertj.core.api.Assertions.assertThat; + +class WebMvcServerRequestUtilsTests { + + @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(WebMvcServerRequestUtils.isJsonContentType(requestWithContentType(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", "not a media type" }) + void rejectsNonJsonContentType(String contentType) { + assertThat(WebMvcServerRequestUtils.isJsonContentType(requestWithContentType(contentType))).isFalse(); + } + + private static ServerRequest requestWithContentType(String contentType) { + MockHttpServletRequest servletRequest = new MockHttpServletRequest("POST", "/mcp"); + if (contentType != null) { + servletRequest.addHeader("Content-Type", contentType); + } + return ServerRequest.create(servletRequest, List.of(new StringHttpMessageConverter())); + } + +} 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 2640ce19e..b886255d1 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; @@ -35,6 +41,8 @@ import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; import org.junit.jupiter.params.provider.MethodSource; +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; @@ -224,6 +232,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().name("tool1").inputSchema(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 sseResponse = httpClient.send(sseRequest, HttpResponse.BodyHandlers.ofInputStream()); + try (var reader = new BufferedReader(new InputStreamReader(sseResponse.body(), StandardCharsets.UTF_8))) { + var sessionIdFuture = CompletableFuture.supplyAsync(() -> reader.lines() + .filter(line -> line.startsWith("data:") && line.contains("sessionId=")) + .map(line -> line.substring(line.indexOf("sessionId=") + "sessionId=".length()).strip()) + .findFirst() + .orElseThrow(() -> new IllegalStateException("sessionId not found in SSE stream"))); + String sessionId = sessionIdFuture.get(5, TimeUnit.SECONDS); + + // CORS-safelisted content types can be sent cross-origin by a browser without + // a + // preflight, so they must be rejected before the message is handled + var request = HttpRequest.newBuilder() + .uri(URI.create("http://localhost:" + PORT + CUSTOM_MESSAGE_ENDPOINT + "?sessionId=" + sessionId)) + .header("Content-Type", contentType) + .POST(HttpRequest.BodyPublishers.ofString( + """ + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""")) + .build(); + + var response = httpClient.send(request, HttpResponse.BodyHandlers.ofString()); + + assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE); + assertThat(response.body()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + } + static McpTransportContextExtractor TEST_CONTEXT_EXTRACTOR = (r) -> McpTransportContext .create(Map.of("important", "value")); diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStatelessIntegrationTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStatelessIntegrationTests.java index ce74e43f3..3562ed0e2 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStatelessIntegrationTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStatelessIntegrationTests.java @@ -14,6 +14,7 @@ import java.util.List; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; import java.util.function.BiFunction; @@ -787,6 +788,45 @@ public void cancel() { } } + @ParameterizedTest + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) + void rejectsNonJsonContentType(String contentType) throws Exception { + AtomicBoolean toolCalled = new AtomicBoolean(); + var mcpServer = McpServer.sync(mcpStatelessServerTransport) + .capabilities(ServerCapabilities.builder().tools(false).build()) + .tools(McpStatelessServerFeatures.SyncToolSpecification.builder() + .tool(Tool.builder().name("tool1").inputSchema(EMPTY_JSON_SCHEMA).build()) + .callHandler((transportContext, request) -> { + toolCalled.set(true); + return CallToolResult.builder().build(); + }) + .build()) + .build(); + + try { + // CORS-safelisted content types can be sent cross-origin by a browser without + // a preflight, so they must be rejected before the message is handled + var request = HttpRequest.newBuilder() + .uri(URI.create("http://localhost:" + PORT + CUSTOM_MESSAGE_ENDPOINT)) + .header("Content-Type", contentType) + .header("Accept", APPLICATION_JSON + ", " + TEXT_EVENT_STREAM) + .POST(HttpRequest.BodyPublishers.ofString( + """ + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""")) + .build(); + + var response = HttpClient.newHttpClient().send(request, HttpResponse.BodyHandlers.ofString()); + + assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE); + assertThatJson(response.body()).inPath("message") + .isEqualTo("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + finally { + mcpServer.closeGracefully(); + } + } + private double evaluateExpression(String expression) { // Simple expression evaluator for testing return switch (expression) { diff --git a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java index 60ad658d4..6cb532ffc 100644 --- a/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java +++ b/mcp-test/src/test/java/io/modelcontextprotocol/server/HttpServletStreamableIntegrationTests.java @@ -4,6 +4,7 @@ package io.modelcontextprotocol.server; +import java.io.IOException; import java.net.URI; import java.net.http.HttpClient; import java.net.http.HttpRequest; @@ -12,6 +13,7 @@ import java.nio.charset.StandardCharsets; import java.time.Duration; import java.util.Map; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Stream; import io.modelcontextprotocol.AbstractMcpClientServerIntegrationTests; @@ -22,19 +24,22 @@ import io.modelcontextprotocol.server.McpServer.SyncSpecification; import io.modelcontextprotocol.server.transport.HttpServletStreamableServerTransportProvider; import io.modelcontextprotocol.server.transport.TomcatTestUtil; +import io.modelcontextprotocol.spec.HttpHeaders; import io.modelcontextprotocol.spec.McpSchema; import jakarta.servlet.http.HttpServletRequest; import jakarta.servlet.http.HttpServletResponse; import org.apache.catalina.LifecycleException; import org.apache.catalina.LifecycleState; import org.apache.catalina.startup.Tomcat; -import org.junit.Test; import org.junit.jupiter.api.AfterEach; 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.MethodSource; +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; @@ -192,4 +197,83 @@ 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 rejectsInitializeWithNonJsonContentType(String contentType) throws Exception { + var httpClient = HttpClient.newHttpClient(); + prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0").build(); + + // CORS-safelisted content types can be sent cross-origin by a browser without a + // preflight, so they must be rejected before a session is created + var initialize = HttpRequest.newBuilder() + .uri(URI.create("http://localhost:" + PORT + MESSAGE_ENDPOINT)) + .header("Content-Type", contentType) + .header("Accept", "text/event-stream, application/json") + .POST(HttpRequest.BodyPublishers.ofString(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2025-06-18","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""")) + .build(); + + var response = httpClient.send(initialize, HttpResponse.BodyHandlers.ofString()); + + assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE); + assertThat(response.body()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(response.headers().firstValue(HttpHeaders.MCP_SESSION_ID)).isEmpty(); + } + + @Test + void rejectsToolCallWithNonJsonContentType() throws Exception { + var httpClient = HttpClient.newHttpClient(); + var toolCalled = new AtomicBoolean(); + prepareAsyncServerBuilder().serverInfo("test-server", "1.0.0") + .capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) + .tools(McpServerFeatures.AsyncToolSpecification.builder() + .tool(McpSchema.Tool.builder().name("tool1").inputSchema(EMPTY_JSON_SCHEMA).build()) + .callHandler((exchange, request) -> { + toolCalled.set(true); + return Mono.just(McpSchema.CallToolResult.builder().build()); + }) + .build()) + .build(); + var sessionId = initializeSession(httpClient); + + var toolCall = HttpRequest.newBuilder() + .uri(URI.create("http://localhost:" + PORT + MESSAGE_ENDPOINT)) + .header("Content-Type", "text/plain;charset=UTF-8") + .header("Accept", "text/event-stream, application/json") + .header(HttpHeaders.MCP_SESSION_ID, sessionId) + .POST(HttpRequest.BodyPublishers.ofString(""" + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""")) + .build(); + + var response = httpClient.send(toolCall, HttpResponse.BodyHandlers.ofString()); + + assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE); + assertThat(response.body()).contains("Unsupported Media Type: Content-Type must be application/json"); + assertThat(toolCalled).isFalse(); + } + + private String initializeSession(HttpClient httpClient) { + var initialize = HttpRequest.newBuilder() + .uri(URI.create("http://localhost:" + PORT + MESSAGE_ENDPOINT)) + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream, application/json") + .POST(HttpRequest.BodyPublishers.ofString(""" + {"jsonrpc":"2.0","id":"init","method":"initialize","params":{ + "protocolVersion":"2025-06-18","capabilities":{}, + "clientInfo":{"name":"test-client","version":"1.0.0"}}}""")) + .build(); + + HttpResponse response = null; + try { + response = httpClient.send(initialize, HttpResponse.BodyHandlers.ofString()); + } + catch (IOException | InterruptedException e) { + return null; + } + assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_OK); + return response.headers().firstValue(HttpHeaders.MCP_SESSION_ID).orElse(null); + } + }