diff --git a/spring-graphql-web/src/main/java/org/springframework/boot/graphql/WebFluxGraphQLAutoConfiguration.java b/spring-graphql-web/src/main/java/org/springframework/boot/graphql/WebFluxGraphQLAutoConfiguration.java index 9e157e99..d1e6a770 100644 --- a/spring-graphql-web/src/main/java/org/springframework/boot/graphql/WebFluxGraphQLAutoConfiguration.java +++ b/spring-graphql-web/src/main/java/org/springframework/boot/graphql/WebFluxGraphQLAutoConfiguration.java @@ -16,7 +16,6 @@ package org.springframework.boot.graphql; import java.util.Collections; -import java.util.Map; import graphql.GraphQL; @@ -27,15 +26,11 @@ import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean import org.springframework.boot.autoconfigure.condition.ConditionalOnWebApplication; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; -import org.springframework.core.ResolvableType; -import org.springframework.core.codec.Decoder; -import org.springframework.core.codec.Encoder; import org.springframework.core.io.Resource; import org.springframework.core.io.ResourceLoader; import org.springframework.graphql.WebFluxGraphQLHandler; +import org.springframework.graphql.WebFluxGraphQLWebSocketHandler; import org.springframework.http.MediaType; -import org.springframework.http.codec.DecoderHttpMessageReader; -import org.springframework.http.codec.EncoderHttpMessageWriter; import org.springframework.http.codec.ServerCodecConfigurer; import org.springframework.web.reactive.HandlerMapping; import org.springframework.web.reactive.function.server.RouterFunction; @@ -55,25 +50,16 @@ public class WebFluxGraphQLAutoConfiguration { @Bean @ConditionalOnMissingBean - public WebFluxGraphQLHandler graphQLHandler( + public WebFluxGraphQLHandler graphQLHandler(GraphQL.Builder graphQLBuilder) { + return new WebFluxGraphQLHandler(graphQLBuilder.build(), Collections.emptyList()); + } + + @Bean + @ConditionalOnMissingBean + public WebFluxGraphQLWebSocketHandler graphQLWebSocketHandler( GraphQL.Builder graphQLBuilder, ServerCodecConfigurer configurer) { - ResolvableType mapType = ResolvableType.forClass(Map.class); - - Decoder jsonDecoder = configurer.getReaders().stream() - .filter(reader -> reader.canRead(mapType, MediaType.APPLICATION_JSON)) - .map(reader -> ((DecoderHttpMessageReader) reader).getDecoder()) - .findFirst() - .orElseThrow(() -> new IllegalArgumentException("No JSON Decoder")); - - Encoder jsonEncoder = configurer.getWriters().stream() - .filter(writer -> writer.canWrite(mapType, MediaType.APPLICATION_JSON)) - .map(writer -> ((EncoderHttpMessageWriter) writer).getEncoder()) - .findFirst() - .orElseThrow(() -> new IllegalArgumentException("No JSON Encoder")); - - return new WebFluxGraphQLHandler( - graphQLBuilder.build(), Collections.emptyList(), jsonDecoder, jsonEncoder); + return new WebFluxGraphQLWebSocketHandler(graphQLBuilder.build(), Collections.emptyList(), configurer); } @Bean @@ -90,10 +76,12 @@ public class WebFluxGraphQLAutoConfiguration { } @Bean - public HandlerMapping graphQLWebSocketEndpoint(WebFluxGraphQLHandler handler, GraphQLProperties properties) { + public HandlerMapping graphQLWebSocketEndpoint( + WebFluxGraphQLWebSocketHandler handler, GraphQLProperties properties) { + String path = properties.getWebSocketPath(); SimpleUrlHandlerMapping mapping = new SimpleUrlHandlerMapping(); - mapping.setUrlMap(Collections.singletonMap(path, handler.getSubscriptionWebSocketHandler())); + mapping.setUrlMap(Collections.singletonMap(path, handler)); mapping.setOrder(-1); // Ahead of annotated controllers return mapping; } diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java b/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java index 0a737dad..099ecc04 100644 --- a/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java @@ -15,62 +15,35 @@ */ package org.springframework.graphql; -import java.util.Collections; import java.util.List; import java.util.Map; -import graphql.ExecutionResult; import graphql.GraphQL; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; -import org.reactivestreams.Publisher; import reactor.core.publisher.Mono; -import org.springframework.core.ResolvableType; -import org.springframework.core.codec.Decoder; -import org.springframework.core.codec.Encoder; -import org.springframework.core.io.buffer.DataBuffer; -import org.springframework.core.io.buffer.DataBufferUtils; -import org.springframework.util.CollectionUtils; -import org.springframework.util.MimeTypeUtils; import org.springframework.web.reactive.function.server.ServerRequest; import org.springframework.web.reactive.function.server.ServerResponse; -import org.springframework.web.reactive.socket.HandshakeInfo; -import org.springframework.web.reactive.socket.WebSocketHandler; -import org.springframework.web.reactive.socket.WebSocketMessage; -import org.springframework.web.reactive.socket.WebSocketSession; /** - * GraphQL handler to expose as a WebFlux.fn endpoint via - * {@link org.springframework.web.reactive.function.server.RouterFunctions}. + * WebFlux.fn Handler for GraphQL over HTTP requests. */ public class WebFluxGraphQLHandler { - private static Log logger = LogFactory.getLog(WebFluxGraphQLHandler.class); + private static final Log logger = LogFactory.getLog(WebFluxGraphQLHandler.class); private final WebInterceptorExecutionChain executionChain; - private final Decoder jsonDecoder; - - private final Encoder jsonEncoder; - /** - * Create a handler that executes queries through the given {@link GraphQL} - * and and invokes the given interceptors to customize input to and the - * result from the execution of the query. + * Create a new instance. * @param graphQL the GraphQL instance to use for query execution * @param interceptors 0 or more interceptors to customize input and output - * @param jsonDecoder to decode JSON for subscriptions over WebSocket - * @param jsonEncoder to encode JSON for subscriptions over WebSocket */ - public WebFluxGraphQLHandler(GraphQL graphQL, List interceptors, - Decoder jsonDecoder, Encoder jsonEncoder) { - + public WebFluxGraphQLHandler(GraphQL graphQL, List interceptors) { this.executionChain = new WebInterceptorExecutionChain(graphQL, interceptors); - this.jsonDecoder = jsonDecoder; - this.jsonEncoder = jsonEncoder; } @@ -99,69 +72,4 @@ public class WebFluxGraphQLHandler { }); } - /** - * Return a handler that supports subscriptions over WebSocket. - */ - public WebSocketHandler getSubscriptionWebSocketHandler() { - return new SubscriptionWebSocketHandler(); - } - - - /** - * Handler for subscriptions over WebSocket. - */ - private class SubscriptionWebSocketHandler implements WebSocketHandler { - - @Override - @SuppressWarnings("unchecked") - public Mono handle(WebSocketSession session) { - return session.send(session.receive() - .concatMap(message -> { - Map map = decode(message); - HandshakeInfo handshakeInfo = session.getHandshakeInfo(); - WebInput webInput = new WebInput(handshakeInfo.getUri(), handshakeInfo.getHeaders(), map); - if (logger.isDebugEnabled()) { - logger.debug("Executing: " + webInput); - } - return executionChain.execute(webInput); - }) - .concatMap(output -> { - if (!CollectionUtils.isEmpty(output.getErrors())) { - throw new IllegalStateException( - "Execution failed: " + output.getErrors()); - } - if (!(output.getData() instanceof Publisher)) { - throw new IllegalStateException( - "Expected Publisher: " + output.toSpecification()); - } - if (logger.isDebugEnabled()) { - logger.debug("Execution complete, subscribing for events."); - } - return (Publisher) output.getData(); - }) - .map(result -> { - Object data = result.getData(); - return encode(session, data); - }) - ); - } - - @SuppressWarnings({"unchecked", "ConstantConditions"}) - private Map decode(WebSocketMessage message) { - DataBuffer buffer = message.getPayload(); - return (Map) jsonDecoder.decode( - DataBufferUtils.retain(buffer), WebInput.MAP_RESOLVABLE_TYPE, null, Collections.emptyMap()); - } - - @SuppressWarnings("unchecked") - private WebSocketMessage encode(WebSocketSession session, Object data) { - DataBuffer buffer = ((Encoder) jsonEncoder).encodeValue((T) data, - session.bufferFactory(), - ResolvableType.forInstance(data), - MimeTypeUtils.APPLICATION_JSON, - Collections.emptyMap()); - return new WebSocketMessage(WebSocketMessage.Type.TEXT, buffer); - } - } - } diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLWebSocketHandler.java b/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLWebSocketHandler.java new file mode 100644 index 00000000..14f5d0b7 --- /dev/null +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLWebSocketHandler.java @@ -0,0 +1,145 @@ +/* + * Copyright 2002-2020 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. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.graphql; + +import java.util.Collections; +import java.util.List; +import java.util.Map; + +import graphql.ExecutionResult; +import graphql.GraphQL; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.reactivestreams.Publisher; +import reactor.core.publisher.Mono; + +import org.springframework.core.ResolvableType; +import org.springframework.core.codec.Decoder; +import org.springframework.core.codec.Encoder; +import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferUtils; +import org.springframework.http.MediaType; +import org.springframework.http.codec.DecoderHttpMessageReader; +import org.springframework.http.codec.EncoderHttpMessageWriter; +import org.springframework.http.codec.ServerCodecConfigurer; +import org.springframework.util.CollectionUtils; +import org.springframework.util.MimeTypeUtils; +import org.springframework.web.reactive.socket.HandshakeInfo; +import org.springframework.web.reactive.socket.WebSocketHandler; +import org.springframework.web.reactive.socket.WebSocketMessage; +import org.springframework.web.reactive.socket.WebSocketSession; + +/** + * WebSocketHandler for GraphQL based on + * GraphQL Over WebSocket Protocol + */ +public class WebFluxGraphQLWebSocketHandler implements WebSocketHandler { + + private static final Log logger = LogFactory.getLog(WebFluxGraphQLWebSocketHandler.class); + + private static final ResolvableType MAP_TYPE = ResolvableType.forClass(Map.class); + + + private final WebInterceptorExecutionChain executionChain; + + private final Decoder jsonDecoder; + + private final Encoder jsonEncoder; + + + /** + * Create a new instance. + * @param graphQL the GraphQL instance to use for query execution + * @param interceptors 0 or more interceptors to customize input and output + * @param configurer codec configurer for JSON encoding and decoding + */ + public WebFluxGraphQLWebSocketHandler(GraphQL graphQL, List interceptors, + ServerCodecConfigurer configurer) { + + this.executionChain = new WebInterceptorExecutionChain(graphQL, interceptors); + this.jsonDecoder = initDecoder(configurer); + this.jsonEncoder = initEncoder(configurer); + } + + private static Decoder initDecoder(ServerCodecConfigurer configurer) { + return configurer.getReaders().stream() + .filter(reader -> reader.canRead(MAP_TYPE, MediaType.APPLICATION_JSON)) + .map(reader -> ((DecoderHttpMessageReader) reader).getDecoder()) + .findFirst() + .orElseThrow(() -> new IllegalArgumentException("No JSON Decoder")); + } + + private static Encoder initEncoder(ServerCodecConfigurer configurer) { + return configurer.getWriters().stream() + .filter(writer -> writer.canWrite(MAP_TYPE, MediaType.APPLICATION_JSON)) + .map(writer -> ((EncoderHttpMessageWriter) writer).getEncoder()) + .findFirst() + .orElseThrow(() -> new IllegalArgumentException("No JSON Encoder")); + } + + + @Override + @SuppressWarnings("unchecked") + public Mono handle(WebSocketSession session) { + return session.send(session.receive() + .concatMap(message -> { + Map map = decode(message); + HandshakeInfo handshakeInfo = session.getHandshakeInfo(); + WebInput webInput = new WebInput(handshakeInfo.getUri(), handshakeInfo.getHeaders(), map); + if (logger.isDebugEnabled()) { + logger.debug("Executing: " + webInput); + } + return executionChain.execute(webInput); + }) + .concatMap(output -> { + if (!CollectionUtils.isEmpty(output.getErrors())) { + throw new IllegalStateException( + "Execution failed: " + output.getErrors()); + } + if (!(output.getData() instanceof Publisher)) { + throw new IllegalStateException( + "Expected Publisher: " + output.toSpecification()); + } + if (logger.isDebugEnabled()) { + logger.debug("Execution complete, subscribing for events."); + } + return (Publisher) output.getData(); + }) + .map(result -> { + Object data = result.getData(); + return encode(session, data); + }) + ); + } + + @SuppressWarnings({"unchecked", "ConstantConditions"}) + private Map decode(WebSocketMessage message) { + DataBuffer buffer = message.getPayload(); + return (Map) jsonDecoder.decode( + DataBufferUtils.retain(buffer), WebInput.MAP_RESOLVABLE_TYPE, null, Collections.emptyMap()); + } + + @SuppressWarnings("unchecked") + private WebSocketMessage encode(WebSocketSession session, Object data) { + DataBuffer buffer = ((Encoder) jsonEncoder).encodeValue((T) data, + session.bufferFactory(), + ResolvableType.forInstance(data), + MimeTypeUtils.APPLICATION_JSON, + Collections.emptyMap()); + return new WebSocketMessage(WebSocketMessage.Type.TEXT, buffer); + } + +} diff --git a/spring-graphql-web/src/test/java/org/springframework/boot/graphql/WebFluxApplicationContextTests.java b/spring-graphql-web/src/test/java/org/springframework/boot/graphql/WebFluxApplicationContextTests.java index 58c54b60..14cc55ad 100644 --- a/spring-graphql-web/src/test/java/org/springframework/boot/graphql/WebFluxApplicationContextTests.java +++ b/spring-graphql-web/src/test/java/org/springframework/boot/graphql/WebFluxApplicationContextTests.java @@ -29,6 +29,7 @@ import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.core.io.buffer.DefaultDataBufferFactory; import org.springframework.graphql.GraphQLDataFetchers; import org.springframework.graphql.WebFluxGraphQLHandler; +import org.springframework.graphql.WebFluxGraphQLWebSocketHandler; import org.springframework.http.HttpHeaders; import org.springframework.http.MediaType; import org.springframework.http.codec.json.Jackson2JsonDecoder; @@ -100,8 +101,7 @@ class WebFluxApplicationContextTests { Flux input = Flux.just(new WebSocketMessage(WebSocketMessage.Type.TEXT, buffer)); TestWebSocketSession session = new TestWebSocketSession("1", URI.create(BASE_URL), input); - context.getBean(WebFluxGraphQLHandler.class) - .getSubscriptionWebSocketHandler().handle(session).block(); + context.getBean(WebFluxGraphQLWebSocketHandler.class).handle(session).block(); StepVerifier.create(session.getOutput()) .consumeNextWith(message -> assertThat(extractBook(message)).containsEntry("id", "book-2"))