Refactor GraphQlMessage and add GraphQlMessageType enum
See gh-270
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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));
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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));
|
||||
}));
|
||||
}
|
||||
|
||||
|
||||
@@ -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 + "\"," +
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user