Cancel SSE publisher in case of async timeouts

Prior to this commit, the MVC `GraphQlSseHandler` would not react to
async request timeouts thrown by the Servlet container. This means that
when such timeouts happened, the SSE handler would still try to write to
the underlying response, whereas it was already recycled. This would
lead to NullPointerException thrown by the container.

This commit ensures that the SSE handler registers an async listener to
be notified of async timeouts and cancels the publisher as a result.
The SSE completion is not performed so as to let the client know that
the exchange did not complete and that it should re-subscribe.

Fixes gh-1067
This commit is contained in:
Brian Clozel
2024-10-16 21:49:16 +02:00
parent 2650eb5863
commit a210f77353
2 changed files with 39 additions and 6 deletions

View File

@@ -33,6 +33,7 @@ import org.springframework.graphql.server.WebGraphQlHandler;
import org.springframework.graphql.server.WebGraphQlResponse;
import org.springframework.util.AlternativeJdkIdGenerator;
import org.springframework.util.IdGenerator;
import org.springframework.web.context.request.async.AsyncRequestTimeoutException;
import org.springframework.web.servlet.function.ServerRequest;
import org.springframework.web.servlet.function.ServerResponse;
@@ -93,6 +94,12 @@ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler {
private SseSubscriber(ServerResponse.SseBuilder sseBuilder) {
this.sseBuilder = sseBuilder;
this.sseBuilder.onTimeout(this::onTimeout);
}
private void onTimeout() {
this.cancel();
this.sseBuilder.error(new AsyncRequestTimeoutException());
}
@Override
@@ -116,11 +123,11 @@ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler {
if (ex instanceof SubscriptionPublisherException spe) {
ExecutionResult result = ExecutionResult.newExecutionResult().errors(spe.getErrors()).build();
writeResult(result.toSpecification());
hookOnComplete();
}
else {
this.sseBuilder.error(ex);
}
hookOnComplete();
}
@Override

View File

@@ -24,6 +24,8 @@ import java.util.List;
import java.util.concurrent.atomic.AtomicBoolean;
import graphql.schema.DataFetcher;
import jakarta.servlet.AsyncEvent;
import jakarta.servlet.AsyncListener;
import jakarta.servlet.ServletException;
import jakarta.servlet.ServletOutputStream;
import jakarta.servlet.http.HttpServletResponse;
@@ -35,6 +37,7 @@ import org.springframework.graphql.GraphQlSetup;
import org.springframework.http.MediaType;
import org.springframework.http.converter.HttpMessageConverter;
import org.springframework.http.converter.json.MappingJackson2HttpMessageConverter;
import org.springframework.mock.web.MockAsyncContext;
import org.springframework.mock.web.MockHttpServletRequest;
import org.springframework.mock.web.MockHttpServletResponse;
import org.springframework.web.servlet.function.AsyncServerResponse;
@@ -72,7 +75,7 @@ class GraphQlSseHandlerTests {
void shouldRejectQueryOperations() throws Exception {
GraphQlSseHandler handler = createSseHandler(SEARCH_DATA_FETCHER);
MockHttpServletRequest request = createServletRequest("{ \"query\": \"{ bookById(id: 42) {name} }\"}");
MockHttpServletResponse response = handleRequest(request, handler);
MockHttpServletResponse response = handleAndAwait(request, handler);
assertThat(response.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE);
assertThat(response.getContentAsString()).isEqualTo("""
@@ -91,7 +94,7 @@ class GraphQlSseHandlerTests {
MockHttpServletRequest request = createServletRequest("""
{ "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" }
""");
MockHttpServletResponse response = handleRequest(request, handler);
MockHttpServletResponse response = handleAndAwait(request, handler);
assertThat(response.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE);
assertThat(response.getContentAsString()).isEqualTo("""
@@ -117,7 +120,7 @@ class GraphQlSseHandlerTests {
MockHttpServletRequest request = createServletRequest("""
{ "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" }
""");
MockHttpServletResponse response = handleRequest(request, handler);
MockHttpServletResponse response = handleAndAwait(request, handler);
assertThat(response.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE);
assertThat(response.getContentAsString()).isEqualTo("""
@@ -153,7 +156,26 @@ class GraphQlSseHandlerTests {
response.writeTo(servletRequest, servletResponse, new DefaultContext());
await().atMost(Duration.ofMillis(500)).until(DATA_FETCHER_CANCELLED::get);
}
@Test
void shouldCancelDataFetcherWhenAsyncTimeout() throws Exception {
DataFetcher<?> errorDataFetcher = env -> Flux.just(BookSource.getBook(1L))
.delayElements(Duration.ofMillis(500)).doOnCancel(() -> DATA_FETCHER_CANCELLED.set(true));
GraphQlSseHandler handler = createSseHandler(errorDataFetcher);
MockHttpServletRequest servletRequest = createServletRequest("""
{ "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" }
""");
MockHttpServletResponse servletResponse = handleRequest(servletRequest, handler);
for (AsyncListener listener : ((MockAsyncContext) servletRequest.getAsyncContext()).getListeners()) {
listener.onTimeout(new AsyncEvent(servletRequest.getAsyncContext()));
}
assertThat(DATA_FETCHER_CANCELLED.get()).isTrue();
assertThat(servletResponse.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE);
assertThat(servletResponse.getContentAsString()).isEmpty();
}
private GraphQlSseHandler createSseHandler(DataFetcher<?> dataFetcher) {
@@ -174,15 +196,19 @@ class GraphQlSseHandlerTests {
private MockHttpServletResponse handleRequest(
MockHttpServletRequest servletRequest, GraphQlSseHandler handler) throws ServletException, IOException {
ServerRequest request = ServerRequest.create(servletRequest, MESSAGE_READERS);
ServerResponse response = handler.handleRequest(request);
if (response instanceof AsyncServerResponse asyncResponse) {
asyncResponse.block();
}
MockHttpServletResponse servletResponse = new MockHttpServletResponse();
response.writeTo(servletRequest, servletResponse, new DefaultContext());
return servletResponse;
}
private MockHttpServletResponse handleAndAwait(
MockHttpServletRequest servletRequest, GraphQlSseHandler handler) throws ServletException, IOException {
MockHttpServletResponse servletResponse = handleRequest(servletRequest, handler);
await().atMost(Duration.ofMillis(500)).until(() -> servletResponse.getContentAsString().contains("complete"));
return servletResponse;
}