Refactor GraphQlMessage and add GraphQlMessageType enum

See gh-270
This commit is contained in:
rstoyanchev
2022-03-08 12:50:32 +00:00
parent ec00547130
commit 878c643672
14 changed files with 619 additions and 421 deletions

View File

@@ -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 <T> WebSocketMessage encode(WebSocketSession session, GraphQlWebSocketMessage message) {
public <T> WebSocketMessage encode(WebSocketSession session, GraphQlMessage message) {
DataBuffer buffer = ((Encoder<T>) 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);
}
}

View File

@@ -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<Map<String, Object>> connectionAckHandler;
@@ -181,7 +182,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
@Nullable Object connectionInitPayload, Consumer<Map<String, Object>> 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<GraphQlWebSocketMessage> requestSink = Sinks.many().unicast().onBackpressureBuffer();
private final Sinks.Many<GraphQlMessage> requestSink = Sinks.many().unicast().onBackpressureBuffer();
private final Map<String, Sinks.One<ExecutionResult>> 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<GraphQlWebSocketMessage> getRequestFlux() {
public Flux<GraphQlMessage> getRequestFlux() {
return this.requestSink.asFlux();
}
public Mono<ExecutionResult> 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<ExecutionResult> sink = Sinks.one();
this.resultSinks.put(id, sink);
trySend(message);
@@ -403,7 +404,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
public Flux<ExecutionResult> 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<ExecutionResult> 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<ExecutionResult> 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<ExecutionResult> sink = this.resultSinks.remove(id);
Sinks.Many<ExecutionResult> streamingSink = this.streamingSinks.get(id);
@@ -459,7 +460,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
return;
}
Map<String, Object> resultMap = message.getPayloadOrDefault(Collections.emptyMap());
Map<String, Object> 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<ExecutionResult> sink = this.resultSinks.remove(id);
Sinks.Many<ExecutionResult> streamingSink = this.streamingSinks.remove(id);
@@ -487,7 +488,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
return;
}
List<Map<String, Object>> payload = message.getPayloadOrDefault(Collections.emptyList());
List<Map<String, Object>> 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<ExecutionResult> resultSink = this.resultSinks.remove(message.getId());
Sinks.Many<ExecutionResult> streamingResultSink = this.streamingSinks.remove(message.getId());

View File

@@ -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 <a href="https://github.com/enisdenjo/graphql-ws/blob/master/PROTOCOL.md">GraphQL Over WebSocket Protocol</a>
*/
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> 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<Map<String, Object>> 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);
}
}

View File

@@ -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 <a href="https://github.com/enisdenjo/graphql-ws/blob/master/PROTOCOL.md">GraphQL Over WebSocket Protocol</a>
*/
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;
}
}

View File

@@ -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;

View File

@@ -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 <T> WebSocketMessage encode(WebSocketSession session, GraphQlWebSocketMessage message) {
public <T> WebSocketMessage encode(WebSocketSession session, GraphQlMessage message) {
DataBuffer buffer = ((Encoder<T>) 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));
}
}

View File

@@ -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<String, Object> payload = message.getPayloadOrDefault(Collections.emptyMap());
switch (message.getType()) {
case "subscribe":
Map<String, Object> 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);
}

View File

@@ -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> P getPayload() {
return (P) this.payload;
}
@SuppressWarnings("unchecked")
public <P> 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);
}
}

View File

@@ -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<String, Object> payload = message.getPayloadOrDefault(Collections.emptyMap());
Map<String, Object> 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<GraphQlWebSocketMessage>) this.converter)
.read(GraphQlWebSocketMessage.class, null, new HttpInputMessageAdapter(message));
private GraphQlMessage decode(TextMessage message) throws IOException {
return ((GenericHttpMessageConverter<GraphQlMessage>) 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 <T> TextMessage encode(GraphQlWebSocketMessage message) {
private <T> TextMessage encode(GraphQlMessage message) {
try {
HttpOutputMessageAdapter outputMessage = new HttpOutputMessageAdapter();
((HttpMessageConverter<T>) this.converter).write((T) message, null, outputMessage);

View File

@@ -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<GraphQlWebSocketMessage> handleMessage(GraphQlWebSocketMessage message) {
if ("connection_init".equals(message.getType())) {
if (this.connectionInitHandler == null) {
return Flux.just(GraphQlWebSocketMessage.connectionAck(null));
}
else {
Map<String, Object> payload = message.getPayload();
return this.connectionInitHandler.apply(payload).map(GraphQlWebSocketMessage::connectionAck);
}
private Publisher<GraphQlMessage> handleMessage(GraphQlMessage message) {
switch (message.resolvedType()) {
case CONNECTION_INIT:
if (this.connectionInitHandler == null) {
return Flux.just(GraphQlMessage.connectionAck(null));
}
else {
Map<String, Object> 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));
}

View File

@@ -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<GraphQlWebSocketMessage> actualMessages = connection.getClientMessages().stream()
List<GraphQlMessage> actualMessages = connection.getClientMessages().stream()
.map(CODEC_DELEGATE::decode)
.collect(Collectors.toList());
@@ -320,14 +322,14 @@ public class MockWebSocketGraphQlTransportTests {
public Mono<Void> 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));
}));
}

View File

@@ -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 + "\"," +

View File

@@ -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.<Map<String, Object>>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<WebSocketMessage, 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.<Map<String, Object>>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<WebSocketMessage> 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.<Map<String, Object>>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<GraphQlWebSocketMessage> messages = new ArrayList<>();
List<GraphQlMessage> 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.<Map<String, Object>>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.<List<Map<String, Object>>>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);
}
}

View File

@@ -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.<Map<String, Object>>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<WebSocketMessage<?>, 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.<Map<String, Object>>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.<Map<String, Object>>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<GraphQlWebSocketMessage> messages = new ArrayList<>();
List<GraphQlMessage> 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.<Map<String, Object>>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.<List<Map<String, Object>>>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<GraphQlWebSocketMessage>) converter)
.read(GraphQlWebSocketMessage.class, null, inputMessage);
return ((GenericHttpMessageConverter<GraphQlMessage>) 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);
}
}