diff --git a/spring-graphql/src/main/java/org/springframework/graphql/execution/SubscriptionPublisherException.java b/spring-graphql/src/main/java/org/springframework/graphql/execution/SubscriptionPublisherException.java index 851f7d9b..75a40b3d 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/execution/SubscriptionPublisherException.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/execution/SubscriptionPublisherException.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2022 the original author or authors. + * Copyright 2002-2024 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -17,10 +17,13 @@ package org.springframework.graphql.execution; import java.util.List; +import java.util.Map; +import graphql.ExecutionResult; import graphql.GraphQLError; import org.springframework.core.NestedRuntimeException; +import org.springframework.lang.Nullable; /** * An exception raised after a GraphQL subscription @@ -47,7 +50,7 @@ public final class SubscriptionPublisherException extends NestedRuntimeException * @param errors the list of resolved GraphQL errors * @param cause the original exception */ - public SubscriptionPublisherException(List errors, Throwable cause) { + public SubscriptionPublisherException(List errors, @Nullable Throwable cause) { super("GraphQL subscription ended with error(s): " + errors, cause); this.errors = errors; } @@ -62,4 +65,12 @@ public final class SubscriptionPublisherException extends NestedRuntimeException return this.errors; } + /** + * Return an {@link ExecutionResult} specification map with the GraphQL errors. + * @since 1.3.0 + */ + public Map toMap() { + return ExecutionResult.newExecutionResult().errors(this.errors).build().toSpecification(); + } + } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlHttpHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlHttpHandler.java index 099e6f1e..16d846ef 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlHttpHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlHttpHandler.java @@ -16,7 +16,6 @@ package org.springframework.graphql.server.webflux; -import java.util.Arrays; import java.util.List; import org.apache.commons.logging.Log; @@ -25,6 +24,7 @@ import reactor.core.publisher.Mono; import org.springframework.graphql.server.WebGraphQlHandler; import org.springframework.graphql.server.WebGraphQlRequest; +import org.springframework.graphql.server.WebGraphQlResponse; import org.springframework.http.MediaType; import org.springframework.http.codec.CodecConfigurer; import org.springframework.web.reactive.function.server.ServerRequest; @@ -42,8 +42,8 @@ public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { private static final Log logger = LogFactory.getLog(GraphQlHttpHandler.class); @SuppressWarnings("removal") - private static final List SUPPORTED_MEDIA_TYPES = - Arrays.asList(MediaType.APPLICATION_GRAPHQL_RESPONSE, MediaType.APPLICATION_JSON, MediaType.APPLICATION_GRAPHQL); + private static final List SUPPORTED_MEDIA_TYPES = List.of( + MediaType.APPLICATION_GRAPHQL_RESPONSE, MediaType.APPLICATION_JSON, MediaType.APPLICATION_GRAPHQL); /** @@ -66,18 +66,18 @@ public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { /** * Handle GraphQL requests over HTTP. - * @param serverRequest the incoming HTTP request + * @param request the incoming HTTP request * @return the HTTP response */ - public Mono handleRequest(ServerRequest serverRequest) { - return readRequest(serverRequest) + public Mono handleRequest(ServerRequest request) { + return readRequest(request) .flatMap((body) -> { WebGraphQlRequest graphQlRequest = new WebGraphQlRequest( - serverRequest.uri(), serverRequest.headers().asHttpHeaders(), - serverRequest.cookies(), serverRequest.remoteAddress().orElse(null), - serverRequest.attributes(), body, - serverRequest.exchange().getRequest().getId(), - serverRequest.exchange().getLocaleContext().getLocale()); + request.uri(), request.headers().asHttpHeaders(), + request.cookies(), request.remoteAddress().orElse(null), + request.attributes(), body, + request.exchange().getRequest().getId(), + request.exchange().getLocaleContext().getLocale()); if (logger.isDebugEnabled()) { logger.debug("Executing: " + graphQlRequest); } @@ -85,20 +85,20 @@ public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { }) .flatMap((response) -> { if (logger.isDebugEnabled()) { - logger.debug("Execution complete"); - } - ServerResponse.BodyBuilder builder = ServerResponse.ok(); - builder.headers((headers) -> headers.putAll(response.getResponseHeaders())); - builder.contentType(selectResponseMediaType(serverRequest)); - if (this.codecDelegate != null) { - return builder.bodyValue(this.codecDelegate.encode(response)); - } - else { - return builder.bodyValue(response.toMap()); + logger.debug("Execution result ready"); } + return prepareResponse(request, response); }); } + protected Mono prepareResponse(ServerRequest serverRequest, WebGraphQlResponse response) { + ServerResponse.BodyBuilder builder = ServerResponse.ok(); + builder.headers((headers) -> headers.putAll(response.getResponseHeaders())); + builder.contentType(selectResponseMediaType(serverRequest)); + return builder.bodyValue((this.codecDelegate != null) ? + this.codecDelegate.encode(response) : response.toMap()); + } + private static MediaType selectResponseMediaType(ServerRequest serverRequest) { for (MediaType accepted : serverRequest.headers().accept()) { if (SUPPORTED_MEDIA_TYPES.contains(accepted)) { diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlSseHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlSseHandler.java index d656bd44..fff26edc 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlSseHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlSseHandler.java @@ -18,6 +18,7 @@ package org.springframework.graphql.server.webflux; import java.util.Collections; +import java.util.List; import java.util.Map; import graphql.ErrorType; @@ -29,6 +30,7 @@ import org.reactivestreams.Publisher; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import org.springframework.graphql.ResponseError; import org.springframework.graphql.execution.SubscriptionPublisherException; import org.springframework.graphql.server.WebGraphQlHandler; import org.springframework.graphql.server.WebGraphQlRequest; @@ -46,13 +48,15 @@ import org.springframework.web.reactive.function.server.ServerResponse; * {@link org.springframework.web.reactive.function.server.RouterFunctions}. * * @author Brian Clozel + * @author Rossen Stoyanchev * @since 1.3.0 */ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler { private static final Log logger = LogFactory.getLog(GraphQlSseHandler.class); - private static final Mono>> COMPLETE_EVENT = Mono.just(ServerSentEvent.>builder(Collections.emptyMap()).event("complete").build()); + private static final Mono>> COMPLETE_EVENT = Mono.just( + ServerSentEvent.>builder(Collections.emptyMap()).event("complete").build()); public GraphQlSseHandler(WebGraphQlHandler graphQlHandler) { @@ -66,7 +70,7 @@ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler { */ @SuppressWarnings("unchecked") public Mono handleRequest(ServerRequest serverRequest) { - Flux>> data = readRequest(serverRequest) + return readRequest(serverRequest) .flatMap((body) -> { WebGraphQlRequest graphQlRequest = new WebGraphQlRequest( serverRequest.uri(), serverRequest.headers().asHttpHeaders(), @@ -79,35 +83,39 @@ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler { } return this.graphQlHandler.handleRequest(graphQlRequest); }) - .flatMapMany((response) -> { + .flatMap((response) -> { if (logger.isDebugEnabled()) { - logger.debug("Execution result ready" - + (!CollectionUtils.isEmpty(response.getErrors()) ? " with errors: " + response.getErrors() : "") - + "."); + List errors = response.getErrors(); + logger.debug("Execution result " + + (!CollectionUtils.isEmpty(errors) ? "has errors: " + errors : "is ready") + "."); } + Flux> resultFlux; if (response.getData() instanceof Publisher) { - // Subscription - return Flux.from((Publisher) response.getData()).map(ExecutionResult::toSpecification); + resultFlux = Flux.from((Publisher) response.getData()) + .map(ExecutionResult::toSpecification) + .onErrorResume(SubscriptionPublisherException.class, (ex) -> Mono.just(ex.toMap())); } - if (logger.isDebugEnabled()) { - logger.debug("Only subscriptions are supported, DataFetcher must return a Publisher type"); + else { + if (logger.isDebugEnabled()) { + logger.debug("A subscription DataFetcher must return a Publisher: " + response.getData()); + } + resultFlux = Flux.just(ExecutionResult.newExecutionResult() + .addError(GraphQLError.newError() + .errorType(ErrorType.OperationNotSupported) + .message("SSE handler supports only subscriptions") + .build()) + .build() + .toSpecification()); } - // Single response (query or mutation) are not supported - String errorMessage = "SSE transport only supports Subscription operations"; - GraphQLError unsupportedOperationError = GraphQLError.newError().errorType(ErrorType.OperationNotSupported) - .message(errorMessage).build(); - return Flux.error(new SubscriptionPublisherException(Collections.singletonList(unsupportedOperationError), - new IllegalArgumentException(errorMessage))); - }) - .onErrorResume(SubscriptionPublisherException.class, (exc) -> { - ExecutionResult errorResult = ExecutionResult.newExecutionResult().errors(exc.getErrors()).build(); - return Flux.just(errorResult.toSpecification()); - }) - .map((event) -> ServerSentEvent.builder(event).event("next").build()); - Flux>> body = data.concatWith(COMPLETE_EVENT); - return ServerResponse.ok().contentType(MediaType.TEXT_EVENT_STREAM).body(BodyInserters.fromServerSentEvents(body)) - .onErrorResume(Throwable.class, (exc) -> ServerResponse.badRequest().build()); + Flux>> sseFlux = + resultFlux.map((event) -> ServerSentEvent.builder(event).event("next").build()); + + return ServerResponse.ok() + .contentType(MediaType.TEXT_EVENT_STREAM) + .body(BodyInserters.fromServerSentEvents(sseFlux.concatWith(COMPLETE_EVENT))) + .onErrorResume(Throwable.class, (ex) -> ServerResponse.badRequest().build()); + }); } } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlHttpHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlHttpHandler.java index 2604c8c4..6dfd8353 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlHttpHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlHttpHandler.java @@ -16,18 +16,16 @@ package org.springframework.graphql.server.webmvc; -import java.util.Arrays; import java.util.List; +import java.util.Map; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutionException; import jakarta.servlet.ServletException; import org.springframework.context.i18n.LocaleContextHolder; -import org.springframework.graphql.GraphQlResponse; import org.springframework.graphql.server.WebGraphQlHandler; import org.springframework.graphql.server.WebGraphQlRequest; -import org.springframework.http.HttpMethod; import org.springframework.http.MediaType; import org.springframework.http.converter.HttpMessageConverter; import org.springframework.http.server.ServletServerHttpResponse; @@ -47,8 +45,8 @@ import org.springframework.web.servlet.function.ServerResponse; public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { @SuppressWarnings("removal") - private static final List SUPPORTED_MEDIA_TYPES = - Arrays.asList(MediaType.APPLICATION_GRAPHQL_RESPONSE, MediaType.APPLICATION_JSON, MediaType.APPLICATION_GRAPHQL); + private static final List SUPPORTED_MEDIA_TYPES = List.of( + MediaType.APPLICATION_GRAPHQL_RESPONSE, MediaType.APPLICATION_JSON, MediaType.APPLICATION_GRAPHQL); /** @@ -62,27 +60,28 @@ public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { /** * Create a new instance with a custom message converter. *

If no converter is provided, this will use - * {@link org.springframework.web.servlet.config.annotation.WebMvcConfigurer#configureMessageConverters(List) the one configured in the web framework}. + * {@link org.springframework.web.servlet.config.annotation.WebMvcConfigurer#configureMessageConverters(List) + * the one configured in the web framework}. * @param graphQlHandler common handler for GraphQL over HTTP requests - * @param messageConverter custom {@link HttpMessageConverter} to be used for encoding and decoding GraphQL payloads + * @param converter custom {@link HttpMessageConverter} to read and write GraphQL payloads */ - public GraphQlHttpHandler(WebGraphQlHandler graphQlHandler, @Nullable HttpMessageConverter messageConverter) { - super(graphQlHandler, messageConverter); + public GraphQlHttpHandler(WebGraphQlHandler graphQlHandler, @Nullable HttpMessageConverter converter) { + super(graphQlHandler, converter); } /** * Handle GraphQL requests over HTTP. - * @param serverRequest the incoming HTTP request + * @param request the incoming HTTP request * @return the HTTP response * @throws ServletException may be raised when reading the request body, e.g. * {@link HttpMediaTypeNotSupportedException}. */ - public ServerResponse handleRequest(ServerRequest serverRequest) throws ServletException { + public ServerResponse handleRequest(ServerRequest request) throws ServletException { WebGraphQlRequest graphQlRequest = new WebGraphQlRequest( - serverRequest.uri(), serverRequest.headers().asHttpHeaders(), initCookies(serverRequest), - serverRequest.remoteAddress().orElse(null), - serverRequest.attributes(), readBody(serverRequest), this.idGenerator.generateId().toString(), + request.uri(), request.headers().asHttpHeaders(), initCookies(request), + request.remoteAddress().orElse(null), + request.attributes(), readBody(request), this.idGenerator.generateId().toString(), LocaleContextHolder.getLocale()); if (logger.isDebugEnabled()) { @@ -94,13 +93,13 @@ public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { if (logger.isDebugEnabled()) { logger.debug("Execution complete"); } - MediaType contentType = selectResponseMediaType(serverRequest); + MediaType contentType = selectResponseMediaType(request); ServerResponse.BodyBuilder builder = ServerResponse.ok(); builder.headers((headers) -> headers.putAll(response.getResponseHeaders())); builder.contentType(contentType); if (this.messageConverter != null) { - return builder.build(writeFunction(contentType, response)); + return builder.build(writeFunction(this.messageConverter, contentType, response.toMap())); } else { return builder.body(response.toMap()); @@ -132,16 +131,13 @@ public class GraphQlHttpHandler extends AbstractGraphQlHttpHandler { return MediaType.APPLICATION_JSON; } - private ServerResponse.HeadersBuilder.WriteFunction writeFunction(MediaType contentType, GraphQlResponse response) { + private static ServerResponse.HeadersBuilder.WriteFunction writeFunction( + HttpMessageConverter converter, MediaType contentType, Map resultMap) { + return (servletRequest, servletResponse) -> { - if (messageConverter != null) { - ServletServerHttpResponse httpResponse = new ServletServerHttpResponse(servletResponse); - messageConverter.write(response.toMap(), contentType, httpResponse); - return null; - } - else { - throw new HttpMediaTypeNotSupportedException(contentType, SUPPORTED_MEDIA_TYPES, HttpMethod.POST); - } + ServletServerHttpResponse httpResponse = new ServletServerHttpResponse(servletResponse); + converter.write(resultMap, contentType, httpResponse); + return null; }; } 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 c21649e7..c21eea74 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 @@ -17,8 +17,9 @@ package org.springframework.graphql.server.webmvc; import java.io.IOException; -import java.util.Collections; +import java.util.List; import java.util.Map; +import java.util.function.Consumer; import graphql.ErrorType; import graphql.ExecutionResult; @@ -29,10 +30,10 @@ import reactor.core.publisher.BaseSubscriber; import reactor.core.publisher.Flux; import org.springframework.context.i18n.LocaleContextHolder; +import org.springframework.graphql.ResponseError; import org.springframework.graphql.execution.SubscriptionPublisherException; import org.springframework.graphql.server.WebGraphQlHandler; import org.springframework.graphql.server.WebGraphQlRequest; -import org.springframework.graphql.server.WebGraphQlResponse; import org.springframework.util.AlternativeJdkIdGenerator; import org.springframework.util.CollectionUtils; import org.springframework.util.IdGenerator; @@ -47,6 +48,7 @@ import org.springframework.web.servlet.function.ServerResponse; * {@link org.springframework.web.servlet.function.RouterFunctions}. * * @author Brian Clozel + * @author Rossen Stoyanchev * @since 1.3.0 */ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler { @@ -60,89 +62,91 @@ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler { /** * Handle GraphQL requests over HTTP using the Server-Sent Events protocol. - * @param serverRequest the incoming HTTP request + * @param request the incoming HTTP request * @return the HTTP response * @throws ServletException may be raised when reading the request body, e.g. * {@link HttpMediaTypeNotSupportedException}. */ - public ServerResponse handleRequest(ServerRequest serverRequest) throws ServletException { + @SuppressWarnings("unchecked") + public ServerResponse handleRequest(ServerRequest request) throws ServletException { WebGraphQlRequest graphQlRequest = new WebGraphQlRequest( - serverRequest.uri(), serverRequest.headers().asHttpHeaders(), initCookies(serverRequest), - serverRequest.remoteAddress().orElse(null), serverRequest.attributes(), - readBody(serverRequest), this.idGenerator.generateId().toString(), + request.uri(), request.headers().asHttpHeaders(), initCookies(request), + request.remoteAddress().orElse(null), request.attributes(), + readBody(request), this.idGenerator.generateId().toString(), LocaleContextHolder.getLocale()); if (logger.isDebugEnabled()) { logger.debug("Executing: " + graphQlRequest); } - return ServerResponse.sse((sseBuilder) -> this.graphQlHandler.handleRequest(graphQlRequest) - .flatMapMany(this::handleResponse) - .subscribe(new SendMessageSubscriber(graphQlRequest.getId(), sseBuilder))); + + Flux> resultFlux = this.graphQlHandler.handleRequest(graphQlRequest) + .flatMapMany((response) -> { + if (logger.isDebugEnabled()) { + List errors = response.getErrors(); + logger.debug("Execution result " + + (!CollectionUtils.isEmpty(errors) ? "has errors: " + errors : "is ready") + "."); + } + if (response.getData() instanceof Publisher) { + return Flux.from((Publisher) response.getData()) + .map(ExecutionResult::toSpecification); + } + else { + if (logger.isDebugEnabled()) { + logger.debug("A subscription DataFetcher must return a Publisher: " + response.getData()); + } + return Flux.just(ExecutionResult.newExecutionResult() + .addError(GraphQLError.newError() + .errorType(ErrorType.OperationNotSupported) + .message("SSE handler supports only subscriptions") + .build()) + .build() + .toSpecification()); + } + }); + + return ServerResponse.sse(SseSubscriber.connect(resultFlux)); } - @SuppressWarnings("unchecked") - private Publisher> handleResponse(WebGraphQlResponse response) { - if (logger.isDebugEnabled()) { - logger.debug("Execution result ready" - + (!CollectionUtils.isEmpty(response.getErrors()) ? " with errors: " + response.getErrors() : "") - + "."); - } - if (response.getData() instanceof Publisher) { - // Subscription - return Flux.from((Publisher) response.getData()).map(ExecutionResult::toSpecification); - } - if (logger.isDebugEnabled()) { - logger.debug("Only subscriptions are supported, DataFetcher must return a Publisher type"); - } - // Single response (query or mutation) are not supported - String errorMessage = "SSE transport only supports Subscription operations"; - GraphQLError unsupportedOperationError = GraphQLError.newError().errorType(ErrorType.OperationNotSupported) - .message(errorMessage).build(); - return Flux.error(new SubscriptionPublisherException(Collections.singletonList(unsupportedOperationError), - new IllegalArgumentException(errorMessage))); - } + /** + * {@link org.reactivestreams.Subscriber} that writes to {@link ServerResponse.SseBuilder}. + */ + private static final class SseSubscriber extends BaseSubscriber> { + private final ServerResponse.SseBuilder sseBuilder; - private static class SendMessageSubscriber extends BaseSubscriber> { - - final String id; - - final ServerResponse.SseBuilder sseBuilder; - - SendMessageSubscriber(String id, ServerResponse.SseBuilder sseBuilder) { - this.id = id; + private SseSubscriber(ServerResponse.SseBuilder sseBuilder) { this.sseBuilder = sseBuilder; } @Override protected void hookOnNext(Map value) { - writeNext(value); + writeResult(value); } - @Override - protected void hookOnError(Throwable throwable) { - if (throwable instanceof SubscriptionPublisherException subscriptionException) { - ExecutionResult errorResult = ExecutionResult.newExecutionResult().errors(subscriptionException.getErrors()).build(); - writeNext(errorResult.toSpecification()); - } - else { - this.sseBuilder.error(throwable); - } - this.hookOnComplete(); - } - - private void writeNext(Map value) { + private void writeResult(Map value) { try { this.sseBuilder.event("next"); this.sseBuilder.data(value); } catch (IOException exception) { - this.onError(exception); + onError(exception); } } + @Override + protected void hookOnError(Throwable ex) { + if (ex instanceof SubscriptionPublisherException spe) { + ExecutionResult result = ExecutionResult.newExecutionResult().errors(spe.getErrors()).build(); + writeResult(result.toSpecification()); + } + else { + this.sseBuilder.error(ex); + } + hookOnComplete(); + } + @Override protected void hookOnComplete() { try { @@ -154,6 +158,12 @@ public class GraphQlSseHandler extends AbstractGraphQlHttpHandler { this.sseBuilder.complete(); } + static Consumer connect(Flux> resultFlux) { + return (sseBuilder) -> { + SseSubscriber subscriber = new SseSubscriber(sseBuilder); + resultFlux.subscribe(subscriber); + }; + } } } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlSseHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlSseHandlerTests.java index 1b38cf06..ff0fcce6 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlSseHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlSseHandlerTests.java @@ -50,12 +50,12 @@ import static org.assertj.core.api.Assertions.assertThat; */ class GraphQlSseHandlerTests { - private static final List> MESSAGE_WRITERS = Collections.singletonList(new ServerSentEventHttpMessageWriter(new Jackson2JsonEncoder())); + private static final List> MESSAGE_WRITERS = + List.of(new ServerSentEventHttpMessageWriter(new Jackson2JsonEncoder())); - private static final DataFetcher BOOK_SEARCH = environment -> { - String author = environment.getArgument("author"); - return Flux.fromIterable(BookSource.books()) - .filter((book) -> book.getAuthor().getFullName().contains(author)); + private static final DataFetcher SEARCH_DATA_FETCHER = env -> { + String author = env.getArgument("author"); + return Flux.fromIterable(BookSource.books()).filter((book) -> book.getAuthor().getFullName().contains(author)); }; private final MockServerHttpRequest httpRequest = MockServerHttpRequest.post("/graphql") @@ -66,70 +66,75 @@ class GraphQlSseHandlerTests { @Test void shouldRejectQueryOperations() { SerializableGraphQlRequest request = initRequest("{ bookById(id: 42) {name} }"); - GraphQlSseHandler sseHandler = createSseHandler(BOOK_SEARCH); - MockServerHttpResponse httpResponse = handleRequest(this.httpRequest, sseHandler, request); + GraphQlSseHandler handler = createHandler(SEARCH_DATA_FETCHER); + MockServerHttpResponse response = handleRequest(this.httpRequest, handler, request); - assertThat(httpResponse.getHeaders().getContentType().isCompatibleWith(MediaType.TEXT_EVENT_STREAM)).isTrue(); - assertThat(httpResponse.getBodyAsString().block()).isEqualTo( - """ - event:next - data:{"errors":[{"message":"SSE transport only supports Subscription operations","locations":[],"extensions":{"classification":"OperationNotSupported"}}]} + assertThat(response.getHeaders().getContentType().isCompatibleWith(MediaType.TEXT_EVENT_STREAM)).isTrue(); + assertThat(response.getBodyAsString().block()).isEqualTo(""" + event:next + data:{"errors":[{"message":"SSE handler supports only subscriptions","locations":[],"extensions":{"classification":"OperationNotSupported"}}]} - event:complete - data:{} + event:complete + data:{} - """); + """); } @Test void shouldWriteMultipleEventsForSubscription() { - SerializableGraphQlRequest request = initRequest("subscription TestSubscription { bookSearch(author:\"Orwell\") { id name } }"); - GraphQlSseHandler sseHandler = createSseHandler(BOOK_SEARCH); - MockServerHttpResponse httpResponse = handleRequest(this.httpRequest, sseHandler, request); - assertThat(httpResponse.getHeaders().getContentType().isCompatibleWith(MediaType.TEXT_EVENT_STREAM)).isTrue(); - assertThat(httpResponse.getBodyAsString().block()).isEqualTo( - """ - event:next - data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} + SerializableGraphQlRequest request = initRequest( + "subscription TestSubscription { bookSearch(author:\"Orwell\") { id name } }"); - event:next - data:{"data":{"bookSearch":{"id":"5","name":"Animal Farm"}}} + GraphQlSseHandler handler = createHandler(SEARCH_DATA_FETCHER); + MockServerHttpResponse response = handleRequest(this.httpRequest, handler, request); - event:complete - data:{} + assertThat(response.getHeaders().getContentType().isCompatibleWith(MediaType.TEXT_EVENT_STREAM)).isTrue(); + assertThat(response.getBodyAsString().block()).isEqualTo(""" + event:next + data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} - """); + event:next + data:{"data":{"bookSearch":{"id":"5","name":"Animal Farm"}}} + + event:complete + data:{} + + """); } @Test void shouldWriteEventsAndTerminalError() { - SerializableGraphQlRequest request = initRequest("subscription TestSubscription { bookSearch(author:\"Orwell\") { id name } }"); - DataFetcher errorDataFetcher = env -> Flux.just(BookSource.getBook(1L)) - .concatWith(Flux.error(new IllegalStateException("test error"))); - GraphQlSseHandler sseHandler = createSseHandler(errorDataFetcher); - MockServerHttpResponse httpResponse = handleRequest(this.httpRequest, sseHandler, request); - assertThat(httpResponse.getHeaders().getContentType().isCompatibleWith(MediaType.TEXT_EVENT_STREAM)).isTrue(); - assertThat(httpResponse.getBodyAsString().block()).isEqualTo( - """ - event:next - data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} + SerializableGraphQlRequest request = initRequest( + "subscription TestSubscription { bookSearch(author:\"Orwell\") { id name } }"); - event:next - data:{"errors":[{"message":"Subscription error","locations":[],"extensions":{"classification":"INTERNAL_ERROR"}}]} + DataFetcher errorDataFetcher = env -> + Flux.just(BookSource.getBook(1L)).concatWith(Flux.error(new IllegalStateException("test error"))); - event:complete - data:{} + GraphQlSseHandler handler = createHandler(errorDataFetcher); + MockServerHttpResponse response = handleRequest(this.httpRequest, handler, request); - """); + assertThat(response.getHeaders().getContentType().isCompatibleWith(MediaType.TEXT_EVENT_STREAM)).isTrue(); + assertThat(response.getBodyAsString().block()).isEqualTo(""" + event:next + data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} + + event:next + data:{"errors":[{"message":"Subscription error","locations":[],"extensions":{"classification":"INTERNAL_ERROR"}}]} + + event:complete + data:{} + + """); } - private GraphQlSseHandler createSseHandler(DataFetcher subscriptionDataFetcher) { - return new GraphQlSseHandler(GraphQlSetup.schemaResource(BookSource.schema) - .queryFetcher("bookById", (env) -> BookSource.getBookWithoutAuthor(1L)) - .subscriptionFetcher("bookSearch", subscriptionDataFetcher) - .toWebGraphQlHandler()); + private GraphQlSseHandler createHandler(DataFetcher subscriptionDataFetcher) { + return new GraphQlSseHandler( + GraphQlSetup.schemaResource(BookSource.schema) + .queryFetcher("bookById", (env) -> BookSource.getBookWithoutAuthor(1L)) + .subscriptionFetcher("bookSearch", subscriptionDataFetcher) + .toWebGraphQlHandler()); } private static SerializableGraphQlRequest initRequest(String document) { @@ -139,9 +144,9 @@ class GraphQlSseHandlerTests { } private MockServerHttpResponse handleRequest( - MockServerHttpRequest httpRequest, GraphQlSseHandler handler, GraphQlRequest body) { + MockServerHttpRequest request, GraphQlSseHandler handler, GraphQlRequest body) { - MockServerWebExchange exchange = MockServerWebExchange.from(httpRequest); + MockServerWebExchange exchange = MockServerWebExchange.from(request); MockServerRequest serverRequest = MockServerRequest.builder() .exchange(exchange) @@ -172,5 +177,4 @@ class GraphQlSseHandlerTests { } - } 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 9ac23bff..49bf9a00 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 @@ -20,7 +20,6 @@ package org.springframework.graphql.server.webmvc; import java.io.IOException; import java.nio.charset.StandardCharsets; import java.time.Duration; -import java.util.Collections; import java.util.List; import graphql.schema.DataFetcher; @@ -50,98 +49,93 @@ import static org.awaitility.Awaitility.await; class GraphQlSseHandlerTests { private static final List> MESSAGE_READERS = - Collections.singletonList(new MappingJackson2HttpMessageConverter()); + List.of(new MappingJackson2HttpMessageConverter()); - private static final DataFetcher BOOK_SEARCH = environment -> { - String author = environment.getArgument("author"); - return Flux.fromIterable(BookSource.books()) - .filter((book) -> book.getAuthor().getFullName().contains(author)); + private static final DataFetcher SEARCH_DATA_FETCHER = env -> { + String author = env.getArgument("author"); + return Flux.fromIterable(BookSource.books()).filter((book) -> book.getAuthor().getFullName().contains(author)); }; + @Test void shouldRejectQueryOperations() throws Exception { - GraphQlSseHandler sseHandler = createSseHandler(BOOK_SEARCH); + GraphQlSseHandler handler = createSseHandler(SEARCH_DATA_FETCHER); MockHttpServletRequest request = createServletRequest("{ \"query\": \"{ bookById(id: 42) {name} }\"}"); - MockHttpServletResponse response = handleRequest(request, sseHandler); + MockHttpServletResponse response = handleRequest(request, handler); assertThat(response.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE); - assertThat(response.getContentAsString()).isEqualTo( - """ - event:next - data:{"errors":[{"message":"SSE transport only supports Subscription operations","locations":[],"extensions":{"classification":"OperationNotSupported"}}]} + assertThat(response.getContentAsString()).isEqualTo(""" + event:next + data:{"errors":[{"message":"SSE handler supports only subscriptions","locations":[],"extensions":{"classification":"OperationNotSupported"}}]} - event:complete - data: + event:complete + data: - """); + """); } @Test void shouldWriteMultipleEventsForSubscription() throws Exception { - GraphQlSseHandler sseHandler = createSseHandler(BOOK_SEARCH); + GraphQlSseHandler handler = createSseHandler(SEARCH_DATA_FETCHER); MockHttpServletRequest request = createServletRequest(""" - { - "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" - } + { "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" } """); - MockHttpServletResponse response = handleRequest(request, sseHandler); + MockHttpServletResponse response = handleRequest(request, handler); assertThat(response.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE); - assertThat(response.getContentAsString()).isEqualTo( - """ - event:next - data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} + assertThat(response.getContentAsString()).isEqualTo(""" + event:next + data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} - event:next - data:{"data":{"bookSearch":{"id":"5","name":"Animal Farm"}}} + event:next + data:{"data":{"bookSearch":{"id":"5","name":"Animal Farm"}}} - event:complete - data: + event:complete + data: - """); + """); } @Test void shouldWriteEventsAndTerminalError() throws Exception { + DataFetcher errorDataFetcher = env -> Flux.just(BookSource.getBook(1L)) .concatWith(Flux.error(new IllegalStateException("test error"))); - GraphQlSseHandler sseHandler = createSseHandler(errorDataFetcher); + + GraphQlSseHandler handler = createSseHandler(errorDataFetcher); MockHttpServletRequest request = createServletRequest(""" - { - "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" - } + { "query": "subscription TestSubscription { bookSearch(author:\\\"Orwell\\\") { id name } }" } """); - MockHttpServletResponse response = handleRequest(request, sseHandler); + MockHttpServletResponse response = handleRequest(request, handler); assertThat(response.getContentType()).isEqualTo(MediaType.TEXT_EVENT_STREAM_VALUE); - assertThat(response.getContentAsString()).isEqualTo( - """ - event:next - data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} + assertThat(response.getContentAsString()).isEqualTo(""" + event:next + data:{"data":{"bookSearch":{"id":"1","name":"Nineteen Eighty-Four"}}} - event:next - data:{"errors":[{"message":"Subscription error","locations":[],"extensions":{"classification":"INTERNAL_ERROR"}}]} + event:next + data:{"errors":[{"message":"Subscription error","locations":[],"extensions":{"classification":"INTERNAL_ERROR"}}]} - event:complete - data: + event:complete + data: - """); + """); } - private GraphQlSseHandler createSseHandler(DataFetcher subscriptionDataFetcher) { + private GraphQlSseHandler createSseHandler(DataFetcher dataFetcher) { return new GraphQlSseHandler(GraphQlSetup.schemaResource(BookSource.schema) .queryFetcher("bookById", (env) -> BookSource.getBookWithoutAuthor(1L)) - .subscriptionFetcher("bookSearch", subscriptionDataFetcher) + .subscriptionFetcher("bookSearch", dataFetcher) .toWebGraphQlHandler()); } private MockHttpServletRequest createServletRequest(String query) { - MockHttpServletRequest servletRequest = new MockHttpServletRequest("POST", "/"); - servletRequest.setContentType(MediaType.APPLICATION_JSON_VALUE); - servletRequest.setContent(query.getBytes(StandardCharsets.UTF_8)); - servletRequest.addHeader("Accept", MediaType.TEXT_EVENT_STREAM_VALUE); - servletRequest.setAsyncSupported(true); - return servletRequest; + MockHttpServletRequest request = new MockHttpServletRequest("POST", "/"); + request.setContentType(MediaType.APPLICATION_JSON_VALUE); + request.setContent(query.getBytes(StandardCharsets.UTF_8)); + request.addHeader("Accept", MediaType.TEXT_EVENT_STREAM_VALUE); + request.setAsyncSupported(true); + return request; } private MockHttpServletResponse handleRequest(