diff --git a/graphql-spring-boot-starter/build.gradle b/graphql-spring-boot-starter/build.gradle index 147a8152..a0bfe327 100644 --- a/graphql-spring-boot-starter/build.gradle +++ b/graphql-spring-boot-starter/build.gradle @@ -33,8 +33,9 @@ dependencies { api 'org.springframework.boot:spring-boot-starter' compileOnly 'org.springframework:spring-webflux' - compileOnly 'javax.servlet:javax.servlet-api' compileOnly 'org.springframework:spring-webmvc' + compileOnly 'org.springframework:spring-websocket' + compileOnly 'javax.servlet:javax.servlet-api' compileOnly 'io.micrometer:micrometer-core' compileOnly 'org.springframework.boot:spring-boot-actuator-autoconfigure' diff --git a/graphql-spring-boot-starter/src/main/java/org/springframework/graphql/boot/WebMvcGraphQLAutoConfiguration.java b/graphql-spring-boot-starter/src/main/java/org/springframework/graphql/boot/WebMvcGraphQLAutoConfiguration.java index 79eb604a..e60ef8a3 100644 --- a/graphql-spring-boot-starter/src/main/java/org/springframework/graphql/boot/WebMvcGraphQLAutoConfiguration.java +++ b/graphql-spring-boot-starter/src/main/java/org/springframework/graphql/boot/WebMvcGraphQLAutoConfiguration.java @@ -1,5 +1,5 @@ /* - * Copyright 2020-2020 the original author or authors. + * Copyright 2020-2021 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. @@ -16,6 +16,7 @@ package org.springframework.graphql.boot; import java.util.Collections; +import java.util.Map; import graphql.GraphQL; @@ -24,15 +25,22 @@ import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; import org.springframework.boot.autoconfigure.condition.ConditionalOnWebApplication; +import org.springframework.boot.autoconfigure.http.HttpMessageConverters; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.core.io.Resource; import org.springframework.core.io.ResourceLoader; import org.springframework.graphql.webmvc.GraphQLHttpHandler; +import org.springframework.graphql.webmvc.GraphQLWebSocketHandler; import org.springframework.http.MediaType; +import org.springframework.http.converter.HttpMessageConverter; +import org.springframework.web.servlet.HandlerMapping; import org.springframework.web.servlet.function.RouterFunction; import org.springframework.web.servlet.function.RouterFunctions; import org.springframework.web.servlet.function.ServerResponse; +import org.springframework.web.servlet.handler.SimpleUrlHandlerMapping; +import org.springframework.web.socket.server.support.DefaultHandshakeHandler; +import org.springframework.web.socket.server.support.WebSocketHttpRequestHandler; import static org.springframework.web.servlet.function.RequestPredicates.accept; import static org.springframework.web.servlet.function.RequestPredicates.contentType; @@ -50,6 +58,22 @@ public class WebMvcGraphQLAutoConfiguration { return new GraphQLHttpHandler(graphQLBuilder.build(), Collections.emptyList()); } + @Bean + @ConditionalOnMissingBean + public GraphQLWebSocketHandler graphQLWebSocketHandler( + GraphQL.Builder graphQLBuilder, GraphQLProperties properties, HttpMessageConverters converters) { + + HttpMessageConverter converter = converters.getConverters().stream() + .filter(candidate -> candidate.canRead(Map.class, MediaType.APPLICATION_JSON)) + .findFirst() + .orElseThrow(() -> new IllegalStateException("No JSON converter")); + + return new GraphQLWebSocketHandler( + graphQLBuilder.build(), Collections.emptyList(), + converter, properties.getConnectionInitTimeoutDuration() + ); + } + @Bean public RouterFunction graphQLQueryEndpoint( ResourceLoader resourceLoader, GraphQLHttpHandler handler, GraphQLProperties properties) { @@ -63,4 +87,16 @@ public class WebMvcGraphQLAutoConfiguration { .build(); } + @Bean + public HandlerMapping graphQLWebSocketEndpoint(GraphQLWebSocketHandler handler, GraphQLProperties properties) { + WebSocketHttpRequestHandler httpRequestHandler = + new WebSocketHttpRequestHandler(handler, new DefaultHandshakeHandler()); + + String path = properties.getWebSocketPath(); + SimpleUrlHandlerMapping mapping = new SimpleUrlHandlerMapping(); + mapping.setUrlMap(Collections.singletonMap(path, httpRequestHandler)); + mapping.setOrder(-1); // Ahead of annotated controllers + return mapping; + } + } diff --git a/spring-graphql-web/build.gradle b/spring-graphql-web/build.gradle index 3e25920c..39dd505c 100644 --- a/spring-graphql-web/build.gradle +++ b/spring-graphql-web/build.gradle @@ -29,11 +29,14 @@ dependencies { compileOnly "javax.annotation:javax.annotation-api:1.3.2" compileOnly 'org.springframework:spring-webflux' compileOnly 'org.springframework:spring-webmvc' + compileOnly 'org.springframework:spring-websocket' compileOnly 'javax.servlet:javax.servlet-api:4.0.1' testImplementation 'com.fasterxml.jackson.core:jackson-databind' testImplementation 'org.springframework:spring-webflux' testImplementation 'org.springframework:spring-webmvc' + testImplementation 'org.springframework:spring-websocket' + testImplementation 'org.springframework:spring-test' testImplementation 'io.projectreactor:reactor-test' testImplementation 'javax.servlet:javax.servlet-api:4.0.1' diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/webmvc/GraphQLWebSocketHandler.java b/spring-graphql-web/src/main/java/org/springframework/graphql/webmvc/GraphQLWebSocketHandler.java new file mode 100644 index 00000000..8b977219 --- /dev/null +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/webmvc/GraphQLWebSocketHandler.java @@ -0,0 +1,480 @@ +/* + * Copyright 2002-2021 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.webmvc; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.net.URI; +import java.time.Duration; +import java.util.Arrays; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + +import graphql.ErrorType; +import graphql.ExecutionResult; +import graphql.GraphQL; +import graphql.GraphqlErrorBuilder; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.reactivestreams.Publisher; +import org.reactivestreams.Subscription; +import reactor.core.publisher.BaseSubscriber; +import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; +import reactor.core.scheduler.Scheduler; +import reactor.core.scheduler.Schedulers; + +import org.springframework.graphql.WebInterceptor; +import org.springframework.graphql.WebInterceptorExecutionChain; +import org.springframework.graphql.WebOutput; +import org.springframework.graphql.WebSocketMessageInput; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpInputMessage; +import org.springframework.http.HttpOutputMessage; +import org.springframework.http.converter.HttpMessageConverter; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; +import org.springframework.util.CollectionUtils; +import org.springframework.web.socket.CloseStatus; +import org.springframework.web.socket.SubProtocolCapable; +import org.springframework.web.socket.TextMessage; +import org.springframework.web.socket.WebSocketSession; +import org.springframework.web.socket.handler.ExceptionWebSocketHandlerDecorator; +import org.springframework.web.socket.handler.TextWebSocketHandler; + +/** + * WebSocketHandler for GraphQL based on + * GraphQL Over WebSocket Protocol + * and for use on a Servlet container with {@code spring-websocket}. + */ +public class GraphQLWebSocketHandler extends TextWebSocketHandler implements SubProtocolCapable { + + private static final Log logger = LogFactory.getLog(GraphQLWebSocketHandler.class); + + private static final List SUB_PROTOCOL_LIST = + Arrays.asList("graphql-transport-ws", "subscriptions-transport-ws"); + + + private final WebInterceptorExecutionChain executionChain; + + private final Duration initTimeoutDuration; + + private final HttpMessageConverter converter; + + private final Map sessionInfoMap = new ConcurrentHashMap<>(); + + + /** + * 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 converter for JSON encoding and decoding + * @param initTimeoutDuration the time within which the {@code CONNECTION_INIT} + * type message must be received. + */ + public GraphQLWebSocketHandler(GraphQL graphQL, List interceptors, + HttpMessageConverter converter, Duration initTimeoutDuration) { + + Assert.notNull(converter, "HttpMessageConverter for JSON is required"); + this.executionChain = new WebInterceptorExecutionChain(graphQL, interceptors); + this.initTimeoutDuration = initTimeoutDuration; + this.converter = converter; + } + + + @Override + public List getSubProtocols() { + return SUB_PROTOCOL_LIST; + } + + + @Override + public void afterConnectionEstablished(WebSocketSession session) { + if ("subscriptions-transport-ws".equalsIgnoreCase(session.getAcceptedProtocol())) { + if (logger.isDebugEnabled()) { + logger.debug("apollographql/subscriptions-transport-ws is not supported, nor maintained. " + + "Please, use https://github.com/enisdenjo/graphql-ws."); + } + GraphQLStatus.closeSession(session, GraphQLStatus.INVALID_MESSAGE_STATUS); + return; + } + + SessionState sessionState = new SessionState(session.getId()); + this.sessionInfoMap.put(session.getId(), sessionState); + + Mono.delay(this.initTimeoutDuration) + .then(Mono.fromRunnable(() -> { + if (!sessionState.isConnectionInitProcessed()) { + GraphQLStatus.closeSession(session, GraphQLStatus.INIT_TIMEOUT_STATUS); + } + })) + .subscribe(); + } + + @Override + @SuppressWarnings("unchecked") + protected void handleTextMessage(WebSocketSession session, TextMessage message) throws Exception { + Map map = decode(message, Map.class); + String id = (String) map.get("id"); + MessageType messageType = MessageType.resolve((String) map.get("type")); + if (messageType == null) { + GraphQLStatus.closeSession(session, GraphQLStatus.INVALID_MESSAGE_STATUS); + return; + } + SessionState sessionState = getSessionInfo(session); + switch (messageType) { + case SUBSCRIBE: + if (!sessionState.isConnectionInitProcessed()) { + GraphQLStatus.closeSession(session, GraphQLStatus.UNAUTHORIZED_STATUS); + return; + } + if (id == null) { + GraphQLStatus.closeSession(session, GraphQLStatus.INVALID_MESSAGE_STATUS); + return; + } + URI uri = session.getUri(); + HttpHeaders headers = session.getHandshakeHeaders(); + WebSocketMessageInput input = new WebSocketMessageInput(uri, headers, id, getPayload(map)); + if (logger.isDebugEnabled()) { + logger.debug("Executing: " + input); + } + this.executionChain.execute(input) + .flatMapMany(output -> handleWebOutput(session, input.requestId(), output)) + .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread + .subscribe(new SendMessageSubscriber(id, session, sessionState)); + return; + case COMPLETE: + if (id != null) { + Subscription subscription = sessionState.getSubscriptions().remove(id); + if (subscription != null) { + subscription.cancel(); + } + } + return; + case CONNECTION_INIT: + if (sessionState.setConnectionInitProcessed()) { + GraphQLStatus.closeSession(session, GraphQLStatus.TOO_MANY_INIT_REQUESTS_STATUS); + return; + } + TextMessage outputMessage = encode(session, null, MessageType.CONNECTION_ACK, null); + session.sendMessage(outputMessage); + return; + default: + GraphQLStatus.closeSession(session, GraphQLStatus.INVALID_MESSAGE_STATUS); + } + + } + + @SuppressWarnings("unchecked") + private T decode(TextMessage message, Class targetClass) throws IOException { + return ((HttpMessageConverter) this.converter).read(targetClass, new HttpInputMessageAdapter(message)); + } + + @SuppressWarnings("unchecked") + private static Map getPayload(Map message) { + Map payload = (Map) message.get("payload"); + Assert.notNull(payload, "No \"payload\" in message: " + message); + return payload; + } + + private SessionState getSessionInfo(WebSocketSession session) { + SessionState info = this.sessionInfoMap.get(session.getId()); + Assert.notNull(info, "No SessionInfo for " + session); + return info; + } + + @SuppressWarnings("unchecked") + private Flux handleWebOutput(WebSocketSession session, String id, WebOutput output) { + if (logger.isDebugEnabled()) { + logger.debug("Execution result ready" + + (!CollectionUtils.isEmpty(output.getErrors()) ? + " with errors: " + output.getErrors() : "") + "."); + } + Flux outputFlux; + if (output.getData() instanceof Publisher) { + // Subscription + outputFlux = Flux.from((Publisher) output.getData()) + .doOnSubscribe(subscription -> { + Subscription prev = getSessionInfo(session).getSubscriptions().putIfAbsent(id, subscription); + if (prev != null) { + throw new SubscriptionExistsException(); + } + }); + } + else { + // Query + outputFlux = (CollectionUtils.isEmpty(output.getErrors()) ? + Flux.just(output) : + Flux.error(new IllegalStateException("Execution failed: " + output.getErrors()))); + } + + return outputFlux + .map(result -> { + Map dataMap = result.toSpecification(); + return encode(session, id, MessageType.NEXT, dataMap); + }) + .concatWith(Mono.fromCallable(() -> encode(session, id, MessageType.COMPLETE, null))) + .onErrorResume(ex -> { + if (ex instanceof SubscriptionExistsException) { + CloseStatus status = new CloseStatus(4409, "Subscriber for " + id + " already exists"); + GraphQLStatus.closeSession(session, status); + return Flux.empty(); + } + ErrorType errorType = ErrorType.DataFetchingException; + String message = ex.getMessage(); + Map errorMap = GraphqlErrorBuilder.newError() + .errorType(errorType) + .message(message) + .build() + .toSpecification(); + return Mono.just(encode(session, id, MessageType.ERROR, errorMap)); + }); + } + + @SuppressWarnings("unchecked") + private TextMessage encode( + WebSocketSession session, @Nullable String id, MessageType messageType, @Nullable Object payload) { + + Map payloadMap = new HashMap<>(3); + payloadMap.put("type", messageType.getType()); + if (id != null) { + payloadMap.put("id", id); + } + if (payload != null) { + payloadMap.put("payload", payload); + } + try { + HttpOutputMessageAdapter outputMessage = new HttpOutputMessageAdapter(); + ((HttpMessageConverter) this.converter).write((T) payloadMap, null, outputMessage); + return new TextMessage(outputMessage.toByteArray()); + } + catch (IOException ex) { + throw new IllegalStateException("Failed to write " + payloadMap + " as JSON", ex); + } + } + + @Override + public void handleTransportError(WebSocketSession session, Throwable exception) { + SessionState info = this.sessionInfoMap.remove(session.getId()); + if (info != null) { + info.dispose(); + } + } + + @Override + public void afterConnectionClosed(WebSocketSession session, CloseStatus closeStatus) { + SessionState info = this.sessionInfoMap.remove(session.getId()); + if (info != null) { + info.dispose(); + } + } + + @Override + public boolean supportsPartialMessages() { + return false; + } + + + private enum MessageType { + + CONNECTION_INIT("connection_init"), + CONNECTION_ACK("connection_ack"), + SUBSCRIBE("subscribe"), + NEXT("next"), + ERROR("error"), + COMPLETE("complete"); + + + private static final Map messageTypes = new HashMap<>(6); + + static { + for (MessageType messageType : MessageType.values()) { + messageTypes.put(messageType.getType(), messageType); + } + } + + + private final String type; + + MessageType(String type) { + this.type = type; + } + + public String getType() { + return this.type; + } + + @Nullable + public static MessageType resolve(@Nullable String type) { + return (type != null ? messageTypes.get(type) : null); + } + } + + + private static class GraphQLStatus { + + private static final CloseStatus INVALID_MESSAGE_STATUS = new CloseStatus(4400, "Invalid message"); + + private static final CloseStatus UNAUTHORIZED_STATUS = new CloseStatus(4401, "Unauthorized"); + + private static final CloseStatus INIT_TIMEOUT_STATUS = new CloseStatus(4408, "Connection initialisation timeout"); + + private static final CloseStatus TOO_MANY_INIT_REQUESTS_STATUS = new CloseStatus(4429, "Too many initialisation requests"); + + + public static void closeSession(WebSocketSession session, CloseStatus status) { + try { + session.close(status); + } + catch (IOException ex) { + if (logger.isDebugEnabled()) { + logger.debug("Error while closing session with status: " + status, ex); + } + } + } + } + + + private static class HttpInputMessageAdapter extends ByteArrayInputStream implements HttpInputMessage { + + HttpInputMessageAdapter(TextMessage message) { + super(message.asBytes()); + } + + @Override + public InputStream getBody() { + return this; + } + + @Override + public HttpHeaders getHeaders() { + return HttpHeaders.EMPTY; + } + } + + + private static class HttpOutputMessageAdapter extends ByteArrayOutputStream implements HttpOutputMessage { + + private static final HttpHeaders noOpHeaders = new HttpHeaders(); + + @Override + public OutputStream getBody() { + return this; + } + + @Override + public HttpHeaders getHeaders() { + return noOpHeaders; + } + } + + + private static class SessionState { + + private boolean connectionInitProcessed; + + private final Map subscriptions = new ConcurrentHashMap<>(); + + private final Scheduler scheduler; + + public SessionState(String sessionId) { + this.scheduler = Schedulers.newSingle("GraphQL-WsSession-" + sessionId); + } + + public boolean isConnectionInitProcessed() { + return this.connectionInitProcessed; + } + + public synchronized boolean setConnectionInitProcessed() { + boolean previousValue = this.connectionInitProcessed; + this.connectionInitProcessed = true; + return previousValue; + } + + public Map getSubscriptions() { + return this.subscriptions; + } + + public void dispose() { + for (Map.Entry entry : this.subscriptions.entrySet()) { + try { + entry.getValue().cancel(); + } + catch (Throwable ex) { + // Ignore and keep on + } + } + this.subscriptions.clear(); + this.scheduler.dispose(); + } + + public Scheduler getScheduler() { + return this.scheduler; + } + } + + + private class SendMessageSubscriber extends BaseSubscriber { + + private final String subscriptionId; + + private final WebSocketSession session; + + private final SessionState sessionState; + + public SendMessageSubscriber(String subscriptionId, WebSocketSession session, SessionState sessionState) { + this.subscriptionId = subscriptionId; + this.session = session; + this.sessionState = sessionState; + } + + @Override + protected void hookOnSubscribe(Subscription subscription) { + subscription.request(1); + } + + @Override + protected void hookOnNext(TextMessage nextMessage) { + try { + this.session.sendMessage(nextMessage); + request(1); + } + catch (IOException ex) { + ExceptionWebSocketHandlerDecorator.tryCloseWithError(this.session, ex, logger); + } + } + + @Override + public void hookOnError(Throwable ex) { + ExceptionWebSocketHandlerDecorator.tryCloseWithError(this.session, ex, logger); + } + + @Override + public void hookOnComplete() { + this.sessionState.getSubscriptions().remove(this.subscriptionId); + } + } + + + private static class SubscriptionExistsException extends RuntimeException { + } + +} diff --git a/spring-graphql-web/src/test/java/org/springframework/graphql/webmvc/GraphQLWebSocketHandlerTests.java b/spring-graphql-web/src/test/java/org/springframework/graphql/webmvc/GraphQLWebSocketHandlerTests.java new file mode 100644 index 00000000..ac1294a1 --- /dev/null +++ b/spring-graphql-web/src/test/java/org/springframework/graphql/webmvc/GraphQLWebSocketHandlerTests.java @@ -0,0 +1,327 @@ +/* + * Copyright 2002-2021 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.webmvc; + +import java.io.ByteArrayInputStream; +import java.io.File; +import java.io.IOException; +import java.io.InputStream; +import java.time.Duration; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.function.BiConsumer; +import java.util.function.Consumer; + +import graphql.GraphQL; +import graphql.schema.GraphQLSchema; +import graphql.schema.idl.RuntimeWiring; +import graphql.schema.idl.SchemaGenerator; +import graphql.schema.idl.SchemaParser; +import graphql.schema.idl.TypeDefinitionRegistry; +import org.junit.jupiter.api.Test; +import reactor.test.StepVerifier; + +import org.springframework.graphql.ConsumeOneAndNeverCompleteInterceptor; +import org.springframework.graphql.GraphQLDataFetchers; +import org.springframework.graphql.WebInterceptor; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpInputMessage; +import org.springframework.http.converter.HttpMessageConverter; +import org.springframework.http.converter.json.MappingJackson2HttpMessageConverter; +import org.springframework.lang.Nullable; +import org.springframework.util.ResourceUtils; +import org.springframework.web.socket.CloseStatus; +import org.springframework.web.socket.TextMessage; +import org.springframework.web.socket.WebSocketMessage; + +import static graphql.schema.idl.TypeRuntimeWiring.newTypeWiring; +import static org.assertj.core.api.Assertions.as; +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.InstanceOfAssertFactories.map; + +/** + * Unit tests for {@link GraphQLWebSocketHandler}. + */ +public class GraphQLWebSocketHandlerTests { + + private static final String SUBSCRIPTION_ID = "123"; + + private static final String BOOK_SEARCH_QUERY = "{" + + "\"id\":\"" + SUBSCRIPTION_ID + "\"," + + "\"type\":\"subscribe\"," + + "\"payload\":{\"query\": \"" + + " subscription TestSubscription {" + + " bookSearch(minPages: 200) {" + + " id" + + " name" + + " pageCount" + + " author" + + " }}\"}" + + "}"; + + private static final HttpMessageConverter converter = new MappingJackson2HttpMessageConverter(); + + + private final TestWebSocketSession session = new TestWebSocketSession(); + + private final GraphQLWebSocketHandler handler = + initWebSocketHandler(Collections.emptyList(), Duration.ofSeconds(60)); + + + @Test + void query() throws Exception { + String bookQuery = "{" + + "\"id\":\"" + SUBSCRIPTION_ID + "\"," + + "\"type\":\"subscribe\"," + + "\"payload\":{\"query\": \"" + + " query TestQuery {" + + " bookById(id: \\\"book-1\\\"){ " + + " id" + + " name" + + " pageCount" + + " author" + + " }}\"}" + + "}"; + + this.handler.afterConnectionEstablished(session); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + this.handler.handleTextMessage(session, new TextMessage(bookQuery)); + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, "connection_ack")) + .consumeNextWith(message -> + assertThat(decode(message)) + .hasSize(3) + .containsEntry("id", SUBSCRIPTION_ID) + .containsEntry("type", "next") + .extractingByKey("payload", as(map(String.class, Object.class))) + .extractingByKey("data", as(map(String.class, Object.class))) + .extractingByKey("bookById", as(map(String.class, Object.class))) + .containsEntry("name", "GraphQL for beginners")) + .consumeNextWith(message -> assertMessageType(message, "complete")) + .then(session::close) // Complete output Flux + .verifyComplete(); + } + + @Test + void subscription() throws Exception { + this.handler.afterConnectionEstablished(session); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + this.handler.handleTextMessage(session, new TextMessage(BOOK_SEARCH_QUERY)); + + BiConsumer, String> bookPayloadAssertion = (message, bookId) -> + assertThat(decode(message)) + .hasSize(3) + .containsEntry("id", SUBSCRIPTION_ID) + .containsEntry("type", "next") + .extractingByKey("payload", as(map(String.class, Object.class))) + .extractingByKey("data", as(map(String.class, Object.class))) + .extractingByKey("bookSearch", as(map(String.class, Object.class))) + .containsEntry("id", bookId); + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, "connection_ack")) + .consumeNextWith(message -> bookPayloadAssertion.accept(message, "book-2")) + .consumeNextWith(message -> bookPayloadAssertion.accept(message, "book-3")) + .consumeNextWith(message -> bookPayloadAssertion.accept(message, "book-3")) + .consumeNextWith(message -> assertMessageType(message, "complete")) + .then(session::close) // Complete output Flux + .verifyComplete(); + } + + @Test + void unauthorizedWithoutMessageType() throws Exception { + this.handler.afterConnectionEstablished(session); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + this.handler.handleTextMessage(session, new TextMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\"}")); // No message type + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, "connection_ack")) + .verifyComplete(); + + assertThat(session.getCloseStatus()).isEqualTo(new CloseStatus(4400, "Invalid message")); + } + + @Test + void invalidMessageWithoutId() throws Exception { + this.handler.afterConnectionEstablished(session); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"subscribe\"}")); // No message id + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, "connection_ack")) + .verifyComplete(); + + assertThat(session.getCloseStatus()).isEqualTo(new CloseStatus(4400, "Invalid message")); + } + + @Test + void unauthorizedWithoutConnectionInit() throws Exception { + this.handler.afterConnectionEstablished(session); + this.handler.handleTextMessage(session, new TextMessage(BOOK_SEARCH_QUERY)); + + StepVerifier.create(session.getOutput()).verifyComplete(); + assertThat(session.getCloseStatus()).isEqualTo(new CloseStatus(4401, "Unauthorized")); + } + + @Test + void tooManyConnectionInitRequests() throws Exception { + this.handler.afterConnectionEstablished(session); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + this.handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, "connection_ack")) + .verifyComplete(); + + assertThat(session.getCloseStatus()) + .isEqualTo(new CloseStatus(4429, "Too many initialisation requests")); + } + + @Test + void connectionInitTimeout() { + GraphQLWebSocketHandler handler = initWebSocketHandler(Collections.emptyList(), Duration.ofMillis(50)); + handler.afterConnectionEstablished(session); + + StepVerifier.create(session.closeStatus()) + .expectNext(new CloseStatus(4408, "Connection initialisation timeout")) + .verifyComplete(); + } + + @Test + void subscriptionExists() throws Exception { + GraphQLWebSocketHandler handler = initWebSocketHandler( + Collections.singletonList(new ConsumeOneAndNeverCompleteInterceptor()), null); + + handler.afterConnectionEstablished(session); + handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + handler.handleTextMessage(session, new TextMessage(BOOK_SEARCH_QUERY)); + handler.handleTextMessage(session, new TextMessage(BOOK_SEARCH_QUERY)); + + // Collect messages until session closed + List> messages = new ArrayList<>(); + session.getOutput().subscribe(message -> messages.add(decode(message))); + + StepVerifier.create(session.closeStatus()) + .expectNext(new CloseStatus(4409, "Subscriber for " + SUBSCRIPTION_ID + " already exists")) + .verifyComplete(); + + assertThat(messages.size()).isEqualTo(2); + assertThat(messages.get(0).get("type")).isEqualTo("connection_ack"); + assertThat(messages.get(1).get("type")).isEqualTo("next"); + } + + @Test + void clientCompletion() throws Exception { + GraphQLWebSocketHandler handler = initWebSocketHandler( + Collections.singletonList(new ConsumeOneAndNeverCompleteInterceptor()), null); + + handler.afterConnectionEstablished(session); + handler.handleTextMessage(session, new TextMessage("{\"type\":\"connection_init\"}")); + handler.handleTextMessage(session, new TextMessage(BOOK_SEARCH_QUERY)); + + String completeMessage = "{\"id\":\"" + SUBSCRIPTION_ID + "\",\"type\":\"complete\"}"; + Consumer messageSender = body -> { + try { + handler.handleTextMessage(session, new TextMessage(body)); + } + catch (Exception ex) { + throw new IllegalStateException(ex); + } + }; + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, "connection_ack")) + .consumeNextWith(message -> assertMessageType(message, "next")) + .then(() -> messageSender.accept(completeMessage)) + .as("Second subscription with same id is possible only if the first was properly removed") + .then(() -> messageSender.accept(BOOK_SEARCH_QUERY)) + .consumeNextWith(message -> assertMessageType(message, "next")) + .then(() -> messageSender.accept(completeMessage)) + .verifyTimeout(Duration.ofMillis(500)); + } + + + private GraphQLWebSocketHandler initWebSocketHandler( + @Nullable List interceptors, @Nullable Duration initTimeoutDuration) { + + try { + return new GraphQLWebSocketHandler(initGraphQL(), + (interceptors != null ? interceptors : Collections.emptyList()), converter, + (initTimeoutDuration != null ? initTimeoutDuration : Duration.ofSeconds(60))); + } + catch (Exception ex) { + throw new IllegalStateException(ex); + } + } + + private static GraphQL initGraphQL() throws Exception { + File schemaFile = ResourceUtils.getFile("classpath:books/schema.graphqls"); + TypeDefinitionRegistry typeDefinitionRegistry = new SchemaParser().parse(schemaFile); + + RuntimeWiring.Builder builder = RuntimeWiring.newRuntimeWiring(); + builder.type(newTypeWiring("Query").dataFetcher("bookById", GraphQLDataFetchers.getBookByIdDataFetcher())); + builder.type(newTypeWiring("Subscription").dataFetcher("bookSearch", GraphQLDataFetchers.getBooksOnSale())); + RuntimeWiring runtimeWiring = builder.build(); + + GraphQLSchema schema = new SchemaGenerator().makeExecutableSchema(typeDefinitionRegistry, runtimeWiring); + return GraphQL.newGraphQL(schema).build(); + } + + private void assertMessageType(WebSocketMessage message, String messageType) { + Map map = decode(message, Map.class); + assertThat(map).containsEntry("type", messageType); + if (!messageType.equals("connection_ack")) { + assertThat(map).containsEntry("id", SUBSCRIPTION_ID); + } + } + + @SuppressWarnings("unchecked") + private Map decode(WebSocketMessage message) { + return decode(message, Map.class); + } + + @SuppressWarnings("unchecked") + private T decode(WebSocketMessage message, Class targetClass) { + try { + HttpInputMessageAdapter inputMessage = new HttpInputMessageAdapter((TextMessage) message); + return ((HttpMessageConverter) converter).read(targetClass, inputMessage); + } + catch (IOException ex) { + throw new IllegalStateException(ex); + } + } + + + private static class HttpInputMessageAdapter extends ByteArrayInputStream implements HttpInputMessage { + + HttpInputMessageAdapter(TextMessage message) { + super(message.asBytes()); + } + + @Override + public InputStream getBody() { + return this; + } + + @Override + public HttpHeaders getHeaders() { + return HttpHeaders.EMPTY; + } + } +} diff --git a/spring-graphql-web/src/test/java/org/springframework/graphql/webmvc/TestWebSocketSession.java b/spring-graphql-web/src/test/java/org/springframework/graphql/webmvc/TestWebSocketSession.java new file mode 100644 index 00000000..1a9a8774 --- /dev/null +++ b/spring-graphql-web/src/test/java/org/springframework/graphql/webmvc/TestWebSocketSession.java @@ -0,0 +1,161 @@ +/* + * Copyright 2002-2021 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.webmvc; + +import java.net.InetSocketAddress; +import java.net.URI; +import java.security.Principal; +import java.util.List; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + +import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; +import reactor.core.publisher.Sinks; + +import org.springframework.http.HttpHeaders; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; +import org.springframework.web.socket.CloseStatus; +import org.springframework.web.socket.WebSocketExtension; +import org.springframework.web.socket.WebSocketMessage; +import org.springframework.web.socket.WebSocketSession; + +/** + * + */ +public class TestWebSocketSession implements WebSocketSession { + + private final URI uri = URI.create("https://example.org/graphql"); + + private final HttpHeaders headers = new HttpHeaders(); + + private final Map attributes = new ConcurrentHashMap<>(); + + private final Sinks.Many> messagesSink = Sinks.many().unicast().onBackpressureBuffer(); + + private Sinks.One statusSink = Sinks.one(); + + private boolean closed; + + + @Override + public String getId() { + return "1"; + } + + @Override + public URI getUri() { + return this.uri; + } + + @Override + public HttpHeaders getHandshakeHeaders() { + return this.headers; + } + + @Override + public Map getAttributes() { + return this.attributes; + } + + @Override + public Principal getPrincipal() { + throw new UnsupportedOperationException(); + } + + @Override + public InetSocketAddress getLocalAddress() { + throw new UnsupportedOperationException(); + } + + @Override + public InetSocketAddress getRemoteAddress() { + throw new UnsupportedOperationException(); + } + + @Override + public String getAcceptedProtocol() { + return "graphql-transport-ws"; + } + + @Override + public void setTextMessageSizeLimit(int messageSizeLimit) { + throw new UnsupportedOperationException(); + } + + @Override + public int getTextMessageSizeLimit() { + throw new UnsupportedOperationException(); + } + + @Override + public void setBinaryMessageSizeLimit(int messageSizeLimit) { + throw new UnsupportedOperationException(); + } + + @Override + public int getBinaryMessageSizeLimit() { + throw new UnsupportedOperationException(); + } + + @Override + public List getExtensions() { + throw new UnsupportedOperationException(); + } + + @Override + public void sendMessage(WebSocketMessage message) { + emitMessagesSignal(this.messagesSink.tryEmitNext(message)); + } + + private void emitMessagesSignal(Sinks.EmitResult result) { + Assert.state(result == Sinks.EmitResult.OK, "Emit failed: " + result); + } + + public Flux> getOutput() { + return this.messagesSink.asFlux(); + } + + @Override + public boolean isOpen() { + return !this.closed; + } + + @Override + public void close() { + this.closed = true; + emitMessagesSignal(this.messagesSink.tryEmitComplete()); + this.statusSink.tryEmitEmpty(); + } + + @Override + public void close(CloseStatus status) { + this.closed = true; + emitMessagesSignal(this.messagesSink.tryEmitComplete()); + this.statusSink.tryEmitValue(status); + } + + @Nullable + public CloseStatus getCloseStatus() { + return (this.closed ? this.statusSink.asMono().block() : null); + } + + public Mono closeStatus() { + return this.statusSink.asMono(); + } + +}