From 878c643672077fd8ad0cb016a652b6b1a80b9675 Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Tue, 8 Mar 2022 12:50:32 +0000 Subject: [PATCH] Refactor GraphQlMessage and add GraphQlMessageType enum See gh-270 --- .../graphql/client/CodecDelegate.java | 10 +- .../client/WebSocketGraphQlTransport.java | 43 ++-- .../graphql/web/support/GraphQlMessage.java | 221 ++++++++++++++++++ .../web/support/GraphQlMessageType.java | 99 ++++++++ .../graphql/web/support/package-info.java | 25 ++ .../graphql/web/webflux/CodecDelegate.java | 44 ++-- .../web/webflux/GraphQlWebSocketHandler.java | 17 +- .../web/webflux/GraphQlWebSocketMessage.java | 165 ------------- .../web/webmvc/GraphQlWebSocketHandler.java | 135 +++++------ .../client/MockGraphQlWebSocketServer.java | 54 ++--- .../MockWebSocketGraphQlTransportTests.java | 54 ++--- .../web/WebSocketHandlerTestSupport.java | 36 +-- .../webflux/GraphQlWebSocketHandlerTests.java | 68 +++--- .../webmvc/GraphQlWebSocketHandlerTests.java | 69 +++--- 14 files changed, 619 insertions(+), 421 deletions(-) create mode 100644 spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessage.java create mode 100644 spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessageType.java create mode 100644 spring-graphql/src/main/java/org/springframework/graphql/web/support/package-info.java delete mode 100644 spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketMessage.java diff --git a/spring-graphql/src/main/java/org/springframework/graphql/client/CodecDelegate.java b/spring-graphql/src/main/java/org/springframework/graphql/client/CodecDelegate.java index 0e183176..47b22556 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/client/CodecDelegate.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/client/CodecDelegate.java @@ -20,7 +20,7 @@ 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.graphql.web.webflux.GraphQlWebSocketMessage; +import org.springframework.graphql.web.support.GraphQlMessage; import org.springframework.http.MediaType; import org.springframework.http.codec.ClientCodecConfigurer; import org.springframework.http.codec.CodecConfigurer; @@ -40,7 +40,7 @@ import org.springframework.web.reactive.socket.WebSocketSession; */ final class CodecDelegate { - private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlWebSocketMessage.class); + private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlMessage.class); private final CodecConfigurer codecConfigurer; @@ -84,7 +84,7 @@ final class CodecDelegate { @SuppressWarnings("unchecked") - public WebSocketMessage encode(WebSocketSession session, GraphQlWebSocketMessage message) { + public WebSocketMessage encode(WebSocketSession session, GraphQlMessage message) { DataBuffer buffer = ((Encoder) this.encoder).encodeValue( (T) message, session.bufferFactory(), MESSAGE_TYPE, MimeTypeUtils.APPLICATION_JSON, null); @@ -93,9 +93,9 @@ final class CodecDelegate { } @SuppressWarnings("ConstantConditions") - public GraphQlWebSocketMessage decode(WebSocketMessage webSocketMessage) { + public GraphQlMessage decode(WebSocketMessage webSocketMessage) { DataBuffer buffer = DataBufferUtils.retain(webSocketMessage.getPayload()); - return (GraphQlWebSocketMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null); + return (GraphQlMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null); } } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java b/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java index 7400c0a4..8baf517e 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java @@ -36,7 +36,8 @@ import reactor.core.publisher.Sinks; import org.springframework.graphql.GraphQlRequest; import org.springframework.graphql.support.MapExecutionResult; import org.springframework.graphql.support.MapGraphQlError; -import org.springframework.graphql.web.webflux.GraphQlWebSocketMessage; +import org.springframework.graphql.web.support.GraphQlMessage; +import org.springframework.graphql.web.support.GraphQlMessageType; import org.springframework.http.HttpHeaders; import org.springframework.http.codec.CodecConfigurer; import org.springframework.lang.Nullable; @@ -168,7 +169,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { private final CodecDelegate codecDelegate; - private final GraphQlWebSocketMessage connectionInitMessage; + private final GraphQlMessage connectionInitMessage; private final Consumer> connectionAckHandler; @@ -181,7 +182,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { @Nullable Object connectionInitPayload, Consumer> connectionAckHandler) { this.codecDelegate = new CodecDelegate(codecConfigurer); - this.connectionInitMessage = GraphQlWebSocketMessage.connectionInit(connectionInitPayload); + this.connectionInitMessage = GraphQlMessage.connectionInit(connectionInitPayload); this.connectionAckHandler = connectionAckHandler; this.graphQlSessionSink = Sinks.unsafe().one(); } @@ -240,8 +241,8 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { .flatMap(webSocketMessage -> { if (sessionNotInitialized()) { try { - GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage); - Assert.state(message.getType().equals("connection_ack"), + GraphQlMessage message = this.codecDelegate.decode(webSocketMessage); + Assert.state(message.resolvedType() == GraphQlMessageType.CONNECTION_ACK, () -> "Unexpected message before connection_ack: " + message); this.connectionAckHandler.accept(message.getPayload()); if (logger.isDebugEnabled()) { @@ -259,15 +260,15 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { } } else { - GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage); - switch (message.getType()) { - case "next": + GraphQlMessage message = this.codecDelegate.decode(webSocketMessage); + switch (message.resolvedType()) { + case NEXT: graphQlSession.handleNext(message); break; - case "error": + case ERROR: graphQlSession.handleError(message); break; - case "complete": + case COMPLETE: graphQlSession.handleComplete(message); break; default: @@ -366,7 +367,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { private final AtomicLong requestIndex = new AtomicLong(); - private final Sinks.Many requestSink = Sinks.many().unicast().onBackpressureBuffer(); + private final Sinks.Many requestSink = Sinks.many().unicast().onBackpressureBuffer(); private final Map> resultSinks = new ConcurrentHashMap<>(); @@ -381,14 +382,14 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { /** * Return the {@code Flux} of GraphQL requests to send as WebSocket messages. */ - public Flux getRequestFlux() { + public Flux getRequestFlux() { return this.requestSink.asFlux(); } public Mono execute(GraphQlRequest request) { String id = String.valueOf(this.requestIndex.incrementAndGet()); try { - GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request); + GraphQlMessage message = GraphQlMessage.subscribe(id, request); Sinks.One sink = Sinks.one(); this.resultSinks.put(id, sink); trySend(message); @@ -403,7 +404,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { public Flux executeSubscription(GraphQlRequest request) { String id = String.valueOf(this.requestIndex.incrementAndGet()); try { - GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request); + GraphQlMessage message = GraphQlMessage.subscribe(id, request); Sinks.Many sink = Sinks.many().unicast().onBackpressureBuffer(); this.streamingSinks.put(id, sink); trySend(message); @@ -417,7 +418,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { // TODO: queue to serialize sending? - private void trySend(GraphQlWebSocketMessage message) { + private void trySend(GraphQlMessage message) { Sinks.EmitResult emitResult = null; for (int i = 0; i < 100; i++) { emitResult = this.requestSink.tryEmitNext(message); @@ -432,7 +433,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { Sinks.Many streamSink = this.streamingSinks.remove(id); if (streamSink != null) { try { - trySend(GraphQlWebSocketMessage.complete(id)); + trySend(GraphQlMessage.complete(id)); } catch (Exception ex) { if (logger.isErrorEnabled()) { @@ -447,7 +448,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { /** * Handle a "next" message and route to its recipient. */ - public void handleNext(GraphQlWebSocketMessage message) { + public void handleNext(GraphQlMessage message) { String id = message.getId(); Sinks.One sink = this.resultSinks.remove(id); Sinks.Many streamingSink = this.streamingSinks.get(id); @@ -459,7 +460,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { return; } - Map resultMap = message.getPayloadOrDefault(Collections.emptyMap()); + Map resultMap = message.getPayload(); ExecutionResult result = MapExecutionResult.from(resultMap); Sinks.EmitResult emitResult = (sink != null ? sink.tryEmitValue(result) : streamingSink.tryEmitNext(result)); @@ -475,7 +476,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { * Handle an "error" message, turning it into an {@link ExecutionResult} * for a single result response, or signaling an error to streams. */ - public void handleError(GraphQlWebSocketMessage message) { + public void handleError(GraphQlMessage message) { String id = message.getId(); Sinks.One sink = this.resultSinks.remove(id); Sinks.Many streamingSink = this.streamingSinks.remove(id); @@ -487,7 +488,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { return; } - List> payload = message.getPayloadOrDefault(Collections.emptyList()); + List> payload = message.getPayload(); Sinks.EmitResult emitResult; if (sink != null) { @@ -508,7 +509,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { /** * Handle a "complete" message. */ - public void handleComplete(GraphQlWebSocketMessage message) { + public void handleComplete(GraphQlMessage message) { Sinks.One resultSink = this.resultSinks.remove(message.getId()); Sinks.Many streamingResultSink = this.streamingSinks.remove(message.getId()); diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessage.java b/spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessage.java new file mode 100644 index 00000000..a1c696e7 --- /dev/null +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessage.java @@ -0,0 +1,221 @@ +/* + * Copyright 2002-2022 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.web.support; + +import java.util.Collections; +import java.util.List; +import java.util.Map; + +import graphql.ExecutionResult; +import graphql.GraphQLError; + +import org.springframework.graphql.GraphQlRequest; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; +import org.springframework.util.ObjectUtils; + + +/** + * Represents a GraphQL over WebSocket protocol message. + * + * @author Rossen Stoyanchev + * @since 1.0.0 + * @see GraphQL Over WebSocket Protocol + */ +public class GraphQlMessage { + + @Nullable + private String id; + + @Nullable + private GraphQlMessageType type; + + @Nullable + private Object payload; + + + /** + * Private constructor. See static factory methods. + */ + private GraphQlMessage(@Nullable String id, GraphQlMessageType type, @Nullable Object payload) { + Assert.notNull(type, "GraphQlMessageType is required"); + Assert.isTrue(payload != null || type.doesNotRequirePayload(), "Payload is required for [" + type + "]"); + this.id = id; + this.type = type; + this.payload = payload; + } + + + /** + * Constructor for deserialization. + */ + GraphQlMessage() { + this.type = GraphQlMessageType.NOT_SPECIFIED; + } + + + /** + * Return the request id that is applicable to messages associated with a + * request, or {@code null} for connection level messages. + */ + @Nullable + public String getId() { + return this.id; + } + + /** + * Return the message type value as it should appear on the wire. + */ + public String getType() { + Assert.notNull(this.type, "Type is required"); + return this.type.value(); + } + + /** + * Return the message type as an emum. + */ + public GraphQlMessageType resolvedType() { + Assert.state(this.type != null, "GraphQlWebSocketMessage does not have a type"); + return this.type; + } + + /** + * Return the payload. For a deserialized message, this is typically a + * {@code Map} or {@code List} for an {@code "error"} message. + */ + @SuppressWarnings("unchecked") + public

P getPayload() { + if (this.payload == null) { + Assert.state(resolvedType().doesNotRequirePayload(), this.type + " requires a payload"); + return (P) Collections.emptyMap(); + } + return (P) this.payload; + } + + public void setId(@Nullable String id) { + this.id = id; + } + + public void setType(String type) { + this.type = GraphQlMessageType.fromValue(type); + } + + public void setPayload(@Nullable Object payload) { + this.payload = payload; + } + + + @Override + public int hashCode() { + int hashCode = (this.type != null ? this.type.hashCode() : 0); + hashCode = 31 * hashCode + ObjectUtils.nullSafeHashCode(this.id); + hashCode = 31 * hashCode + ObjectUtils.nullSafeHashCode(this.payload); + return hashCode; + } + + @Override + public boolean equals(Object o) { + if (!(o instanceof GraphQlMessage)) { + return false; + } + GraphQlMessage other = (GraphQlMessage) o; + return (ObjectUtils.nullSafeEquals(this.type, other.type) && + (ObjectUtils.nullSafeEquals(this.id, other.id) || (this.id == null && other.id == null)) && + (ObjectUtils.nullSafeEquals(getPayload(), other.getPayload()))); + } + + @Override + public String toString() { + return "GraphQlWebSocketMessage[" + + (this.id != null ? "id=\"" + this.id + "\"" + ", " : "") + + "type=\"" + this.type + "\"" + + (this.payload != null ? ", payload=" + this.payload : "") + "]"; + } + + + /** + * Create a {@code "connection_init"} client message. + * @param payload an optional payload + */ + public static GraphQlMessage connectionInit(@Nullable Object payload) { + return new GraphQlMessage(null, GraphQlMessageType.CONNECTION_INIT, payload); + } + + /** + * Create a {@code "connection_ack"} server message. + * @param payload an optional payload + */ + public static GraphQlMessage connectionAck(@Nullable Object payload) { + return new GraphQlMessage(null, GraphQlMessageType.CONNECTION_ACK, payload); + } + + /** + * Create a {@code "subscribe"} client message. + * @param id unique request id + * @param request the request to add as the message payload + */ + public static GraphQlMessage subscribe(String id, GraphQlRequest request) { + Assert.notNull(request, "GraphQlRequest is required"); + return new GraphQlMessage(id, GraphQlMessageType.SUBSCRIBE, request.toMap()); + } + + /** + * Create a {@code "next"} server message. + * @param id unique request id + * @param result the result from request execution to add as the message payload + */ + public static GraphQlMessage next(String id, ExecutionResult result) { + Assert.notNull(result, "ExecutionResult is required"); + return new GraphQlMessage(id, GraphQlMessageType.NEXT, result.toSpecification()); + } + + /** + * Create an {@code "error"} server message. + * @param id unique request id + * @param error the error to add as the message payload + */ + public static GraphQlMessage error(String id, GraphQLError error) { + Assert.notNull(error, "GraphQlError is required"); + List> errors = Collections.singletonList(error.toSpecification()); + return new GraphQlMessage(id, GraphQlMessageType.ERROR, errors); + } + + /** + * Create a {@code "complete"} server message. + * @param id unique request id + */ + public static GraphQlMessage complete(String id) { + return new GraphQlMessage(id, GraphQlMessageType.COMPLETE, null); + } + + /** + * Create a {@code "ping"} client or server message. + * @param payload an optional payload + */ + public static GraphQlMessage ping(@Nullable Object payload) { + return new GraphQlMessage(null, GraphQlMessageType.PING, payload); + } + + /** + * Create a {@code "pong"} client or server message. + * @param payload an optional payload + */ + public static GraphQlMessage pong(@Nullable Object payload) { + return new GraphQlMessage(null, GraphQlMessageType.PONG, payload); + } + +} diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessageType.java b/spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessageType.java new file mode 100644 index 00000000..c9ca68c8 --- /dev/null +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/support/GraphQlMessageType.java @@ -0,0 +1,99 @@ +/* + * Copyright 2002-2022 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.web.support; + + +/** + * Enum for a message type as defined in the GraphQL over WebSocket spec proposal. + * + * @author Rossen Stoyanchev + * @since 1.0.0 + * @see GraphQL Over WebSocket Protocol + */ +public enum GraphQlMessageType { + + CONNECTION_INIT("connection_init", false), + + CONNECTION_ACK("connection_ack", false), + + PING("ping", false), + + PONG("pong", false), + + SUBSCRIBE("subscribe", true), + + NEXT("next", true), + + ERROR("error", true), + + COMPLETE("complete", false), + + /** + * Indicates the GraphQL message did not have a message type. + */ + NOT_SPECIFIED("", false); + + + private static final GraphQlMessageType[] VALUES; + + static { + VALUES = values(); + } + + + private final String value; + + private final boolean requiresPayload; + + + GraphQlMessageType(String value, boolean requiresPayload) { + this.value = value; + this.requiresPayload = requiresPayload; + } + + + /** + * The protocol value for the message type. + */ + public String value() { + return this.value; + } + + /** + * Return {@code } if the message type has a payload, and it is required. + */ + public boolean doesNotRequirePayload() { + return !this.requiresPayload; + } + + + public static GraphQlMessageType fromValue(String value) { + for (GraphQlMessageType type : VALUES) { + if (type.value.equals(value)) { + return type; + } + } + throw new IllegalArgumentException("No matching constant for [" + value + "]"); + } + + + @Override + public String toString() { + return this.value; + } + +} diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/support/package-info.java b/spring-graphql/src/main/java/org/springframework/graphql/web/support/package-info.java new file mode 100644 index 00000000..39fc277e --- /dev/null +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/support/package-info.java @@ -0,0 +1,25 @@ +/* + * 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. + * 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. + */ + +/** + * Support classes for Web transports. + */ +@NonNullApi +@NonNullFields +package org.springframework.graphql.web.support; + +import org.springframework.lang.NonNullApi; +import org.springframework.lang.NonNullFields; diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/CodecDelegate.java b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/CodecDelegate.java index a9ca8714..9cd05e9f 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/CodecDelegate.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/CodecDelegate.java @@ -24,6 +24,7 @@ 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.graphql.web.support.GraphQlMessage; import org.springframework.http.MediaType; import org.springframework.http.codec.CodecConfigurer; import org.springframework.http.codec.DecoderHttpMessageReader; @@ -42,7 +43,7 @@ import org.springframework.web.reactive.socket.WebSocketSession; */ final class CodecDelegate { - private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlWebSocketMessage.class); + private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlMessage.class); private final Decoder decoder; @@ -73,25 +74,8 @@ final class CodecDelegate { } - public WebSocketMessage encodeConnectionAck(WebSocketSession session, Object ackPayload) { - return encode(session, GraphQlWebSocketMessage.connectionAck(ackPayload)); - } - - public WebSocketMessage encodeNext(WebSocketSession session, String id, ExecutionResult result) { - return encode(session, GraphQlWebSocketMessage.next(id, result)); - } - - public WebSocketMessage encodeError(WebSocketSession session, String id, Throwable ex) { - GraphQLError error = GraphqlErrorBuilder.newError().message(ex.getMessage()).build(); - return encode(session, GraphQlWebSocketMessage.error(id, error)); - } - - public WebSocketMessage encodeComplete(WebSocketSession session, String id) { - return encode(session, GraphQlWebSocketMessage.complete(id)); - } - @SuppressWarnings("unchecked") - public WebSocketMessage encode(WebSocketSession session, GraphQlWebSocketMessage message) { + public WebSocketMessage encode(WebSocketSession session, GraphQlMessage message) { DataBuffer buffer = ((Encoder) this.encoder).encodeValue( (T) message, session.bufferFactory(), MESSAGE_TYPE, MimeTypeUtils.APPLICATION_JSON, null); @@ -100,9 +84,27 @@ final class CodecDelegate { } @SuppressWarnings("ConstantConditions") - public GraphQlWebSocketMessage decode(WebSocketMessage webSocketMessage) { + public GraphQlMessage decode(WebSocketMessage webSocketMessage) { DataBuffer buffer = DataBufferUtils.retain(webSocketMessage.getPayload()); - return (GraphQlWebSocketMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null); + return (GraphQlMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null); } + public WebSocketMessage encodeConnectionAck(WebSocketSession session, Object ackPayload) { + return encode(session, GraphQlMessage.connectionAck(ackPayload)); + } + + public WebSocketMessage encodeNext(WebSocketSession session, String id, ExecutionResult result) { + return encode(session, GraphQlMessage.next(id, result)); + } + + public WebSocketMessage encodeError(WebSocketSession session, String id, Throwable ex) { + GraphQLError error = GraphqlErrorBuilder.newError().message(ex.getMessage()).build(); + return encode(session, GraphQlMessage.error(id, error)); + } + + public WebSocketMessage encodeComplete(WebSocketSession session, String id) { + return encode(session, GraphQlMessage.complete(id)); + } + + } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java index 982a5759..12909254 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java @@ -36,6 +36,7 @@ import org.springframework.graphql.web.WebGraphQlHandler; import org.springframework.graphql.web.WebInput; import org.springframework.graphql.web.WebOutput; import org.springframework.graphql.web.WebSocketInterceptor; +import org.springframework.graphql.web.support.GraphQlMessage; import org.springframework.http.codec.CodecConfigurer; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; @@ -73,8 +74,8 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { * Create a new instance. * @param graphQlHandler common handler for GraphQL over WebSocket requests * @param codecConfigurer codec configurer for JSON encoding and decoding - * @param connectionInitTimeout the time within which the {@code CONNECTION_INIT} type - * message must be received. + * @param connectionInitTimeout how long to wait after the establishment of + * the WebSocket for the {@code "connection_ini"} message from the client. */ public GraphQlWebSocketHandler( WebGraphQlHandler graphQlHandler, CodecConfigurer codecConfigurer, Duration connectionInitTimeout) { @@ -127,11 +128,11 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { .subscribe(); return session.send(session.receive().flatMap(webSocketMessage -> { - GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage); + GraphQlMessage message = this.codecDelegate.decode(webSocketMessage); String id = message.getId(); - Map payload = message.getPayloadOrDefault(Collections.emptyMap()); - switch (message.getType()) { - case "subscribe": + Map payload = message.getPayload(); + switch (message.resolvedType()) { + case SUBSCRIBE: if (connectionInitPayloadRef.get() == null) { return GraphQlStatus.close(session, GraphQlStatus.UNAUTHORIZED_STATUS); } @@ -146,7 +147,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { return this.graphQlHandler.handleRequest(input) .flatMapMany((output) -> handleWebOutput(session, id, subscriptions, output)) .doOnTerminate(() -> subscriptions.remove(id)); - case "complete": + case COMPLETE: if (id != null) { Subscription subscription = subscriptions.remove(id); if (subscription != null) { @@ -156,7 +157,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { .thenMany(Flux.empty()); } return Flux.empty(); - case "connection_init": + case CONNECTION_INIT: if (!connectionInitPayloadRef.compareAndSet(null, payload)) { return GraphQlStatus.close(session, GraphQlStatus.TOO_MANY_INIT_REQUESTS_STATUS); } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketMessage.java b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketMessage.java deleted file mode 100644 index 261b0681..00000000 --- a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketMessage.java +++ /dev/null @@ -1,165 +0,0 @@ -/* - * Copyright 2002-2022 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.web.webflux; - -import java.util.Collections; - -import graphql.ExecutionResult; -import graphql.GraphQLError; - -import org.springframework.graphql.GraphQlRequest; -import org.springframework.lang.Nullable; -import org.springframework.util.ObjectUtils; - -/** - * Representation of a GraphQL over WebSocket protocol message. - * - * @author Rossen Stoyanchev - * @since 1.0.0 - */ -public class GraphQlWebSocketMessage { - - @Nullable - private String id; - - private String type; - - @Nullable - private Object payload; - - - /** - * Private constructor for static factory methods. - */ - private GraphQlWebSocketMessage(@Nullable String id, String type, @Nullable Object payload) { - this.id = id; - this.type = type; - this.payload = payload; - } - - /** - * Constructor for deserialization. - */ - GraphQlWebSocketMessage() { - this.type = ""; - } - - - @Nullable - public String getId() { - return this.id; - } - - public String getType() { - return this.type; - } - - @SuppressWarnings("unchecked") - @Nullable - public

P getPayload() { - return (P) this.payload; - } - - @SuppressWarnings("unchecked") - public

P getPayloadOrDefault(P defaultPayload) { - return (this.payload != null ? (P) this.payload : defaultPayload); - } - - public void setId(@Nullable String id) { - this.id = id; - } - - public void setType(String type) { - this.type = type; - } - - public void setPayload(@Nullable Object payload) { - this.payload = payload; - } - - - @Override - public int hashCode() { - int hashCode = this.type.hashCode(); - hashCode = 31 * hashCode + ObjectUtils.nullSafeHashCode(this.id); - hashCode = 31 * hashCode + ObjectUtils.nullSafeHashCode(this.payload); - return hashCode; - } - - @Override - public boolean equals(Object o) { - if (!(o instanceof GraphQlWebSocketMessage)) { - return false; - } - GraphQlWebSocketMessage other = (GraphQlWebSocketMessage) o; - return (this.type.equals(other.type) && - (ObjectUtils.nullSafeEquals(this.id, other.id) || (this.id == null && other.id == null)) && - (ObjectUtils.nullSafeEquals(this.payload, other.payload) || (this.payload == null && other.payload == null))); - } - - @Override - public String toString() { - return "GraphQlWebSocketMessage[" + - (this.id != null ? "id=\"" + this.id + "\"" + ", " : "") + - "type=\"" + this.type + "\"" + - (this.payload != null ? ", payload=" + this.payload : "") + "]"; - } - - - /** - * Create a "connection_init" message. - */ - public static GraphQlWebSocketMessage connectionInit(@Nullable Object payload) { - return new GraphQlWebSocketMessage(null, "connection_init", payload); - } - - /** - * Create a "connection_ack" message. - */ - public static GraphQlWebSocketMessage connectionAck(@Nullable Object payload) { - return new GraphQlWebSocketMessage(null, "connection_ack", payload); - } - - /** - * Create a "subscribe" message. - */ - public static GraphQlWebSocketMessage subscribe(String id, GraphQlRequest request) { - return new GraphQlWebSocketMessage(id, "subscribe", request.toMap()); - } - - /** - * Create a "next" message. - */ - public static GraphQlWebSocketMessage next(String id, ExecutionResult result) { - return new GraphQlWebSocketMessage(id, "next", result.toSpecification()); - } - - /** - * Create an "error" message. - */ - public static GraphQlWebSocketMessage error(String id, GraphQLError error) { - return new GraphQlWebSocketMessage(id, "error", Collections.singletonList(error.toSpecification())); - } - - /** - * Create a "complete" message. - */ - public static GraphQlWebSocketMessage complete(String id) { - return new GraphQlWebSocketMessage(id, "complete", null); - } - -} diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java index 27212cbf..b9a9cc11 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java @@ -47,7 +47,8 @@ import org.springframework.graphql.web.WebGraphQlHandler; import org.springframework.graphql.web.WebInput; import org.springframework.graphql.web.WebOutput; import org.springframework.graphql.web.WebSocketInterceptor; -import org.springframework.graphql.web.webflux.GraphQlWebSocketMessage; +import org.springframework.graphql.web.support.GraphQlMessage; +import org.springframework.graphql.web.support.GraphQlMessageType; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpInputMessage; import org.springframework.http.HttpOutputMessage; @@ -93,8 +94,8 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub * Create a new instance. * @param graphQlHandler common handler for GraphQL over WebSocket requests * @param converter for JSON encoding and decoding - * @param connectionInitTimeout the time within which the {@code CONNECTION_INIT} type - * message must be received. + * @param connectionInitTimeout how long to wait after the establishment of + * the WebSocket for the {@code "connection_ini"} message from the client. */ public GraphQlWebSocketHandler( WebGraphQlHandler graphQlHandler, HttpMessageConverter converter, Duration connectionInitTimeout) { @@ -139,74 +140,74 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub @Override protected void handleTextMessage(WebSocketSession session, TextMessage webSocketMessage) throws Exception { - GraphQlWebSocketMessage message = decode(webSocketMessage); + GraphQlMessage message = decode(webSocketMessage); String id = message.getId(); - Map payload = message.getPayloadOrDefault(Collections.emptyMap()); + Map payload = message.getPayload(); SessionState sessionState = getSessionInfo(session); - switch (message.getType()) { - case "subscribe": - if (sessionState.getConnectionInitPayload() == null) { - GraphQlStatus.closeSession(session, GraphQlStatus.UNAUTHORIZED_STATUS); - return; - } - if (id == null) { - GraphQlStatus.closeSession(session, GraphQlStatus.INVALID_MESSAGE_STATUS); - return; - } - URI uri = session.getUri(); - Assert.notNull(uri, "Expected handshake url"); - HttpHeaders headers = session.getHandshakeHeaders(); - WebInput input = new WebInput(uri, headers, payload, id, null); - if (logger.isDebugEnabled()) { - logger.debug("Executing: " + input); - } - this.graphQlHandler.handleRequest(input) - .flatMapMany((output) -> handleWebOutput(session, input.getId(), 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(); + switch (message.resolvedType()) { + case SUBSCRIBE: + if (sessionState.getConnectionInitPayload() == null) { + GraphQlStatus.closeSession(session, GraphQlStatus.UNAUTHORIZED_STATUS); + return; } - this.webSocketInterceptor.handleCancelledSubscription(session.getId(), id) - .block(Duration.ofSeconds(10)); - } - return; - case "connection_init": - if (!sessionState.setConnectionInitPayload(payload)) { - GraphQlStatus.closeSession(session, GraphQlStatus.TOO_MANY_INIT_REQUESTS_STATUS); + if (id == null) { + GraphQlStatus.closeSession(session, GraphQlStatus.INVALID_MESSAGE_STATUS); + return; + } + URI uri = session.getUri(); + Assert.notNull(uri, "Expected handshake url"); + HttpHeaders headers = session.getHandshakeHeaders(); + WebInput input = new WebInput(uri, headers, payload, id, null); + if (logger.isDebugEnabled()) { + logger.debug("Executing: " + input); + } + this.graphQlHandler.handleRequest(input) + .flatMapMany((output) -> handleWebOutput(session, input.getId(), output)) + .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread + .subscribe(new SendMessageSubscriber(id, session, sessionState)); return; - } - this.webSocketInterceptor.handleConnectionInitialization(session.getId(), payload) - .defaultIfEmpty(Collections.emptyMap()) - .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread - .doOnNext(ackPayload -> { - TextMessage outputMessage = encode(GraphQlWebSocketMessage.connectionAck(ackPayload)); - try { - session.sendMessage(outputMessage); - } - catch (IOException ex) { - throw new IllegalStateException(ex); - } - }) - .onErrorResume(ex -> { - GraphQlStatus.closeSession(session, GraphQlStatus.UNAUTHORIZED_STATUS); - return Mono.empty(); - }) - .block(Duration.ofSeconds(10)); - return; - default: - GraphQlStatus.closeSession(session, GraphQlStatus.INVALID_MESSAGE_STATUS); + case COMPLETE: + if (id != null) { + Subscription subscription = sessionState.getSubscriptions().remove(id); + if (subscription != null) { + subscription.cancel(); + } + this.webSocketInterceptor.handleCancelledSubscription(session.getId(), id) + .block(Duration.ofSeconds(10)); + } + return; + case CONNECTION_INIT: + if (!sessionState.setConnectionInitPayload(payload)) { + GraphQlStatus.closeSession(session, GraphQlStatus.TOO_MANY_INIT_REQUESTS_STATUS); + return; + } + this.webSocketInterceptor.handleConnectionInitialization(session.getId(), payload) + .defaultIfEmpty(Collections.emptyMap()) + .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread + .doOnNext(ackPayload -> { + TextMessage outputMessage = encode(GraphQlMessage.connectionAck(ackPayload)); + try { + session.sendMessage(outputMessage); + } + catch (IOException ex) { + throw new IllegalStateException(ex); + } + }) + .onErrorResume(ex -> { + GraphQlStatus.closeSession(session, GraphQlStatus.UNAUTHORIZED_STATUS); + return Mono.empty(); + }) + .block(Duration.ofSeconds(10)); + return; + default: + GraphQlStatus.closeSession(session, GraphQlStatus.INVALID_MESSAGE_STATUS); } } @SuppressWarnings("unchecked") - private GraphQlWebSocketMessage decode(TextMessage message) throws IOException { - return ((GenericHttpMessageConverter) this.converter) - .read(GraphQlWebSocketMessage.class, null, new HttpInputMessageAdapter(message)); + private GraphQlMessage decode(TextMessage message) throws IOException { + return ((GenericHttpMessageConverter) this.converter) + .read(GraphQlMessage.class, null, new HttpInputMessageAdapter(message)); } private SessionState getSessionInfo(WebSocketSession session) { @@ -240,8 +241,8 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub } return outputFlux - .map(result -> encode(GraphQlWebSocketMessage.next(id, result))) - .concatWith(Mono.fromCallable(() -> encode(GraphQlWebSocketMessage.complete(id)))) + .map(result -> encode(GraphQlMessage.next(id, result))) + .concatWith(Mono.fromCallable(() -> encode(GraphQlMessage.complete(id)))) .onErrorResume((ex) -> { if (ex instanceof SubscriptionExistsException) { CloseStatus status = new CloseStatus(4409, "Subscriber for " + id + " already exists"); @@ -250,12 +251,12 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub } String message = ex.getMessage(); GraphQLError error = GraphqlErrorBuilder.newError().message(message).build(); - return Mono.just(encode(GraphQlWebSocketMessage.error(id, error))); + return Mono.just(encode(GraphQlMessage.error(id, error))); }); } @SuppressWarnings("unchecked") - private TextMessage encode(GraphQlWebSocketMessage message) { + private TextMessage encode(GraphQlMessage message) { try { HttpOutputMessageAdapter outputMessage = new HttpOutputMessageAdapter(); ((HttpMessageConverter) this.converter).write((T) message, null, outputMessage); diff --git a/spring-graphql/src/test/java/org/springframework/graphql/client/MockGraphQlWebSocketServer.java b/spring-graphql/src/test/java/org/springframework/graphql/client/MockGraphQlWebSocketServer.java index b4245f76..2fecff65 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/client/MockGraphQlWebSocketServer.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/client/MockGraphQlWebSocketServer.java @@ -29,7 +29,7 @@ import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; import org.springframework.graphql.GraphQlRequest; -import org.springframework.graphql.web.webflux.GraphQlWebSocketMessage; +import org.springframework.graphql.web.support.GraphQlMessage; import org.springframework.lang.Nullable; import org.springframework.web.reactive.socket.WebSocketHandler; import org.springframework.web.reactive.socket.WebSocketSession; @@ -83,33 +83,33 @@ public final class MockGraphQlWebSocketServer implements WebSocketHandler { } @SuppressWarnings("SuspiciousMethodCalls") - private Publisher handleMessage(GraphQlWebSocketMessage message) { - if ("connection_init".equals(message.getType())) { - if (this.connectionInitHandler == null) { - return Flux.just(GraphQlWebSocketMessage.connectionAck(null)); - } - else { - Map payload = message.getPayload(); - return this.connectionInitHandler.apply(payload).map(GraphQlWebSocketMessage::connectionAck); - } + private Publisher handleMessage(GraphQlMessage message) { + switch (message.resolvedType()) { + case CONNECTION_INIT: + if (this.connectionInitHandler == null) { + return Flux.just(GraphQlMessage.connectionAck(null)); + } + else { + Map payload = message.getPayload(); + return this.connectionInitHandler.apply(payload).map(GraphQlMessage::connectionAck); + } + case SUBSCRIBE: + String id = message.getId(); + Exchange request = expectedExchanges.get(message.getPayload()); + if (id == null || request == null) { + return Flux.error(new IllegalStateException("Unexpected request: " + message)); + } + return request.getResponseFlux() + .map(result -> GraphQlMessage.next(id, result)) + .concatWithValues( + request.getError() != null ? + GraphQlMessage.error(id, request.getError()) : + GraphQlMessage.complete(id)); + case COMPLETE: + return Flux.empty(); + default: + return Flux.error(new IllegalStateException("Unexpected message: " + message)); } - if ("subscribe".equals(message.getType())) { - String id = message.getId(); - Exchange request = expectedExchanges.get(message.getPayload()); - if (id == null || request == null) { - return Flux.error(new IllegalStateException("Unexpected request: " + message)); - } - return request.getResponseFlux() - .map(result -> GraphQlWebSocketMessage.next(id, result)) - .concatWithValues( - request.getError() != null ? - GraphQlWebSocketMessage.error(id, request.getError()) : - GraphQlWebSocketMessage.complete(id)); - } - if ("complete".equals(message.getType())) { - return Flux.empty(); - } - return Flux.error(new IllegalStateException("Unexpected message: " + message)); } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java b/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java index 20370603..d458d5e5 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java @@ -34,10 +34,11 @@ import reactor.core.publisher.Mono; import reactor.test.StepVerifier; import org.springframework.graphql.GraphQlRequest; +import org.springframework.graphql.support.MapExecutionResult; import org.springframework.graphql.web.TestWebSocketClient; import org.springframework.graphql.web.TestWebSocketConnection; -import org.springframework.graphql.support.MapExecutionResult; -import org.springframework.graphql.web.webflux.GraphQlWebSocketMessage; +import org.springframework.graphql.web.support.GraphQlMessage; +import org.springframework.graphql.web.support.GraphQlMessageType; import org.springframework.http.HttpHeaders; import org.springframework.http.codec.ClientCodecConfigurer; import org.springframework.web.reactive.socket.CloseStatus; @@ -83,8 +84,8 @@ public class MockWebSocketGraphQlTransportTests { .verify(TIMEOUT); assertActualClientMessages( - GraphQlWebSocketMessage.connectionInit(null), - GraphQlWebSocketMessage.subscribe("1", request)); + GraphQlMessage.connectionInit(null), + GraphQlMessage.subscribe("1", request)); } @Test @@ -96,8 +97,8 @@ public class MockWebSocketGraphQlTransportTests { .verify(TIMEOUT); assertActualClientMessages( - GraphQlWebSocketMessage.connectionInit(null), - GraphQlWebSocketMessage.subscribe("1", request)); + GraphQlMessage.connectionInit(null), + GraphQlMessage.subscribe("1", request)); } @Test @@ -114,8 +115,8 @@ public class MockWebSocketGraphQlTransportTests { .verify(TIMEOUT); assertActualClientMessages( - GraphQlWebSocketMessage.connectionInit(null), - GraphQlWebSocketMessage.subscribe("1", request)); + GraphQlMessage.connectionInit(null), + GraphQlMessage.subscribe("1", request)); } @Test @@ -132,8 +133,8 @@ public class MockWebSocketGraphQlTransportTests { .verify(TIMEOUT); assertActualClientMessages( - GraphQlWebSocketMessage.connectionInit(null), - GraphQlWebSocketMessage.subscribe("1", request)); + GraphQlMessage.connectionInit(null), + GraphQlMessage.subscribe("1", request)); } @Test @@ -146,8 +147,8 @@ public class MockWebSocketGraphQlTransportTests { .verify(TIMEOUT); assertActualClientMessages( - GraphQlWebSocketMessage.connectionInit(null), - GraphQlWebSocketMessage.subscribe("1", request)); + GraphQlMessage.connectionInit(null), + GraphQlMessage.subscribe("1", request)); } @Test @@ -162,9 +163,9 @@ public class MockWebSocketGraphQlTransportTests { .verify(TIMEOUT); assertActualClientMessages( - GraphQlWebSocketMessage.connectionInit(null), - GraphQlWebSocketMessage.subscribe("1", request), - GraphQlWebSocketMessage.complete("1")); + GraphQlMessage.connectionInit(null), + GraphQlMessage.subscribe("1", request), + GraphQlMessage.complete("1")); } @Test @@ -184,7 +185,7 @@ public class MockWebSocketGraphQlTransportTests { assertThat(client.getConnection(0).isOpen()).isTrue(); assertThat(connectionAckRef.get()).isEqualTo(Collections.singletonMap("key", "valueInitAck")); - assertActualClientMessages(client.getConnection(0), GraphQlWebSocketMessage.connectionInit(initPayload)); + assertActualClientMessages(client.getConnection(0), GraphQlMessage.connectionInit(initPayload)); } @Test @@ -251,7 +252,8 @@ public class MockWebSocketGraphQlTransportTests { IOException ex = new IOException("Connect failure"); WebSocketClient client = mock(WebSocketClient.class); - when(client.execute(any(URI.class), any(HttpHeaders.class), any(WebSocketHandler.class))).thenReturn(Mono.error(ex)); + when(client.execute(any(URI.class), any(HttpHeaders.class), any(WebSocketHandler.class))) + .thenReturn(Mono.error(ex)); StepVerifier.create(createTransport(client).start()) .expectErrorMessage(ex.getMessage()) @@ -293,14 +295,14 @@ public class MockWebSocketGraphQlTransportTests { URI.create("/"), HttpHeaders.EMPTY, client, ClientCodecConfigurer.create(), null, p -> {}); } - private void assertActualClientMessages(GraphQlWebSocketMessage... expectedMessages) { + private void assertActualClientMessages(GraphQlMessage... expectedMessages) { assertActualClientMessages(this.webSocketClient.getConnection(0), expectedMessages); } private void assertActualClientMessages( - TestWebSocketConnection connection, GraphQlWebSocketMessage... expectedMessages) { + TestWebSocketConnection connection, GraphQlMessage... expectedMessages) { - List actualMessages = connection.getClientMessages().stream() + List actualMessages = connection.getClientMessages().stream() .map(CODEC_DELEGATE::decode) .collect(Collectors.toList()); @@ -320,14 +322,14 @@ public class MockWebSocketGraphQlTransportTests { public Mono handle(WebSocketSession session) { return session.send(session.receive().flatMap(webSocketMessage -> { - GraphQlWebSocketMessage requestMessage = this.codecDelegate.decode(webSocketMessage); - String id = requestMessage.getId(); + GraphQlMessage inputMessage = this.codecDelegate.decode(webSocketMessage); + String id = inputMessage.getId(); - GraphQlWebSocketMessage responseMessage = (requestMessage.getType().equals("connection_init") ? - GraphQlWebSocketMessage.connectionAck(null) : - GraphQlWebSocketMessage.subscribe(id, new GraphQlRequest(""))); + GraphQlMessage outputMessage = (inputMessage.resolvedType() == GraphQlMessageType.CONNECTION_INIT ? + GraphQlMessage.connectionAck(null) : + GraphQlMessage.subscribe(id, new GraphQlRequest(""))); - return Flux.just(this.codecDelegate.encode(session, responseMessage)); + return Flux.just(this.codecDelegate.encode(session, outputMessage)); })); } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/web/WebSocketHandlerTestSupport.java b/spring-graphql/src/test/java/org/springframework/graphql/web/WebSocketHandlerTestSupport.java index 51e89472..de43c090 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/web/WebSocketHandlerTestSupport.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/web/WebSocketHandlerTestSupport.java @@ -25,20 +25,28 @@ public abstract class WebSocketHandlerTestSupport { protected static final String SUBSCRIPTION_ID = "1"; - protected static final String BOOK_QUERY = "{" + - "\"id\":\"" + WebSocketHandlerTestSupport.SUBSCRIPTION_ID + "\"," + - "\"type\":\"subscribe\"," + - "\"payload\":{\"query\": \"" + - " query TestQuery {" + - " bookById(id: \\\"1\\\"){ " + - " id" + - " name" + - " author {" + - " firstName" + - " lastName" + - " }" + - " }}\"}" + - "}"; + protected static final String BOOK_QUERY; + + protected static final String BOOK_QUERY_PAYLOAD; + + static { + BOOK_QUERY_PAYLOAD = "{\"query\": \"" + + " query TestQuery {" + + " bookById(id: \\\"1\\\"){ " + + " id" + + " name" + + " author {" + + " firstName" + + " lastName" + + " }" + + " }}\"}"; + + BOOK_QUERY = "{" + + "\"id\":\"" + WebSocketHandlerTestSupport.SUBSCRIPTION_ID + "\"," + + "\"type\":\"subscribe\"," + + "\"payload\":" + BOOK_QUERY_PAYLOAD + + "}"; + } protected static final String BOOK_SUBSCRIPTION = "{" + "\"id\":\"" + SUBSCRIPTION_ID + "\"," + diff --git a/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java index d51c5fad..1a36b3f9 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java @@ -42,6 +42,8 @@ import org.springframework.graphql.web.WebGraphQlHandler; import org.springframework.graphql.web.WebInterceptor; import org.springframework.graphql.web.WebSocketHandlerTestSupport; import org.springframework.graphql.web.WebSocketInterceptor; +import org.springframework.graphql.web.support.GraphQlMessage; +import org.springframework.graphql.web.support.GraphQlMessageType; import org.springframework.http.codec.ServerCodecConfigurer; import org.springframework.http.codec.json.Jackson2JsonDecoder; import org.springframework.web.reactive.socket.CloseStatus; @@ -64,17 +66,17 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { toWebSocketMessage(BOOK_QUERY))); StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("next"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT); assertThat(actual.>getPayload()) .extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class))) .extractingByKey("bookById", as(InstanceOfAssertFactories.map(String.class, Object.class))) .containsEntry("name", "Nineteen Eighty-Four"); }) - .consumeNextWith((message) -> assertMessageType(message, "complete")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) .verifyComplete(); } @@ -85,9 +87,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { toWebSocketMessage(BOOK_SUBSCRIPTION))); BiConsumer bookPayloadAssertion = (message, bookId) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("next"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT); assertThat(actual.>getPayload()) .extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class))) .extractingByKey("bookSearch", as(InstanceOfAssertFactories.map(String.class, Object.class))) @@ -95,10 +97,10 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { }; StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "1")) .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5")) - .consumeNextWith((message) -> assertMessageType(message, "complete")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) .verifyComplete(); } @@ -106,10 +108,10 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { void unauthorizedWithoutMessageType() { TestWebSocketSession session = handle(Flux.just( toWebSocketMessage("{\"type\":\"connection_init\"}"), - toWebSocketMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\"}"))); + toWebSocketMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\", \"payload\":" + BOOK_QUERY_PAYLOAD + "}"))); StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .verifyComplete(); StepVerifier.create(session.closeStatus()) @@ -121,12 +123,12 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { void invalidMessageWithoutId() { Flux input = Flux.just( toWebSocketMessage("{\"type\":\"connection_init\"}"), - toWebSocketMessage("{\"type\":\"subscribe\"}")); // No message id + toWebSocketMessage("{\"type\":\"subscribe\", \"payload\":{}}")); // No message id TestWebSocketSession session = handle(input); StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .verifyComplete(); StepVerifier.create(session.closeStatus()) @@ -149,8 +151,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); - assertThat(actual.getType()).isEqualTo("connection_ack"); + GraphQlMessage actual = decode(message); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); assertThat(actual.>getPayload()).containsEntry("key", "A acknowledged"); }) .verifyComplete(); @@ -213,7 +215,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { toWebSocketMessage("{\"type\":\"connection_init\"}"))); StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .verifyComplete(); StepVerifier.create(session.closeStatus()) @@ -244,7 +246,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { TestWebSocketSession session = handle(messageFlux, new ConsumeOneAndNeverCompleteInterceptor()); // Collect messages until session closed - List messages = new ArrayList<>(); + List messages = new ArrayList<>(); session.getOutput().subscribe((message) -> messages.add(decode(message))); StepVerifier.create(session.closeStatus()) @@ -252,8 +254,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .verifyComplete(); assertThat(messages.size()).isEqualTo(2); - assertThat(messages.get(0).getType()).isEqualTo("connection_ack"); - assertThat(messages.get(1).getType()).isEqualTo("next"); + assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); + assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlMessageType.NEXT); } @Test @@ -267,12 +269,12 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { String completeMessage = "{\"id\":\"" + SUBSCRIPTION_ID + "\",\"type\":\"complete\"}"; StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) - .consumeNextWith((message) -> assertMessageType(message, "next")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.NEXT)) .then(() -> input.tryEmitNext(toWebSocketMessage(completeMessage))) .as("Second subscription with same id is possible only if the first was properly removed") .then(() -> input.tryEmitNext(toWebSocketMessage(BOOK_SUBSCRIPTION))) - .consumeNextWith((message) -> assertMessageType(message, "next")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.NEXT)) .then(() -> input.tryEmitNext(toWebSocketMessage(completeMessage))) .verifyTimeout(Duration.ofMillis(500)); } @@ -306,19 +308,19 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { handler.handle(session).block(); StepVerifier.create(session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("next"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT); assertThat(actual.>getPayload()) .extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class))) .containsEntry("greeting", "a"); }) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("error"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.ERROR); assertThat(actual.>>getPayload()) .asList().hasSize(1) .allSatisfy(theError -> assertThat(theError) @@ -349,15 +351,15 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { } @SuppressWarnings("ConstantConditions") - private GraphQlWebSocketMessage decode(WebSocketMessage message) { - return (GraphQlWebSocketMessage) decoder.decode(DataBufferUtils.retain(message.getPayload()), - ResolvableType.forClass(GraphQlWebSocketMessage.class), null, Collections.emptyMap()); + private GraphQlMessage decode(WebSocketMessage message) { + return (GraphQlMessage) decoder.decode(DataBufferUtils.retain(message.getPayload()), + ResolvableType.forClass(GraphQlMessage.class), null, Collections.emptyMap()); } - private void assertMessageType(WebSocketMessage webSocketMessage, String messageType) { - GraphQlWebSocketMessage message = decode(webSocketMessage); - assertThat(message.getType()).isEqualTo(messageType); - if (!messageType.equals("connection_ack")) { + private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlMessageType messageType) { + GraphQlMessage message = decode(webSocketMessage); + assertThat(message.resolvedType()).isEqualTo(messageType); + if (messageType != GraphQlMessageType.CONNECTION_ACK) { assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID); } } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java index 3674adc1..1c9afac4 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java @@ -40,7 +40,8 @@ import org.springframework.graphql.web.WebGraphQlHandler; import org.springframework.graphql.web.WebInterceptor; import org.springframework.graphql.web.WebSocketHandlerTestSupport; import org.springframework.graphql.web.WebSocketInterceptor; -import org.springframework.graphql.web.webflux.GraphQlWebSocketMessage; +import org.springframework.graphql.web.support.GraphQlMessage; +import org.springframework.graphql.web.support.GraphQlMessageType; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpInputMessage; import org.springframework.http.converter.GenericHttpMessageConverter; @@ -71,17 +72,17 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new TextMessage(BOOK_QUERY)); StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("next"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT); assertThat(actual.>getPayload()) .extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class))) .extractingByKey("bookById", as(InstanceOfAssertFactories.map(String.class, Object.class))) .containsEntry("name", "Nineteen Eighty-Four"); }) - .consumeNextWith((message) -> assertMessageType(message, "complete")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) .then(this.session::close) // Complete output Flux .verifyComplete(); } @@ -91,9 +92,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { handle(this.handler, new TextMessage("{\"type\":\"connection_init\"}"), new TextMessage(BOOK_SUBSCRIPTION)); BiConsumer, String> bookPayloadAssertion = (message, bookId) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("next"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT); assertThat(actual.>getPayload()) .extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class))) .extractingByKey("bookSearch", as(InstanceOfAssertFactories.map(String.class, Object.class))) @@ -101,10 +102,10 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { }; StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "1")) .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5")) - .consumeNextWith((message) -> assertMessageType(message, "complete")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) .then(this.session::close)// Complete output Flux .verifyComplete(); } @@ -113,11 +114,11 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { void unauthorizedWithoutMessageType() throws Exception { handle(this.handler, new TextMessage("{\"type\":\"connection_init\"}"), - new TextMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\"}")); + new TextMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\", \"payload\":" + BOOK_QUERY_PAYLOAD + "}")); // No message type StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .verifyComplete(); assertThat(this.session.getCloseStatus()).isEqualTo(new CloseStatus(4400, "Invalid message")); @@ -127,10 +128,10 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { void invalidMessageWithoutId() throws Exception { handle(this.handler, new TextMessage("{\"type\":\"connection_init\"}"), - new TextMessage("{\"type\":\"subscribe\"}")); // No message id + new TextMessage("{\"type\":\"subscribe\", \"payload\":{}}")); // No message id StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .verifyComplete(); assertThat(this.session.getCloseStatus()).isEqualTo(new CloseStatus(4400, "Invalid message")); @@ -153,8 +154,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()) .consumeNextWith((webSocketMessage) -> { - GraphQlWebSocketMessage message = decode(webSocketMessage); - assertThat(message.getType()).isEqualTo("connection_ack"); + GraphQlMessage message = decode(webSocketMessage); + assertThat(message.resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); assertThat(message.>getPayload()).containsEntry("key", "A acknowledged"); }) .then(this.session::close) // Complete output Flux @@ -223,7 +224,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new TextMessage("{\"type\":\"connection_init\"}")); StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .verifyComplete(); assertThat(this.session.getCloseStatus()).isEqualTo(new CloseStatus(4429, "Too many initialisation requests")); @@ -247,7 +248,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new TextMessage(BOOK_SUBSCRIPTION)); // Collect messages until session closed - List messages = new ArrayList<>(); + List messages = new ArrayList<>(); this.session.getOutput().subscribe((message) -> messages.add(decode(message))); StepVerifier.create(this.session.closeStatus()) @@ -255,8 +256,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .verifyComplete(); assertThat(messages.size()).isEqualTo(2); - assertThat(messages.get(0).getType()).isEqualTo("connection_ack"); - assertThat(messages.get(1).getType()).isEqualTo("next"); + assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); + assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlMessageType.NEXT); } @Test @@ -278,12 +279,12 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { }; StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) - .consumeNextWith((message) -> assertMessageType(message, "next")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.NEXT)) .then(() -> messageSender.accept(completeMessage)) .as("Second subscription with same id is possible only if the first was properly removed") .then(() -> messageSender.accept(BOOK_SUBSCRIPTION)) - .consumeNextWith((message) -> assertMessageType(message, "next")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.NEXT)) .then(() -> messageSender.accept(completeMessage)) .verifyTimeout(Duration.ofMillis(500)); } @@ -313,19 +314,19 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new TextMessage(GREETING_QUERY)); StepVerifier.create(this.session.getOutput()) - .consumeNextWith((message) -> assertMessageType(message, "connection_ack")) + .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("next"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT); assertThat(actual.>getPayload()) .extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class))) .containsEntry("greeting", "a"); }) .consumeNextWith((message) -> { - GraphQlWebSocketMessage actual = decode(message); + GraphQlMessage actual = decode(message); assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID); - assertThat(actual.getType()).isEqualTo("error"); + assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.ERROR); assertThat(actual.>>getPayload()) .asList().hasSize(1) .allSatisfy(theError -> assertThat(theError) @@ -357,21 +358,21 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { } @SuppressWarnings("unchecked") - private GraphQlWebSocketMessage decode(WebSocketMessage message) { + private GraphQlMessage decode(WebSocketMessage message) { try { HttpInputMessageAdapter inputMessage = new HttpInputMessageAdapter((TextMessage) message); - return ((GenericHttpMessageConverter) converter) - .read(GraphQlWebSocketMessage.class, null, inputMessage); + return ((GenericHttpMessageConverter) converter) + .read(GraphQlMessage.class, null, inputMessage); } catch (IOException ex) { throw new IllegalStateException(ex); } } - private void assertMessageType(WebSocketMessage webSocketMessage, String messageType) { - GraphQlWebSocketMessage message = decode(webSocketMessage); - assertThat(message.getType()).isEqualTo(messageType); - if (!messageType.equals("connection_ack")) { + private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlMessageType messageType) { + GraphQlMessage message = decode(webSocketMessage); + assertThat(message.resolvedType()).isEqualTo(messageType); + if (messageType != GraphQlMessageType.CONNECTION_ACK) { assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID); } }