From a210f7735388e7b824701c6ba62d50ddfc722a5b Mon Sep 17 00:00:00 2001 From: Brian Clozel Date: Wed, 16 Oct 2024 21:49:16 +0200 Subject: [PATCH] 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 --- .../server/webmvc/GraphQlSseHandler.java | 9 ++++- .../server/webmvc/GraphQlSseHandlerTests.java | 36 ++++++++++++++++--- 2 files changed, 39 insertions(+), 6 deletions(-) diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlSseHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlSseHandler.java index 957dcc7b..361e2a3d 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlSseHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlSseHandler.java @@ -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 diff --git a/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlSseHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlSseHandlerTests.java index 42e810db..f870df10 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlSseHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlSseHandlerTests.java @@ -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; }