Insert WebSocket in the names of GraphQlMessage[Type]

Those are specific to the GraphQL over WebSocket protocol.

See gh-339
This commit is contained in:
rstoyanchev
2022-03-25 21:59:07 +00:00
parent eb9a369f00
commit 3090328666
11 changed files with 182 additions and 179 deletions

View File

@@ -23,7 +23,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.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.http.MediaType;
import org.springframework.http.codec.ClientCodecConfigurer;
import org.springframework.http.codec.CodecConfigurer;
@@ -42,7 +42,7 @@ import org.springframework.web.reactive.socket.WebSocketSession;
*/
final class CodecDelegate {
private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlMessage.class);
private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlWebSocketMessage.class);
private final CodecConfigurer codecConfigurer;
@@ -104,7 +104,7 @@ final class CodecDelegate {
@SuppressWarnings("unchecked")
public <T> WebSocketMessage encode(WebSocketSession session, GraphQlMessage message) {
public <T> WebSocketMessage encode(WebSocketSession session, GraphQlWebSocketMessage message) {
DataBuffer buffer = ((Encoder<T>) this.encoder).encodeValue(
(T) message, session.bufferFactory(), MESSAGE_TYPE, MimeTypeUtils.APPLICATION_JSON, null);
@@ -113,9 +113,9 @@ final class CodecDelegate {
}
@SuppressWarnings("ConstantConditions")
public GraphQlMessage decode(WebSocketMessage webSocketMessage) {
public GraphQlWebSocketMessage decode(WebSocketMessage webSocketMessage) {
DataBuffer buffer = DataBufferUtils.retain(webSocketMessage.getPayload());
return (GraphQlMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null);
return (GraphQlWebSocketMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null);
}
}

View File

@@ -33,8 +33,8 @@ import reactor.core.publisher.Sinks;
import org.springframework.graphql.GraphQlRequest;
import org.springframework.graphql.GraphQlResponse;
import org.springframework.graphql.ResponseError;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlMessageType;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessageType;
import org.springframework.http.HttpHeaders;
import org.springframework.http.codec.CodecConfigurer;
import org.springframework.lang.Nullable;
@@ -226,9 +226,9 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
GraphQlSession graphQlSession = new GraphQlSession(session);
registerCloseStatusHandling(graphQlSession, session);
Mono<GraphQlMessage> connectionInitMono = this.interceptor.connectionInitPayload()
Mono<GraphQlWebSocketMessage> connectionInitMono = this.interceptor.connectionInitPayload()
.defaultIfEmpty(Collections.emptyMap())
.map(GraphQlMessage::connectionInit);
.map(GraphQlWebSocketMessage::connectionInit);
Mono<Void> sendCompletion =
session.send(connectionInitMono.concatWith(graphQlSession.getRequestFlux())
@@ -238,8 +238,8 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
.flatMap(webSocketMessage -> {
if (sessionNotInitialized()) {
try {
GraphQlMessage message = this.codecDelegate.decode(webSocketMessage);
Assert.state(message.resolvedType() == GraphQlMessageType.CONNECTION_ACK,
GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage);
Assert.state(message.resolvedType() == GraphQlWebSocketMessageType.CONNECTION_ACK,
() -> "Unexpected message before connection_ack: " + message);
return this.interceptor.handleConnectionAck(message.getPayload())
.then(Mono.defer(() -> {
@@ -261,7 +261,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
}
else {
try {
GraphQlMessage message = this.codecDelegate.decode(webSocketMessage);
GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage);
switch (message.resolvedType()) {
case NEXT:
graphQlSession.handleNext(message);
@@ -378,7 +378,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
private final AtomicLong requestIndex = new AtomicLong();
private final Sinks.Many<GraphQlMessage> requestSink = Sinks.many().unicast().onBackpressureBuffer();
private final Sinks.Many<GraphQlWebSocketMessage> requestSink = Sinks.many().unicast().onBackpressureBuffer();
private final Map<String, ResponseState> responseMap = new ConcurrentHashMap<>();
@@ -393,14 +393,14 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
/**
* Return the {@code Flux} of GraphQL requests to send as WebSocket messages.
*/
public Flux<GraphQlMessage> getRequestFlux() {
public Flux<GraphQlWebSocketMessage> getRequestFlux() {
return this.requestSink.asFlux();
}
public Mono<GraphQlResponse> execute(GraphQlRequest request) {
String id = String.valueOf(this.requestIndex.incrementAndGet());
try {
GraphQlMessage message = GraphQlMessage.subscribe(id, request);
GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request);
ResponseState state = new ResponseState(request);
this.responseMap.put(id, state);
trySend(message);
@@ -415,7 +415,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
public Flux<GraphQlResponse> executeSubscription(GraphQlRequest request) {
String id = String.valueOf(this.requestIndex.incrementAndGet());
try {
GraphQlMessage message = GraphQlMessage.subscribe(id, request);
GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request);
SubscriptionState state = new SubscriptionState(request);
this.subscriptionMap.put(id, state);
trySend(message);
@@ -428,13 +428,13 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
}
public void sendPong(@Nullable Map<String, Object> payload) {
GraphQlMessage message = GraphQlMessage.pong(payload);
GraphQlWebSocketMessage message = GraphQlWebSocketMessage.pong(payload);
trySend(message);
}
// TODO: queue to serialize sending?
private void trySend(GraphQlMessage message) {
private void trySend(GraphQlWebSocketMessage message) {
Sinks.EmitResult emitResult = null;
for (int i = 0; i < 100; i++) {
emitResult = this.requestSink.tryEmitNext(message);
@@ -449,7 +449,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
SubscriptionState state = this.subscriptionMap.remove(id);
if (state != null) {
try {
trySend(GraphQlMessage.complete(id));
trySend(GraphQlWebSocketMessage.complete(id));
}
catch (Exception ex) {
if (logger.isErrorEnabled()) {
@@ -465,7 +465,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
/**
* Handle a "next" message and route to its recipient.
*/
public void handleNext(GraphQlMessage message) {
public void handleNext(GraphQlWebSocketMessage message) {
String id = message.getId();
ResponseState responseState = this.responseMap.remove(id);
SubscriptionState subscriptionState = this.subscriptionMap.get(id);
@@ -496,7 +496,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
* Handle an "error" message, turning it into an {@link GraphQlResponse}
* for single responses, or signaling an error for streams.
*/
public void handleError(GraphQlMessage message) {
public void handleError(GraphQlWebSocketMessage message) {
String id = message.getId();
ResponseState responseState = this.responseMap.remove(id);
SubscriptionState subscriptionState = this.subscriptionMap.remove(id);
@@ -529,7 +529,7 @@ final class WebSocketGraphQlTransport implements GraphQlTransport {
/**
* Handle a "complete" message.
*/
public void handleComplete(GraphQlMessage message) {
public void handleComplete(GraphQlWebSocketMessage message) {
ResponseState responseState = this.responseMap.remove(message.getId());
SubscriptionState subscriptionState = this.subscriptionMap.remove(message.getId());

View File

@@ -35,13 +35,13 @@ import org.springframework.util.ObjectUtils;
* @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 {
public class GraphQlWebSocketMessage {
@Nullable
private String id;
@Nullable
private GraphQlMessageType type;
private GraphQlWebSocketMessageType type;
@Nullable
private Object payload;
@@ -50,7 +50,7 @@ public class GraphQlMessage {
/**
* Private constructor. See static factory methods.
*/
private GraphQlMessage(@Nullable String id, GraphQlMessageType type, @Nullable Object payload) {
private GraphQlWebSocketMessage(@Nullable String id, GraphQlWebSocketMessageType type, @Nullable Object payload) {
Assert.notNull(type, "GraphQlMessageType is required");
Assert.isTrue(payload != null || type.doesNotRequirePayload(), "Payload is required for [" + type + "]");
this.id = id;
@@ -63,8 +63,8 @@ public class GraphQlMessage {
* Constructor for deserialization.
*/
@SuppressWarnings("unused")
GraphQlMessage() {
this.type = GraphQlMessageType.NOT_SPECIFIED;
GraphQlWebSocketMessage() {
this.type = GraphQlWebSocketMessageType.NOT_SPECIFIED;
}
@@ -88,7 +88,7 @@ public class GraphQlMessage {
/**
* Return the message type as an emum.
*/
public GraphQlMessageType resolvedType() {
public GraphQlWebSocketMessageType resolvedType() {
Assert.state(this.type != null, "GraphQlWebSocketMessage does not have a type");
return this.type;
}
@@ -111,7 +111,7 @@ public class GraphQlMessage {
}
public void setType(String type) {
this.type = GraphQlMessageType.fromValue(type);
this.type = GraphQlWebSocketMessageType.fromValue(type);
}
public void setPayload(@Nullable Object payload) {
@@ -129,10 +129,10 @@ public class GraphQlMessage {
@Override
public boolean equals(Object o) {
if (!(o instanceof GraphQlMessage)) {
if (!(o instanceof GraphQlWebSocketMessage)) {
return false;
}
GraphQlMessage other = (GraphQlMessage) o;
GraphQlWebSocketMessage other = (GraphQlWebSocketMessage) 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())));
@@ -151,16 +151,16 @@ public class GraphQlMessage {
* 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);
public static GraphQlWebSocketMessage connectionInit(@Nullable Object payload) {
return new GraphQlWebSocketMessage(null, GraphQlWebSocketMessageType.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);
public static GraphQlWebSocketMessage connectionAck(@Nullable Object payload) {
return new GraphQlWebSocketMessage(null, GraphQlWebSocketMessageType.CONNECTION_ACK, payload);
}
/**
@@ -168,9 +168,9 @@ public class GraphQlMessage {
* @param id unique request id
* @param request the request to add as the message payload
*/
public static GraphQlMessage subscribe(String id, GraphQlRequest request) {
public static GraphQlWebSocketMessage subscribe(String id, GraphQlRequest request) {
Assert.notNull(request, "GraphQlRequest is required");
return new GraphQlMessage(id, GraphQlMessageType.SUBSCRIBE, request.toMap());
return new GraphQlWebSocketMessage(id, GraphQlWebSocketMessageType.SUBSCRIBE, request.toMap());
}
/**
@@ -178,9 +178,9 @@ public class GraphQlMessage {
* @param id unique request id
* @param responseMap the response map
*/
public static GraphQlMessage next(String id, Map<String, Object> responseMap) {
public static GraphQlWebSocketMessage next(String id, Map<String, Object> responseMap) {
Assert.notNull(responseMap, "'responseMap' is required");
return new GraphQlMessage(id, GraphQlMessageType.NEXT, responseMap);
return new GraphQlWebSocketMessage(id, GraphQlWebSocketMessageType.NEXT, responseMap);
}
/**
@@ -188,34 +188,34 @@ public class GraphQlMessage {
* @param id unique request id
* @param error the error to add as the message payload
*/
public static GraphQlMessage error(String id, GraphQLError error) {
public static GraphQlWebSocketMessage 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);
return new GraphQlWebSocketMessage(id, GraphQlWebSocketMessageType.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);
public static GraphQlWebSocketMessage complete(String id) {
return new GraphQlWebSocketMessage(id, GraphQlWebSocketMessageType.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);
public static GraphQlWebSocketMessage ping(@Nullable Object payload) {
return new GraphQlWebSocketMessage(null, GraphQlWebSocketMessageType.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);
public static GraphQlWebSocketMessage pong(@Nullable Object payload) {
return new GraphQlWebSocketMessage(null, GraphQlWebSocketMessageType.PONG, payload);
}
}

View File

@@ -24,7 +24,7 @@ package org.springframework.graphql.server.support;
* @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 {
public enum GraphQlWebSocketMessageType {
CONNECTION_INIT("connection_init", false),
@@ -48,7 +48,7 @@ public enum GraphQlMessageType {
NOT_SPECIFIED("", false);
private static final GraphQlMessageType[] VALUES;
private static final GraphQlWebSocketMessageType[] VALUES;
static {
VALUES = values();
@@ -60,7 +60,7 @@ public enum GraphQlMessageType {
private final boolean requiresPayload;
GraphQlMessageType(String value, boolean requiresPayload) {
GraphQlWebSocketMessageType(String value, boolean requiresPayload) {
this.value = value;
this.requiresPayload = requiresPayload;
}
@@ -81,8 +81,8 @@ public enum GraphQlMessageType {
}
public static GraphQlMessageType fromValue(String value) {
for (GraphQlMessageType type : VALUES) {
public static GraphQlWebSocketMessageType fromValue(String value) {
for (GraphQlWebSocketMessageType type : VALUES) {
if (type.value.equals(value)) {
return type;
}

View File

@@ -25,7 +25,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.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.http.MediaType;
import org.springframework.http.codec.CodecConfigurer;
import org.springframework.http.codec.DecoderHttpMessageReader;
@@ -44,7 +44,7 @@ import org.springframework.web.reactive.socket.WebSocketSession;
*/
final class CodecDelegate {
private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlMessage.class);
private static final ResolvableType MESSAGE_TYPE = ResolvableType.forClass(GraphQlWebSocketMessage.class);
private final Decoder<?> decoder;
@@ -76,7 +76,7 @@ final class CodecDelegate {
@SuppressWarnings("unchecked")
public <T> WebSocketMessage encode(WebSocketSession session, GraphQlMessage message) {
public <T> WebSocketMessage encode(WebSocketSession session, GraphQlWebSocketMessage message) {
DataBuffer buffer = ((Encoder<T>) this.encoder).encodeValue(
(T) message, session.bufferFactory(), MESSAGE_TYPE, MimeTypeUtils.APPLICATION_JSON, null);
@@ -85,26 +85,26 @@ final class CodecDelegate {
}
@SuppressWarnings("ConstantConditions")
public GraphQlMessage decode(WebSocketMessage webSocketMessage) {
public GraphQlWebSocketMessage decode(WebSocketMessage webSocketMessage) {
DataBuffer buffer = DataBufferUtils.retain(webSocketMessage.getPayload());
return (GraphQlMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null);
return (GraphQlWebSocketMessage) this.decoder.decode(buffer, MESSAGE_TYPE, null, null);
}
public WebSocketMessage encodeConnectionAck(WebSocketSession session, Object ackPayload) {
return encode(session, GraphQlMessage.connectionAck(ackPayload));
return encode(session, GraphQlWebSocketMessage.connectionAck(ackPayload));
}
public WebSocketMessage encodeNext(WebSocketSession session, String id, Map<String, Object> responseMap) {
return encode(session, GraphQlMessage.next(id, responseMap));
return encode(session, GraphQlWebSocketMessage.next(id, responseMap));
}
public WebSocketMessage encodeError(WebSocketSession session, String id, Throwable ex) {
GraphQLError error = GraphqlErrorBuilder.newError().message(ex.getMessage()).build();
return encode(session, GraphQlMessage.error(id, error));
return encode(session, GraphQlWebSocketMessage.error(id, error));
}
public WebSocketMessage encodeComplete(WebSocketSession session, String id) {
return encode(session, GraphQlMessage.complete(id));
return encode(session, GraphQlWebSocketMessage.complete(id));
}

View File

@@ -36,7 +36,7 @@ import org.springframework.graphql.server.WebGraphQlHandler;
import org.springframework.graphql.server.WebGraphQlRequest;
import org.springframework.graphql.server.WebGraphQlResponse;
import org.springframework.graphql.server.WebSocketGraphQlInterceptor;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.http.codec.CodecConfigurer;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
@@ -128,7 +128,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
.subscribe();
return session.send(session.receive().flatMap(webSocketMessage -> {
GraphQlMessage message = this.codecDelegate.decode(webSocketMessage);
GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage);
String id = message.getId();
Map<String, Object> payload = message.getPayload();
switch (message.resolvedType()) {
@@ -148,7 +148,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
.flatMapMany(response -> handleResponse(session, id, subscriptions, response))
.doOnTerminate(() -> subscriptions.remove(id));
case PING:
return Flux.just(this.codecDelegate.encode(session, GraphQlMessage.pong(null)));
return Flux.just(this.codecDelegate.encode(session, GraphQlWebSocketMessage.pong(null)));
case COMPLETE:
if (id != null) {
Subscription subscription = subscriptions.remove(id);

View File

@@ -47,7 +47,7 @@ import org.springframework.graphql.server.WebGraphQlHandler;
import org.springframework.graphql.server.WebGraphQlRequest;
import org.springframework.graphql.server.WebGraphQlResponse;
import org.springframework.graphql.server.WebSocketGraphQlInterceptor;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.http.HttpHeaders;
import org.springframework.http.HttpInputMessage;
import org.springframework.http.HttpOutputMessage;
@@ -139,7 +139,7 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
@Override
protected void handleTextMessage(WebSocketSession session, TextMessage webSocketMessage) throws Exception {
GraphQlMessage message = decode(webSocketMessage);
GraphQlWebSocketMessage message = decode(webSocketMessage);
String id = message.getId();
Map<String, Object> payload = message.getPayload();
SessionState sessionState = getSessionInfo(session);
@@ -166,7 +166,7 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
.subscribe(new SendMessageSubscriber(id, session, sessionState));
return;
case PING:
session.sendMessage(encode(GraphQlMessage.pong(null)));
session.sendMessage(encode(GraphQlWebSocketMessage.pong(null)));
return;
case COMPLETE:
if (id != null) {
@@ -187,7 +187,7 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
.defaultIfEmpty(Collections.emptyMap())
.publishOn(sessionState.getScheduler()) // Serial blocking send via single thread
.doOnNext(ackPayload -> {
TextMessage outputMessage = encode(GraphQlMessage.connectionAck(ackPayload));
TextMessage outputMessage = encode(GraphQlWebSocketMessage.connectionAck(ackPayload));
try {
session.sendMessage(outputMessage);
}
@@ -207,9 +207,9 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
}
@SuppressWarnings("unchecked")
private GraphQlMessage decode(TextMessage message) throws IOException {
return ((GenericHttpMessageConverter<GraphQlMessage>) this.converter)
.read(GraphQlMessage.class, null, new HttpInputMessageAdapter(message));
private GraphQlWebSocketMessage decode(TextMessage message) throws IOException {
return ((GenericHttpMessageConverter<GraphQlWebSocketMessage>) this.converter)
.read(GraphQlWebSocketMessage.class, null, new HttpInputMessageAdapter(message));
}
private SessionState getSessionInfo(WebSocketSession session) {
@@ -243,8 +243,8 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
}
return responseFlux
.map(responseMap -> encode(GraphQlMessage.next(id, responseMap)))
.concatWith(Mono.fromCallable(() -> encode(GraphQlMessage.complete(id))))
.map(responseMap -> encode(GraphQlWebSocketMessage.next(id, responseMap)))
.concatWith(Mono.fromCallable(() -> encode(GraphQlWebSocketMessage.complete(id))))
.onErrorResume((ex) -> {
if (ex instanceof SubscriptionExistsException) {
CloseStatus status = new CloseStatus(4409, "Subscriber for " + id + " already exists");
@@ -253,12 +253,12 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
}
String message = ex.getMessage();
GraphQLError error = GraphqlErrorBuilder.newError().message(message).build();
return Mono.just(encode(GraphQlMessage.error(id, error)));
return Mono.just(encode(GraphQlWebSocketMessage.error(id, error)));
});
}
@SuppressWarnings("unchecked")
private <T> TextMessage encode(GraphQlMessage message) {
private <T> TextMessage encode(GraphQlWebSocketMessage message) {
try {
HttpOutputMessageAdapter outputMessage = new HttpOutputMessageAdapter();
((HttpMessageConverter<T>) this.converter).write((T) message, null, outputMessage);

View File

@@ -30,7 +30,7 @@ import reactor.core.publisher.Mono;
import org.springframework.graphql.GraphQlRequest;
import org.springframework.graphql.GraphQlResponse;
import org.springframework.graphql.support.DefaultGraphQlRequest;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.lang.Nullable;
import org.springframework.web.reactive.socket.WebSocketHandler;
import org.springframework.web.reactive.socket.WebSocketSession;
@@ -84,15 +84,15 @@ public final class MockGraphQlWebSocketServer implements WebSocketHandler {
}
@SuppressWarnings("SuspiciousMethodCalls")
private Publisher<GraphQlMessage> handleMessage(GraphQlMessage message) {
private Publisher<GraphQlWebSocketMessage> handleMessage(GraphQlWebSocketMessage message) {
switch (message.resolvedType()) {
case CONNECTION_INIT:
if (this.connectionInitHandler == null) {
return Flux.just(GraphQlMessage.connectionAck(null));
return Flux.just(GraphQlWebSocketMessage.connectionAck(null));
}
else {
Map<String, Object> payload = message.getPayload();
return this.connectionInitHandler.apply(payload).map(GraphQlMessage::connectionAck);
return this.connectionInitHandler.apply(payload).map(GraphQlWebSocketMessage::connectionAck);
}
case SUBSCRIBE:
String id = message.getId();
@@ -101,11 +101,11 @@ public final class MockGraphQlWebSocketServer implements WebSocketHandler {
return Flux.error(new IllegalStateException("Unexpected request: " + message));
}
return request.getResponseFlux()
.map(response -> GraphQlMessage.next(id, response.toMap()))
.map(response -> GraphQlWebSocketMessage.next(id, response.toMap()))
.concatWithValues(
request.getError() != null ?
GraphQlMessage.error(id, request.getError()) :
GraphQlMessage.complete(id));
GraphQlWebSocketMessage.error(id, request.getError()) :
GraphQlWebSocketMessage.complete(id));
case COMPLETE:
return Flux.empty();
default:

View File

@@ -37,8 +37,8 @@ import org.springframework.graphql.ResponseError;
import org.springframework.graphql.support.DefaultGraphQlRequest;
import org.springframework.graphql.server.TestWebSocketClient;
import org.springframework.graphql.server.TestWebSocketConnection;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlMessageType;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessageType;
import org.springframework.http.HttpHeaders;
import org.springframework.http.codec.ClientCodecConfigurer;
import org.springframework.web.reactive.socket.CloseStatus;
@@ -86,8 +86,8 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(
GraphQlMessage.connectionInit(null),
GraphQlMessage.subscribe("1", request));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.subscribe("1", request));
}
@Test
@@ -99,8 +99,8 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(
GraphQlMessage.connectionInit(null),
GraphQlMessage.subscribe("1", request));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.subscribe("1", request));
}
@Test
@@ -117,8 +117,8 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(
GraphQlMessage.connectionInit(null),
GraphQlMessage.subscribe("1", request));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.subscribe("1", request));
}
@Test
@@ -135,8 +135,8 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(
GraphQlMessage.connectionInit(null),
GraphQlMessage.subscribe("1", request));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.subscribe("1", request));
}
@Test
@@ -149,8 +149,8 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(
GraphQlMessage.connectionInit(null),
GraphQlMessage.subscribe("1", request));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.subscribe("1", request));
}
@Test
@@ -165,9 +165,9 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(
GraphQlMessage.connectionInit(null),
GraphQlMessage.subscribe("1", request),
GraphQlMessage.complete("1"));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.subscribe("1", request),
GraphQlWebSocketMessage.complete("1"));
}
@Test
@@ -182,9 +182,9 @@ public class MockWebSocketGraphQlTransportTests {
.verify(TIMEOUT);
assertActualClientMessages(client.getConnection(0),
GraphQlMessage.connectionInit(null),
GraphQlMessage.pong(null),
GraphQlMessage.subscribe("1", new DefaultGraphQlRequest("{Query1}")));
GraphQlWebSocketMessage.connectionInit(null),
GraphQlWebSocketMessage.pong(null),
GraphQlWebSocketMessage.subscribe("1", new DefaultGraphQlRequest("{Query1}")));
}
@Test
@@ -218,7 +218,7 @@ public class MockWebSocketGraphQlTransportTests {
assertThat(client.getConnection(0).isOpen()).isTrue();
assertThat(connectionAckRef.get()).isEqualTo(Collections.singletonMap("key", "valueInitAck"));
assertActualClientMessages(client.getConnection(0), GraphQlMessage.connectionInit(initPayload));
assertActualClientMessages(client.getConnection(0), GraphQlWebSocketMessage.connectionInit(initPayload));
}
@Test
@@ -329,14 +329,14 @@ public class MockWebSocketGraphQlTransportTests {
new WebSocketGraphQlClientInterceptor() {});
}
private void assertActualClientMessages(GraphQlMessage... expectedMessages) {
private void assertActualClientMessages(GraphQlWebSocketMessage... expectedMessages) {
assertActualClientMessages(this.webSocketClient.getConnection(0), expectedMessages);
}
private void assertActualClientMessages(
TestWebSocketConnection connection, GraphQlMessage... expectedMessages) {
TestWebSocketConnection connection, GraphQlWebSocketMessage... expectedMessages) {
List<GraphQlMessage> actualMessages = connection.getClientMessages().stream()
List<GraphQlWebSocketMessage> actualMessages = connection.getClientMessages().stream()
.map(CODEC_DELEGATE::decode)
.collect(Collectors.toList());
@@ -361,12 +361,14 @@ public class MockWebSocketGraphQlTransportTests {
public Mono<Void> handle(WebSocketSession session) {
return session.send(session.receive()
.flatMap(webSocketMessage -> {
GraphQlMessage message = this.codecDelegate.decode(webSocketMessage);
GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage);
switch (message.resolvedType()) {
case CONNECTION_INIT:
return Flux.just(GraphQlMessage.connectionAck(null), GraphQlMessage.ping(null));
return Flux.just(
GraphQlWebSocketMessage.connectionAck(null),
GraphQlWebSocketMessage.ping(null));
case SUBSCRIBE:
return Flux.just(GraphQlMessage.next("1", this.response.toMap()));
return Flux.just(GraphQlWebSocketMessage.next("1", this.response.toMap()));
case PONG:
return Flux.empty();
default:
@@ -392,12 +394,13 @@ public class MockWebSocketGraphQlTransportTests {
public Mono<Void> handle(WebSocketSession session) {
return session.send(session.receive().flatMap(webSocketMessage -> {
GraphQlMessage inputMessage = this.codecDelegate.decode(webSocketMessage);
GraphQlWebSocketMessage inputMessage = this.codecDelegate.decode(webSocketMessage);
String id = inputMessage.getId();
GraphQlMessage outputMessage = (inputMessage.resolvedType() == GraphQlMessageType.CONNECTION_INIT ?
GraphQlMessage.connectionAck(null) :
GraphQlMessage.subscribe(id, new DefaultGraphQlRequest("")));
GraphQlWebSocketMessage outputMessage =
(inputMessage.resolvedType() == GraphQlWebSocketMessageType.CONNECTION_INIT ?
GraphQlWebSocketMessage.connectionAck(null) :
GraphQlWebSocketMessage.subscribe(id, new DefaultGraphQlRequest("")));
return Flux.just(this.codecDelegate.encode(session, outputMessage));
}));

View File

@@ -42,8 +42,8 @@ import org.springframework.graphql.server.WebGraphQlHandler;
import org.springframework.graphql.server.WebGraphQlInterceptor;
import org.springframework.graphql.server.WebSocketHandlerTestSupport;
import org.springframework.graphql.server.WebSocketGraphQlInterceptor;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlMessageType;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessageType;
import org.springframework.http.codec.ServerCodecConfigurer;
import org.springframework.http.codec.json.Jackson2JsonDecoder;
import org.springframework.web.reactive.socket.CloseStatus;
@@ -69,17 +69,17 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage(BOOK_QUERY)));
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.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, GraphQlMessageType.COMPLETE))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.COMPLETE))
.expectComplete()
.verify(TIMEOUT);
}
@@ -91,9 +91,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage(BOOK_SUBSCRIPTION)));
BiConsumer<WebSocketMessage, String> bookPayloadAssertion = (message, bookId) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.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 +101,10 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
};
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> bookPayloadAssertion.accept(message, "1"))
.consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5"))
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.COMPLETE))
.expectComplete()
.verify(TIMEOUT);
}
@@ -116,7 +116,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\", \"payload\":" + BOOK_QUERY_PAYLOAD + "}")));
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -135,7 +135,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
TestWebSocketSession session = handle(input);
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -160,8 +160,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.CONNECTION_ACK);
assertThat(actual.<Map<String, Object>>getPayload()).containsEntry("key", "A acknowledged");
})
.expectComplete()
@@ -175,8 +175,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage("{\"type\":\"ping\"}")));
StepVerifier.create(session.getOutput())
.consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.PONG))
.consumeNextWith(message -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith(message -> assertMessageType(message, GraphQlWebSocketMessageType.PONG))
.expectComplete()
.verify(TIMEOUT);
}
@@ -239,7 +239,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage("{\"type\":\"connection_init\"}")));
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -273,7 +273,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
TestWebSocketSession session = handle(messageFlux, new ConsumeOneAndNeverCompleteInterceptor());
// Collect messages until session closed
List<GraphQlMessage> messages = new ArrayList<>();
List<GraphQlWebSocketMessage> messages = new ArrayList<>();
session.getOutput().subscribe((message) -> messages.add(decode(message)));
StepVerifier.create(session.closeStatus())
@@ -282,8 +282,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
.verify(TIMEOUT);
assertThat(messages.size()).isEqualTo(2);
assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK);
assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlWebSocketMessageType.CONNECTION_ACK);
assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlWebSocketMessageType.NEXT);
}
@Test
@@ -297,12 +297,12 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
String completeMessage = "{\"id\":\"" + SUBSCRIPTION_ID + "\",\"type\":\"complete\"}";
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.NEXT))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.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, GraphQlMessageType.NEXT))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.NEXT))
.then(() -> input.tryEmitNext(toWebSocketMessage(completeMessage)))
.verifyTimeout(Duration.ofMillis(500));
}
@@ -336,19 +336,19 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
handler.handle(session).block(TIMEOUT);
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.NEXT);
assertThat(actual.<Map<String, Object>>getPayload())
.extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class)))
.containsEntry("greeting", "a");
})
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.ERROR);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.ERROR);
assertThat(actual.<List<Map<String, Object>>>getPayload())
.asList().hasSize(1)
.allSatisfy(theError -> assertThat(theError)
@@ -380,15 +380,15 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
}
@SuppressWarnings("ConstantConditions")
private GraphQlMessage decode(WebSocketMessage message) {
return (GraphQlMessage) decoder.decode(DataBufferUtils.retain(message.getPayload()),
ResolvableType.forClass(GraphQlMessage.class), null, Collections.emptyMap());
private GraphQlWebSocketMessage decode(WebSocketMessage message) {
return (GraphQlWebSocketMessage) decoder.decode(DataBufferUtils.retain(message.getPayload()),
ResolvableType.forClass(GraphQlWebSocketMessage.class), null, Collections.emptyMap());
}
private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlMessageType messageType) {
GraphQlMessage message = decode(webSocketMessage);
private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlWebSocketMessageType messageType) {
GraphQlWebSocketMessage message = decode(webSocketMessage);
assertThat(message.resolvedType()).isEqualTo(messageType);
if (messageType != GraphQlMessageType.CONNECTION_ACK && messageType != GraphQlMessageType.PONG) {
if (messageType != GraphQlWebSocketMessageType.CONNECTION_ACK && messageType != GraphQlWebSocketMessageType.PONG) {
assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID);
}
}

View File

@@ -38,8 +38,8 @@ import org.springframework.graphql.GraphQlSetup;
import org.springframework.graphql.server.WebGraphQlHandler;
import org.springframework.graphql.server.WebGraphQlInterceptor;
import org.springframework.graphql.server.WebSocketGraphQlInterceptor;
import org.springframework.graphql.server.support.GraphQlMessage;
import org.springframework.graphql.server.support.GraphQlMessageType;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.graphql.server.support.GraphQlWebSocketMessageType;
import org.springframework.graphql.server.ConsumeOneAndNeverCompleteInterceptor;
import org.springframework.graphql.server.WebSocketHandlerTestSupport;
import org.springframework.http.HttpHeaders;
@@ -76,17 +76,17 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
new TextMessage(BOOK_QUERY));
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.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, GraphQlMessageType.COMPLETE))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.COMPLETE))
.then(this.session::close) // Complete output Flux
.expectComplete()
.verify(TIMEOUT);
@@ -97,9 +97,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
handle(this.handler, new TextMessage("{\"type\":\"connection_init\"}"), new TextMessage(BOOK_SUBSCRIPTION));
BiConsumer<WebSocketMessage<?>, String> bookPayloadAssertion = (message, bookId) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.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)))
@@ -107,10 +107,10 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
};
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> bookPayloadAssertion.accept(message, "1"))
.consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5"))
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.COMPLETE))
.then(this.session::close)// Complete output Flux
.expectComplete()
.verify(TIMEOUT);
@@ -124,7 +124,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
// No message type
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -138,7 +138,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
new TextMessage("{\"type\":\"subscribe\", \"payload\":{}}")); // No message id
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -162,8 +162,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
StepVerifier.create(session.getOutput())
.consumeNextWith((webSocketMessage) -> {
GraphQlMessage message = decode(webSocketMessage);
assertThat(message.resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK);
GraphQlWebSocketMessage message = decode(webSocketMessage);
assertThat(message.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.CONNECTION_ACK);
assertThat(message.<Map<String, Object>>getPayload()).containsEntry("key", "A acknowledged");
})
.then(this.session::close) // Complete output Flux
@@ -179,8 +179,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
new TextMessage("{\"type\":\"ping\"}"));
StepVerifier.create(session.getOutput())
.consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.PONG))
.consumeNextWith(message -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith(message -> assertMessageType(message, GraphQlWebSocketMessageType.PONG))
.then(this.session::close) // Complete output Flux
.expectComplete()
.verify(TIMEOUT);
@@ -250,7 +250,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
new TextMessage("{\"type\":\"connection_init\"}"));
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -276,7 +276,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
new TextMessage(BOOK_SUBSCRIPTION));
// Collect messages until session closed
List<GraphQlMessage> messages = new ArrayList<>();
List<GraphQlWebSocketMessage> messages = new ArrayList<>();
this.session.getOutput().subscribe((message) -> messages.add(decode(message)));
StepVerifier.create(this.session.closeStatus())
@@ -285,8 +285,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
.verify(TIMEOUT);
assertThat(messages.size()).isEqualTo(2);
assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK);
assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlWebSocketMessageType.CONNECTION_ACK);
assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlWebSocketMessageType.NEXT);
}
@Test
@@ -308,12 +308,12 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
};
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.NEXT))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.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, GraphQlMessageType.NEXT))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.NEXT))
.then(() -> messageSender.accept(completeMessage))
.verifyTimeout(Duration.ofMillis(500));
}
@@ -343,19 +343,19 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
new TextMessage(GREETING_QUERY));
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.NEXT);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.NEXT);
assertThat(actual.<Map<String, Object>>getPayload())
.extractingByKey("data", as(InstanceOfAssertFactories.map(String.class, Object.class)))
.containsEntry("greeting", "a");
})
.consumeNextWith((message) -> {
GraphQlMessage actual = decode(message);
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.ERROR);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.ERROR);
assertThat(actual.<List<Map<String, Object>>>getPayload())
.asList().hasSize(1)
.allSatisfy(theError -> assertThat(theError)
@@ -388,21 +388,21 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
}
@SuppressWarnings("unchecked")
private GraphQlMessage decode(WebSocketMessage<?> message) {
private GraphQlWebSocketMessage decode(WebSocketMessage<?> message) {
try {
HttpInputMessageAdapter inputMessage = new HttpInputMessageAdapter((TextMessage) message);
return ((GenericHttpMessageConverter<GraphQlMessage>) converter)
.read(GraphQlMessage.class, null, inputMessage);
return ((GenericHttpMessageConverter<GraphQlWebSocketMessage>) converter)
.read(GraphQlWebSocketMessage.class, null, inputMessage);
}
catch (IOException ex) {
throw new IllegalStateException(ex);
}
}
private void assertMessageType(WebSocketMessage<?> webSocketMessage, GraphQlMessageType messageType) {
GraphQlMessage message = decode(webSocketMessage);
private void assertMessageType(WebSocketMessage<?> webSocketMessage, GraphQlWebSocketMessageType messageType) {
GraphQlWebSocketMessage message = decode(webSocketMessage);
assertThat(message.resolvedType()).isEqualTo(messageType);
if (messageType != GraphQlMessageType.CONNECTION_ACK && messageType != GraphQlMessageType.PONG) {
if (messageType != GraphQlWebSocketMessageType.CONNECTION_ACK && messageType != GraphQlWebSocketMessageType.PONG) {
assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID);
}
}