Skip to content

Commit 6a01d9b

Browse files
committed
Ensure errors are priorities over body limits
Also renamed ResponseSubscribers, refactored test code Signed-off-by: Dariusz Jędrzejczyk <dariusz.jedrzejczyk@broadcom.com>
1 parent f8271c2 commit 6a01d9b

10 files changed

Lines changed: 111 additions & 87 deletions

File tree

‎mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransport.java‎

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -396,17 +396,17 @@ public Mono<Void> connect(Function<Mono<JSONRPCMessage>, Mono<JSONRPCMessage>> h
396396
// The body is handed over as a publisher and nothing is read off
397397
// the wire until it is subscribed, so it has to be drained even
398398
// when its content is of no further interest.
399-
return ResponseSubscribers.drain(response.body(), this.maxResponseSize);
399+
return ResponseBodyHandlers.drain(response.body(), this.maxResponseSize);
400400
}
401401

402402
int statusCode = response.statusCode();
403403

404404
if (statusCode >= 200 && statusCode < 300) {
405-
Flux<String> lines = ResponseSubscribers.decodeLines(response.body(), this.maxResponseSize);
406-
return ResponseSubscribers.decodeSseResponse(lines, this.maxResponseSize);
405+
Flux<String> lines = ResponseBodyHandlers.decodeLines(response.body(), this.maxResponseSize);
406+
return ResponseBodyHandlers.decodeSseResponse(lines, this.maxResponseSize);
407407
}
408408
else {
409-
return ResponseSubscribers.drainThenError(response.body(), this.maxResponseSize,
409+
return ResponseBodyHandlers.drainThenError(response.body(), this.maxResponseSize,
410410
new RuntimeException("Failed to connect to SSE stream: " + statusCode));
411411
}
412412
})
@@ -537,10 +537,10 @@ private Mono<Void> sendHttpPost(final String endpoint, final String body) {
537537
.flatMap(response -> {
538538
int statusCode = response.statusCode();
539539
if (statusCode == 200 || statusCode == 201 || statusCode == 202 || statusCode == 206) {
540-
return ResponseSubscribers.drain(response.body(), this.maxResponseSize).then();
540+
return ResponseBodyHandlers.drain(response.body(), this.maxResponseSize).then();
541541
}
542-
return ResponseSubscribers.decodeAggregateResponse(response.body(), this.maxResponseSize)
543-
.flatMap(text -> Mono.error(new RuntimeException(
542+
return ResponseBodyHandlers.decodeAggregateResponse(response.body(), this.maxResponseSize)
543+
.flatMap(text -> Mono.error(new RuntimeException(
544544
"Sending message failed with a non-OK HTTP code: " + statusCode + " - " + text)));
545545
});
546546
});

‎mcp-core/src/main/java/io/modelcontextprotocol/client/transport/HttpClientStreamableHttpTransport.java‎

Lines changed: 15 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -245,7 +245,7 @@ private Publisher<Void> createDelete(String sessionId) {
245245
() -> this.httpClient.sendAsync(requestBuilder.build(), HttpResponse.BodyHandlers.ofPublisher()))
246246
// The response is not inspected, but the body still has to be consumed
247247
// to release the connection.
248-
.flatMapMany(response -> ResponseSubscribers.drain(response.body(), this.maxResponseSize))
248+
.flatMapMany(response -> ResponseBodyHandlers.drain(response.body(), this.maxResponseSize))
249249
.then())
250250
.then();
251251
}
@@ -286,8 +286,8 @@ public Mono<Void> closeGracefully() {
286286
private Flux<McpSchema.JSONRPCMessage> consumeSseStream(
287287
java.util.concurrent.Flow.Publisher<List<java.nio.ByteBuffer>> body,
288288
McpTransportStream<Disposable> existingStream, Runnable onFirstMessage) {
289-
Flux<String> lines = ResponseSubscribers.decodeLines(body, this.maxResponseSize);
290-
return ResponseSubscribers.decodeSseResponse(lines, this.maxResponseSize).flatMap(sseEvent -> {
289+
Flux<String> lines = ResponseBodyHandlers.decodeLines(body, this.maxResponseSize);
290+
return ResponseBodyHandlers.decodeSseResponse(lines, this.maxResponseSize).flatMap(sseEvent -> {
291291
if (!isMessageEvent(sseEvent.event())) {
292292
logger.debug("Received SSE event with type: {}", sseEvent);
293293
if (onFirstMessage != null) {
@@ -434,9 +434,9 @@ else if (statusCode >= 200 && statusCode < 300) {
434434

435435
return proceed ? consumeSseStream(httpResponse.body(), stream, null)
436436
: exception != null
437-
? ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
437+
? ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
438438
exception)
439-
: ResponseSubscribers.drain(httpResponse.body(), this.maxResponseSize);
439+
: ResponseBodyHandlers.drain(httpResponse.body(), this.maxResponseSize);
440440
});
441441
})
442442
.retryWhen(authorizationErrorRetrySpec())
@@ -555,7 +555,7 @@ public Mono<Void> sendMessage(McpSchema.JSONRPCMessage sentMessage) {
555555
var request = requestBuilder.build();
556556
var requestSnapshot = new HttpRequestSnapshot(request.uri(), request.method(),
557557
request.headers());
558-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
558+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
559559
new McpHttpClientTransportAuthorizationException(
560560
"Authorization error when sending message", requestSnapshot,
561561
toResponseInfo(httpResponse)));
@@ -580,7 +580,7 @@ public Mono<Void> sendMessage(McpSchema.JSONRPCMessage sentMessage) {
580580
if (contentType.isBlank() || "0".equals(contentLength) || statusCode == 202) {
581581
logger.debug("No body returned for POST in session {}", sessionRepresentation);
582582
deliveredSink.success();
583-
return ResponseSubscribers.drain(httpResponse.body(), this.maxResponseSize);
583+
return ResponseBodyHandlers.drain(httpResponse.body(), this.maxResponseSize);
584584
}
585585
else if (contentType.contains(TEXT_EVENT_STREAM)) {
586586
AtomicBoolean delivered = new AtomicBoolean();
@@ -591,7 +591,7 @@ else if (contentType.contains(TEXT_EVENT_STREAM)) {
591591
});
592592
}
593593
else if (contentType.contains(APPLICATION_JSON)) {
594-
return ResponseSubscribers
594+
return ResponseBodyHandlers
595595
.decodeAggregateResponse(httpResponse.body(), this.maxResponseSize)
596596
.flatMapMany(data -> {
597597
deliveredSink.success();
@@ -612,34 +612,34 @@ else if (contentType.contains(APPLICATION_JSON)) {
612612

613613
logger.warn("Unknown media type {} returned for POST in session {}", contentType,
614614
sessionRepresentation);
615-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
615+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
616616
new RuntimeException("Unknown media type returned: " + contentType));
617617
}
618618
else if (statusCode == NOT_FOUND) {
619619
if (maybeSessionId.isPresent()) {
620620
logger.debug("Session not found for session ID: {}", sessionRepresentation);
621-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
621+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
622622
new McpTransportSessionNotFoundException(
623623
"Session not found for session ID: " + sessionRepresentation));
624624
}
625-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
625+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
626626
new McpTransportException("Server Not Found. Status code:" + statusCode));
627627
}
628628
else if (statusCode == BAD_REQUEST) {
629629
if (maybeSessionId.isPresent()) {
630-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
630+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
631631
new McpTransportSessionNotFoundException(
632632
"Session not found for session ID: " + sessionRepresentation));
633633
}
634-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
634+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
635635
new McpTransportException("Bad Request. Status code:" + statusCode));
636636
}
637637
else if (statusCode >= 400 && statusCode < 500) {
638-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
638+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
639639
new McpTransportException("Invalid request. Status code: " + statusCode));
640640
}
641641

642-
return ResponseSubscribers.drainThenError(httpResponse.body(), this.maxResponseSize,
642+
return ResponseBodyHandlers.drainThenError(httpResponse.body(), this.maxResponseSize,
643643
new RuntimeException("Failed to send message, status code: " + statusCode));
644644
})
645645
.onErrorMap(CompletionException.class, Throwable::getCause))

mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java renamed to mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseBodyHandlers.java

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,7 @@
3636
* @author Dariusz Jędrzejczyk
3737
* @author Daniel Garnier-Moiroux
3838
*/
39-
class ResponseSubscribers {
39+
class ResponseBodyHandlers {
4040

4141
/**
4242
* Bytes of SSE field framing a single line may carry on top of the message payload:
@@ -137,7 +137,7 @@ static Mono<String> decodeAggregateResponse(Publisher<List<ByteBuffer>> publishe
137137
* @param error the error to propagate once the body has been discarded
138138
*/
139139
static <T> Flux<T> drainThenError(Publisher<List<ByteBuffer>> body, int maxSize, Throwable error) {
140-
return boundTotalBytes(body, maxSize).thenMany(Flux.error(error));
140+
return boundTotalBytes(body, maxSize).onErrorComplete().thenMany(Mono.error(error));
141141
}
142142

143143
/**

‎mcp-core/src/test/java/io/modelcontextprotocol/client/transport/LargeSseEventDecodingTests.java‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717
import reactor.adapter.JdkFlowAdapter;
1818
import reactor.core.publisher.Flux;
1919

20-
import io.modelcontextprotocol.client.transport.ResponseSubscribers.SseEvent;
20+
import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEvent;
2121

2222
import static org.assertj.core.api.Assertions.assertThat;
2323

@@ -48,7 +48,7 @@
4848
*
4949
* <p>
5050
* The middle column read each chunk incrementally, which is ~25x quicker than what 2.0.0
51-
* shipped, but {@link ResponseSubscribers.Utf8LineDecoder} still searched its buffered
51+
* shipped, but {@link ResponseBodyHandlers.Utf8LineDecoder} still searched its buffered
5252
* characters for a line terminator from the start of the buffer on every chunk, so eight
5353
* times the payload cost ~45x the time. Resuming that search where the previous one ended
5454
* gives the third column, which scales with the payload rather than with its square and
@@ -173,8 +173,8 @@ private static long timeDecode(byte[] body, int expectedEvents) {
173173
private static List<SseEvent> decode(byte[] body) {
174174
Flow.Publisher<List<ByteBuffer>> publisher = JdkFlowAdapter
175175
.publisherToFlowPublisher(Flux.fromIterable(chunk(body)));
176-
Flux<String> lines = ResponseSubscribers.decodeLines(publisher, Integer.MAX_VALUE);
177-
return ResponseSubscribers.decodeSseResponse(lines, MAX_SIZE).collectList().block();
176+
Flux<String> lines = ResponseBodyHandlers.decodeLines(publisher, Integer.MAX_VALUE);
177+
return ResponseBodyHandlers.decodeSseResponse(lines, MAX_SIZE).collectList().block();
178178
}
179179

180180
private static List<List<ByteBuffer>> chunk(byte[] body) {

‎mcp-core/src/test/java/io/modelcontextprotocol/client/transport/SseEventParserTests.java‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,8 +8,8 @@
88

99
import org.junit.jupiter.api.Test;
1010

11-
import io.modelcontextprotocol.client.transport.ResponseSubscribers.SseEvent;
12-
import io.modelcontextprotocol.client.transport.ResponseSubscribers.SseEventParser;
11+
import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEvent;
12+
import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.SseEventParser;
1313

1414
import static org.assertj.core.api.Assertions.assertThat;
1515

‎mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderBoundTests.java‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
import java.nio.charset.StandardCharsets;
99
import java.util.List;
1010

11-
import io.modelcontextprotocol.client.transport.ResponseSubscribers.Utf8LineDecoder;
11+
import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.Utf8LineDecoder;
1212
import io.modelcontextprotocol.spec.McpTransportException;
1313
import org.junit.jupiter.api.Test;
1414

‎mcp-core/src/test/java/io/modelcontextprotocol/client/transport/Utf8LineDecoderTests.java‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import java.nio.charset.StandardCharsets;
1010
import java.util.List;
1111

12-
import io.modelcontextprotocol.client.transport.ResponseSubscribers.Utf8LineDecoder;
12+
import io.modelcontextprotocol.client.transport.ResponseBodyHandlers.Utf8LineDecoder;
1313
import org.junit.jupiter.api.Test;
1414

1515
import static org.assertj.core.api.Assertions.assertThat;

‎mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientBoundedReadTestSupport.java‎

Lines changed: 57 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -8,14 +8,18 @@
88
import java.io.OutputStream;
99
import java.net.InetSocketAddress;
1010
import java.nio.charset.StandardCharsets;
11+
import java.time.Duration;
12+
import java.util.Arrays;
13+
import java.util.concurrent.CompletableFuture;
1114
import java.util.concurrent.Executors;
1215

13-
import com.sun.net.httpserver.HttpExchange;
1416
import com.sun.net.httpserver.HttpServer;
1517
import io.modelcontextprotocol.server.transport.TomcatTestUtil;
1618
import org.junit.jupiter.api.AfterEach;
1719
import org.junit.jupiter.api.BeforeEach;
1820

21+
import static org.assertj.core.api.Assertions.assertThat;
22+
1923
/**
2024
* Shared fixture for the transport bounded-read tests: a bare {@link HttpServer} whose
2125
* response body is written by a per-test {@link Responder}.
@@ -62,67 +66,79 @@ void stopServer() {
6266
}
6367

6468
/**
65-
* Registers a handler that responds with the given content type and body.
69+
* Registers a handler that answers {@code method} requests to {@code path} with the
70+
* given content type and a body written by {@code responder}. Any other method gets a
71+
* 405, as from a server offering nothing else there, so that requests the test does
72+
* not target, such as the GET stream the Streamable HTTP transport opens once
73+
* initialized, neither reach the responder nor stand in for the targeted request.
74+
* @return completes once the targeted response has been handled: with the
75+
* {@link IOException} that cut the body short if the client hung up first, or with
76+
* {@code null} if the body was written in full
6677
*/
67-
protected void respondWith(String path, String contentType, Responder responder) {
78+
protected CompletableFuture<IOException> respondWith(String method, String path, String contentType,
79+
Responder responder) {
80+
CompletableFuture<IOException> response = new CompletableFuture<>();
6881
this.server.createContext(path, exchange -> {
69-
exchange.getResponseHeaders().set("Content-Type", contentType);
70-
exchange.sendResponseHeaders(200, 0);
71-
try (OutputStream body = exchange.getResponseBody()) {
72-
responder.respond(body);
73-
}
74-
catch (IOException ignored) {
75-
// The client aborts the response once the limit is exceeded, which closes
76-
// the connection and makes further writes fail. That is the behaviour
77-
// under test.
82+
try {
83+
if (!method.equals(exchange.getRequestMethod())) {
84+
exchange.sendResponseHeaders(405, -1);
85+
return;
86+
}
87+
exchange.getResponseHeaders().set("Content-Type", contentType);
88+
exchange.sendResponseHeaders(200, 0);
89+
try (OutputStream body = exchange.getResponseBody()) {
90+
responder.respond(body);
91+
response.complete(null);
92+
}
93+
catch (IOException ex) {
94+
response.complete(ex);
95+
}
7896
}
7997
finally {
8098
exchange.close();
8199
}
82100
});
101+
return response;
83102
}
84103

85104
/**
86-
* A responder that writes {@code chunks} blocks of {@code 'a'} with no line
87-
* terminator anywhere, so nothing downstream can ever flush a line.
105+
* Asserts that the client hung up on {@code response}, which is how exceeding the
106+
* bound must end: with the endless responders below, a client that read the body in
107+
* full would never let it complete, and one that merely stopped reading would leave
108+
* the server blocked writing into a stalled connection.
88109
*/
89-
protected static Responder unterminatedLine(int chunks) {
90-
return body -> {
91-
byte[] chunk = new byte[MAX_SIZE];
92-
java.util.Arrays.fill(chunk, (byte) 'a');
93-
for (int i = 0; i < chunks; i++) {
94-
body.write(chunk);
95-
body.flush();
96-
}
97-
};
110+
protected static void assertHungUp(CompletableFuture<IOException> response) {
111+
assertThat(response).succeedsWithin(Duration.ofSeconds(5)).isNotNull();
98112
}
99113

100114
/**
101-
* A responder that writes {@code chunks} blocks of {@code 'a'}, each ending in a lone
102-
* CR. The line decoder only flushes a line on LF, so its buffer keeps growing even
103-
* though a CR arrives regularly.
115+
* A responder that streams {@code 'a'} with no line terminator anywhere, so nothing
116+
* downstream can ever flush a line.
104117
*/
105-
protected static Responder carriageReturnTerminatedRuns(int chunks) {
106-
return body -> {
107-
byte[] chunk = new byte[MAX_SIZE];
108-
java.util.Arrays.fill(chunk, (byte) 'a');
109-
chunk[MAX_SIZE - 1] = '\r';
110-
for (int i = 0; i < chunks; i++) {
111-
body.write(chunk);
112-
body.flush();
113-
}
114-
};
118+
protected static Responder unterminatedLine() {
119+
byte[] block = new byte[MAX_SIZE];
120+
Arrays.fill(block, (byte) 'a');
121+
return endlessly(block);
115122
}
116123

117124
/**
118-
* A responder that writes enough short, properly terminated lines to exceed the limit
119-
* in aggregate.
125+
* A responder that streams short, properly terminated lines, each small but exceeding
126+
* the limit in aggregate.
120127
*/
121128
protected static Responder manyShortLines(String prefix) {
129+
return endlessly((prefix + "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa\n").getBytes(StandardCharsets.UTF_8));
130+
}
131+
132+
/**
133+
* A responder that repeats {@code block} until the client hangs up, as a peer
134+
* streaming without end would. There is no amount to tune: however much the socket
135+
* buffers absorb, and whether the client closes with a FIN or a RST, the writes only
136+
* stop once the connection is gone.
137+
*/
138+
private static Responder endlessly(byte[] block) {
122139
return body -> {
123-
byte[] line = (prefix + "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa\n").getBytes(StandardCharsets.UTF_8);
124-
for (int i = 0; i < (MAX_SIZE / line.length) + 64; i++) {
125-
body.write(line);
140+
while (true) {
141+
body.write(block);
126142
body.flush();
127143
}
128144
};

‎mcp-test/src/test/java/io/modelcontextprotocol/client/transport/HttpClientSseClientTransportBoundedReadTests.java‎

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -48,38 +48,41 @@ void releaseStream() {
4848
void shouldRejectSingleLineExceedingMaxSize() {
4949
// A line that never terminates, so the line buffer underneath the SSE parser
5050
// would grow without limit before any event could be flushed.
51-
respondWith(endpoint(), "text/event-stream", unterminatedLine(8));
51+
var response = respondWith("GET", endpoint(), "text/event-stream", unterminatedLine());
5252

5353
StepVerifier.create(connect())
5454
.verifyErrorMatches(t -> messageContains(t, "Inbound line exceeds the maximum allowed size"));
55+
assertHungUp(response);
5556
}
5657

5758
@Test
5859
void shouldRejectEventExceedingMaxSize() {
5960
// Many short, terminated "data:" lines with no blank line to end the event. Each
6061
// line is small, but the accumulated event data would grow without limit.
61-
respondWith(endpoint(), "text/event-stream", manyShortLines("data:"));
62+
var response = respondWith("GET", endpoint(), "text/event-stream", manyShortLines("data:"));
6263

6364
StepVerifier.create(connect())
6465
.verifyErrorMatches(t -> messageContains(t, "Inbound SSE event exceeds the maximum allowed size"));
66+
assertHungUp(response);
6567
}
6668

6769
@Test
68-
void shouldRejectPostResponseExceedingMaxSize() throws Exception {
70+
void shouldRejectPostResponseExceedingMaxSize() {
6971
// The response to a posted message is discarded on success, but a peer must still
7072
// not be able to make the transport read an unbounded one.
71-
respondWith(endpoint(), "text/event-stream", body -> {
73+
respondWith("GET", endpoint(), "text/event-stream", body -> {
7274
body.write(("event:endpoint\ndata:" + MESSAGE_ENDPOINT + "\n\n").getBytes(StandardCharsets.UTF_8));
7375
body.flush();
7476
awaitTeardown();
7577
});
76-
respondWith(MESSAGE_ENDPOINT, "application/json", unterminatedLine(64));
78+
var response = respondWith("POST", MESSAGE_ENDPOINT, "application/json", unterminatedLine());
7779

7880
HttpClientSseClientTransport transport = transport();
7981
transport.connect(Function.identity()).block(Duration.ofSeconds(5));
8082

8183
StepVerifier.create(sendMessage(transport))
8284
.verifyErrorMatches(t -> messageContains(t, "Inbound response body exceeds the maximum allowed size"));
85+
assertHungUp(response);
8386
}
8487

8588
private void awaitTeardown() {

0 commit comments

Comments
 (0)