|
12 | 12 | import java.nio.charset.StandardCharsets; |
13 | 13 | import java.time.Duration; |
14 | 14 | import java.util.Map; |
| 15 | +import java.util.concurrent.CompletableFuture; |
| 16 | +import java.util.concurrent.TimeUnit; |
| 17 | +import java.util.concurrent.atomic.AtomicBoolean; |
15 | 18 | import java.util.stream.Stream; |
16 | 19 |
|
17 | 20 | import io.modelcontextprotocol.AbstractMcpClientServerIntegrationTests; |
|
22 | 25 | import io.modelcontextprotocol.server.McpServer.SyncSpecification; |
23 | 26 | import io.modelcontextprotocol.server.transport.HttpServletSseServerTransportProvider; |
24 | 27 | import io.modelcontextprotocol.server.transport.TomcatTestUtil; |
| 28 | +import io.modelcontextprotocol.spec.McpSchema; |
25 | 29 | import jakarta.servlet.http.HttpServletRequest; |
26 | 30 | import jakarta.servlet.http.HttpServletResponse; |
27 | 31 | import org.apache.catalina.LifecycleException; |
|
33 | 37 | import org.junit.jupiter.api.BeforeEach; |
34 | 38 | import org.junit.jupiter.api.Test; |
35 | 39 | import org.junit.jupiter.api.Timeout; |
| 40 | +import org.junit.jupiter.params.ParameterizedTest; |
36 | 41 | import org.junit.jupiter.params.provider.Arguments; |
| 42 | +import org.junit.jupiter.params.provider.ValueSource; |
| 43 | +import reactor.core.publisher.Mono; |
37 | 44 |
|
| 45 | +import static io.modelcontextprotocol.util.ToolsUtils.EMPTY_JSON_SCHEMA; |
38 | 46 | import static org.assertj.core.api.Assertions.assertThat; |
39 | 47 |
|
40 | 48 | @Timeout(15) |
@@ -125,36 +133,8 @@ protected void prepareClients(int port, String mcpEndpoint) { |
125 | 133 | void rejectsWhenBodyBytesExceedLimitWithoutContentLengthHeader() throws Exception { |
126 | 134 | var httpClient = HttpClient.newHttpClient(); |
127 | 135 |
|
128 | | - // Establish an SSE session to obtain a valid session ID |
129 | 136 | prepareAsyncServerBuilder().build(); |
130 | | - var sseRequest = HttpRequest.newBuilder() |
131 | | - .uri(URI.create("http://localhost:" + PORT + CUSTOM_SSE_ENDPOINT)) |
132 | | - .header("Accept", "text/event-stream") |
133 | | - .GET() |
134 | | - .build(); |
135 | | - var sseResponseRef = new java.util.concurrent.atomic.AtomicReference<HttpResponse<java.io.InputStream>>(); |
136 | | - var sessionIdFuture = new java.util.concurrent.CompletableFuture<String>(); |
137 | | - httpClient.sendAsync(sseRequest, HttpResponse.BodyHandlers.ofInputStream()).thenAccept(response -> { |
138 | | - sseResponseRef.set(response); |
139 | | - try (var reader = new java.io.BufferedReader( |
140 | | - new java.io.InputStreamReader(response.body(), StandardCharsets.UTF_8))) { |
141 | | - String line; |
142 | | - while ((line = reader.readLine()) != null) { |
143 | | - if (line.startsWith("data:") && line.contains("sessionId=")) { |
144 | | - String data = line.substring("data:".length()).strip(); |
145 | | - String sessionId = data.substring(data.indexOf("sessionId=") + "sessionId=".length()); |
146 | | - sessionIdFuture.complete(sessionId); |
147 | | - return; |
148 | | - } |
149 | | - } |
150 | | - sessionIdFuture.completeExceptionally(new RuntimeException("sessionId not found in SSE stream")); |
151 | | - response.body().close(); |
152 | | - } |
153 | | - catch (Exception e) { |
154 | | - sessionIdFuture.completeExceptionally(e); |
155 | | - } |
156 | | - }); |
157 | | - String sessionId = sessionIdFuture.get(5, java.util.concurrent.TimeUnit.SECONDS); |
| 137 | + String sessionId = openSseSession(httpClient); |
158 | 138 |
|
159 | 139 | // Send POST request with an over-sized body |
160 | 140 | byte[] oversizedBody = "a".repeat(MAX_REQUEST_SIZE + 1).getBytes(StandardCharsets.UTF_8); |
@@ -194,6 +174,59 @@ public void cancel() { |
194 | 174 | assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_REQUEST_ENTITY_TOO_LARGE); |
195 | 175 | } |
196 | 176 |
|
| 177 | + @ParameterizedTest |
| 178 | + @ValueSource(strings = { "text/plain;charset=UTF-8", "application/x-www-form-urlencoded", "multipart/form-data" }) |
| 179 | + void rejectsNonJsonContentType(String contentType) throws Exception { |
| 180 | + var httpClient = HttpClient.newHttpClient(); |
| 181 | + var toolCalled = new AtomicBoolean(); |
| 182 | + prepareAsyncServerBuilder().capabilities(McpSchema.ServerCapabilities.builder().tools(true).build()) |
| 183 | + .tools(McpServerFeatures.AsyncToolSpecification.builder() |
| 184 | + .tool(McpSchema.Tool.builder("tool1", EMPTY_JSON_SCHEMA).build()) |
| 185 | + .callHandler((exchange, request) -> { |
| 186 | + toolCalled.set(true); |
| 187 | + return Mono.just(McpSchema.CallToolResult.builder().build()); |
| 188 | + }) |
| 189 | + .build()) |
| 190 | + .build(); |
| 191 | + String sessionId = openSseSession(httpClient); |
| 192 | + |
| 193 | + // CORS-safelisted content types can be sent cross-origin by a browser without a |
| 194 | + // preflight, so they must be rejected before the message is handled |
| 195 | + var request = HttpRequest.newBuilder() |
| 196 | + .uri(URI.create("http://localhost:" + PORT + CUSTOM_MESSAGE_ENDPOINT + "?sessionId=" + sessionId)) |
| 197 | + .header("Content-Type", contentType) |
| 198 | + .POST(HttpRequest.BodyPublishers.ofString(""" |
| 199 | + {"jsonrpc":"2.0","id":"call-1","method":"tools/call","params":{"name":"tool1","arguments":{}}}""")) |
| 200 | + .build(); |
| 201 | + |
| 202 | + var response = httpClient.send(request, HttpResponse.BodyHandlers.ofString()); |
| 203 | + |
| 204 | + assertThat(response.statusCode()).isEqualTo(HttpServletResponse.SC_UNSUPPORTED_MEDIA_TYPE); |
| 205 | + assertThat(response.body()).contains("Unsupported Media Type: Content-Type must be application/json"); |
| 206 | + assertThat(toolCalled).isFalse(); |
| 207 | + } |
| 208 | + |
| 209 | + /** |
| 210 | + * Opens an SSE connection and returns the session ID from the endpoint event. The |
| 211 | + * connection stays open, keeping the session alive, until the server closes it. |
| 212 | + */ |
| 213 | + private String openSseSession(HttpClient httpClient) throws Exception { |
| 214 | + var sseRequest = HttpRequest.newBuilder() |
| 215 | + .uri(URI.create("http://localhost:" + PORT + CUSTOM_SSE_ENDPOINT)) |
| 216 | + .header("Accept", "text/event-stream") |
| 217 | + .GET() |
| 218 | + .build(); |
| 219 | + var sessionIdFuture = new CompletableFuture<String>(); |
| 220 | + httpClient.sendAsync(sseRequest, HttpResponse.BodyHandlers.ofLines()) |
| 221 | + .thenAccept(response -> response.body().forEach(line -> { |
| 222 | + int index = line.indexOf("sessionId="); |
| 223 | + if (line.startsWith("data:") && index != -1) { |
| 224 | + sessionIdFuture.complete(line.substring(index + "sessionId=".length()).strip()); |
| 225 | + } |
| 226 | + })); |
| 227 | + return sessionIdFuture.get(5, TimeUnit.SECONDS); |
| 228 | + } |
| 229 | + |
197 | 230 | static McpTransportContextExtractor<HttpServletRequest> TEST_CONTEXT_EXTRACTOR = (r) -> McpTransportContext |
198 | 231 | .create(Map.of("important", "value")); |
199 | 232 |
|
|
0 commit comments