WebSocket handlers support keepalive PING messages

Closes gh-534
This commit is contained in:
rstoyanchev
2024-04-11 13:54:00 +01:00
parent c8573cb720
commit 80ef9604e2
5 changed files with 179 additions and 29 deletions

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2023 the original author or authors.
* Copyright 2002-2024 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.
@@ -44,6 +44,7 @@ import org.springframework.graphql.server.WebSocketSessionInfo;
import org.springframework.graphql.server.support.GraphQlWebSocketMessage;
import org.springframework.http.HttpHeaders;
import org.springframework.http.codec.CodecConfigurer;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.web.reactive.socket.CloseStatus;
@@ -72,10 +73,13 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
private final WebSocketGraphQlInterceptor webSocketInterceptor;
private final WebSocketCodecDelegate webSocketCodecDelegate;
private final WebSocketCodecDelegate codecDelegate;
private final Duration initTimeoutDuration;
@Nullable
private final Duration keepAliveDuration;
/**
* Create a new instance.
@@ -87,12 +91,30 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
public GraphQlWebSocketHandler(
WebGraphQlHandler graphQlHandler, CodecConfigurer codecConfigurer, Duration connectionInitTimeout) {
this(graphQlHandler, codecConfigurer, connectionInitTimeout, null);
}
/**
* Create a new instance.
* @param graphQlHandler common handler for GraphQL over WebSocket requests
* @param codecConfigurer codec configurer for JSON encoding and decoding
* @param connectionInitTimeout how long to wait after the establishment of
* the WebSocket for the {@code "connection_ini"} message from the client.
* @param keepAliveDuration how frequently to send ping messages; if not
* set then ping messages are not sent.
* @since 1.3
*/
public GraphQlWebSocketHandler(
WebGraphQlHandler graphQlHandler, CodecConfigurer codecConfigurer,
Duration connectionInitTimeout, @Nullable Duration keepAliveDuration) {
Assert.notNull(graphQlHandler, "WebGraphQlHandler is required");
this.graphQlHandler = graphQlHandler;
this.webSocketInterceptor = this.graphQlHandler.getWebSocketInterceptor();
this.webSocketCodecDelegate = new WebSocketCodecDelegate(codecConfigurer);
this.codecDelegate = new WebSocketCodecDelegate(codecConfigurer);
this.initTimeoutDuration = connectionInitTimeout;
this.keepAliveDuration = keepAliveDuration;
}
@@ -137,7 +159,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
.subscribe();
return session.send(session.receive().flatMap((webSocketMessage) -> {
GraphQlWebSocketMessage message = this.webSocketCodecDelegate.decode(webSocketMessage);
GraphQlWebSocketMessage message = this.codecDelegate.decode(webSocketMessage);
String id = message.getId();
Map<String, Object> payload = message.getPayload();
switch (message.resolvedType()) {
@@ -159,7 +181,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
.doOnTerminate(() -> subscriptions.remove(id));
}
case PING -> {
return Flux.just(this.webSocketCodecDelegate.encode(session, GraphQlWebSocketMessage.pong(null)));
return Flux.just(this.codecDelegate.encode(session, GraphQlWebSocketMessage.pong(null)));
}
case COMPLETE -> {
if (id != null) {
@@ -176,11 +198,16 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
if (!connectionInitPayloadRef.compareAndSet(null, payload)) {
return GraphQlStatus.close(session, GraphQlStatus.TOO_MANY_INIT_REQUESTS_STATUS);
}
return this.webSocketInterceptor.handleConnectionInitialization(sessionInfo, payload)
Flux<WebSocketMessage> flux = this.webSocketInterceptor.handleConnectionInitialization(sessionInfo, payload)
.defaultIfEmpty(Collections.emptyMap())
.map((ackPayload) -> this.webSocketCodecDelegate.encodeConnectionAck(session, ackPayload))
.flux()
.onErrorResume((ex) -> GraphQlStatus.close(session, GraphQlStatus.UNAUTHORIZED_STATUS));
.map((ackPayload) -> this.codecDelegate.encodeConnectionAck(session, ackPayload))
.flux();
if (this.keepAliveDuration != null) {
flux = flux.mergeWith(Flux.interval(this.keepAliveDuration, this.keepAliveDuration)
.filter((aLong) -> !this.codecDelegate.checkMessagesEncodedAndClear())
.map((aLong) -> this.codecDelegate.encode(session, GraphQlWebSocketMessage.ping(null))));
}
return flux.onErrorResume((ex) -> GraphQlStatus.close(session, GraphQlStatus.UNAUTHORIZED_STATUS));
}
default -> {
return GraphQlStatus.close(session, GraphQlStatus.INVALID_MESSAGE_STATUS);
@@ -218,14 +245,14 @@ public class GraphQlWebSocketHandler implements WebSocketHandler {
}
return responseFlux
.map((responseMap) -> this.webSocketCodecDelegate.encodeNext(session, id, responseMap))
.concatWith(Mono.fromCallable(() -> this.webSocketCodecDelegate.encodeComplete(session, id)))
.map((responseMap) -> this.codecDelegate.encodeNext(session, id, responseMap))
.concatWith(Mono.fromCallable(() -> this.codecDelegate.encodeComplete(session, id)))
.onErrorResume((ex) -> {
if (ex instanceof SubscriptionExistsException) {
CloseStatus status = new CloseStatus(4409, "Subscriber for " + id + " already exists");
return GraphQlStatus.close(session, status);
}
return Mono.fromCallable(() -> this.webSocketCodecDelegate.encodeError(session, id, ex));
return Mono.fromCallable(() -> this.codecDelegate.encodeError(session, id, ex));
});
}

View File

@@ -54,6 +54,8 @@ final class WebSocketCodecDelegate {
private final Encoder<?> encoder;
private boolean messagesEncoded;
WebSocketCodecDelegate(CodecConfigurer codecConfigurer) {
Assert.notNull(codecConfigurer, "CodecConfigurer is required");
@@ -84,6 +86,8 @@ final class WebSocketCodecDelegate {
DataBuffer buffer = ((Encoder<T>) this.encoder).encodeValue(
(T) message, session.bufferFactory(), MESSAGE_TYPE, MimeTypeUtils.APPLICATION_JSON, null);
this.messagesEncoded = true;
return new WebSocketMessage(WebSocketMessage.Type.TEXT, buffer);
}
@@ -115,4 +119,10 @@ final class WebSocketCodecDelegate {
return encode(session, GraphQlWebSocketMessage.complete(id));
}
boolean checkMessagesEncodedAndClear() {
boolean result = this.messagesEncoded;
this.messagesEncoded = false;
return result;
}
}

View File

@@ -104,8 +104,12 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
private final HttpMessageConverter<?> converter;
@Nullable
private final Duration keepAliveDuration;
private final Map<String, SessionState> sessionInfoMap = new ConcurrentHashMap<>();
/**
* Create a new instance.
* @param graphQlHandler common handler for GraphQL over WebSocket requests
@@ -116,6 +120,23 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
public GraphQlWebSocketHandler(
WebGraphQlHandler graphQlHandler, HttpMessageConverter<?> converter, Duration connectionInitTimeout) {
this(graphQlHandler, converter, connectionInitTimeout, null);
}
/**
* Create a new instance.
* @param graphQlHandler common handler for GraphQL over WebSocket requests
* @param converter for JSON encoding and decoding
* @param connectionInitTimeout how long to wait after the establishment of
* the WebSocket for the {@code "connection_ini"} message from the client.
* @param keepAliveDuration how frequently to send ping messages; if not
* set then ping messages are not sent.
* @since 1.3
*/
public GraphQlWebSocketHandler(
WebGraphQlHandler graphQlHandler, HttpMessageConverter<?> converter,
Duration connectionInitTimeout, @Nullable Duration keepAliveDuration) {
Assert.notNull(graphQlHandler, "WebGraphQlHandler is required");
Assert.notNull(converter, "HttpMessageConverter for JSON is required");
@@ -124,8 +145,10 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
this.webSocketGraphQlInterceptor = this.graphQlHandler.getWebSocketInterceptor();
this.initTimeoutDuration = connectionInitTimeout;
this.converter = converter;
this.keepAliveDuration = keepAliveDuration;
}
@Override
public List<String> getSubProtocols() {
return SUB_PROTOCOL_LIST;
@@ -257,6 +280,21 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
return Mono.empty();
})
.block(Duration.ofSeconds(10));
if (this.keepAliveDuration != null) {
Flux.interval(this.keepAliveDuration, this.keepAliveDuration)
.filter((aLong) -> true)
.publishOn(state.getScheduler()) // Serial blocking send via single thread
.doOnNext((aLong) -> {
try {
session.sendMessage(encode(GraphQlWebSocketMessage.ping(null)));
}
catch (IOException ex) {
ExceptionWebSocketHandlerDecorator.tryCloseWithError(session, ex, logger);
}
})
.subscribe(state.getKeepAliveSubscriber());
}
}
default -> GraphQlStatus.closeSession(session, GraphQlStatus.INVALID_MESSAGE_STATUS);
}
@@ -444,9 +482,12 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
private final Scheduler scheduler;
SessionState(String graphQlSessionId, WebSocketSessionInfo sessionInfo) {
private final KeepAliveSubscriber keepAliveSubscriber;
SessionState(String graphQlSessionId, WebMvcSessionInfo sessionInfo) {
this.sessionInfo = sessionInfo;
this.scheduler = Schedulers.newSingle("GraphQL-WsSession-" + graphQlSessionId);
this.keepAliveSubscriber = new KeepAliveSubscriber(sessionInfo.getSession());
}
WebSocketSessionInfo getSessionInfo() {
@@ -462,12 +503,16 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
return this.connectionInitPayloadRef.compareAndSet(null, payload);
}
KeepAliveSubscriber getKeepAliveSubscriber() {
return this.keepAliveSubscriber;
}
Map<String, Subscription> getSubscriptions() {
return this.subscriptions;
}
void dispose() {
this.keepAliveSubscriber.cancel();
for (Map.Entry<String, Subscription> entry : this.subscriptions.entrySet()) {
try {
entry.getValue().cancel();
@@ -525,6 +570,10 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
public InetSocketAddress getRemoteAddress() {
return this.session.getRemoteAddress();
}
WebSocketSession getSession() {
return this.session;
}
}
@@ -567,9 +616,29 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub
public void hookOnComplete() {
this.sessionState.getSubscriptions().remove(this.subscriptionId);
}
}
private static class KeepAliveSubscriber extends BaseSubscriber<Long> {
private final WebSocketSession session;
KeepAliveSubscriber(WebSocketSession session) {
this.session = session;
}
@Override
protected void hookOnSubscribe(Subscription subscription) {
subscription.request(Integer.MAX_VALUE);
}
@Override
public void hookOnError(Throwable ex) {
ExceptionWebSocketHandlerDecorator.tryCloseWithError(this.session, ex, logger);
}
}
@SuppressWarnings("serial")
private static final class SubscriptionExistsException extends RuntimeException {
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2022 the original author or authors.
* Copyright 2002-2024 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.
@@ -22,6 +22,7 @@ import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.function.BiConsumer;
@@ -57,6 +58,9 @@ import org.springframework.web.reactive.socket.WebSocketMessage;
import static org.assertj.core.api.Assertions.as;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.graphql.server.support.GraphQlWebSocketMessageType.CONNECTION_ACK;
import static org.springframework.graphql.server.support.GraphQlWebSocketMessageType.PING;
import static org.springframework.graphql.server.support.GraphQlWebSocketMessageType.PONG;
/**
* Unit tests for {@link GraphQlWebSocketHandler}.
@@ -75,7 +79,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage(BOOK_QUERY)));
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.consumeNextWith((message) -> {
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
@@ -107,7 +111,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
};
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.consumeNextWith((message) -> bookPayloadAssertion.accept(message, "1"))
.consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5"))
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.COMPLETE))
@@ -115,6 +119,25 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
.verify(TIMEOUT);
}
@Test
void keepAlive() {
GraphQlWebSocketHandler handler = new GraphQlWebSocketHandler(
initHandler(), ServerCodecConfigurer.create(), Duration.ofSeconds(60), Duration.ofMillis(10));
TestWebSocketSession session =
new TestWebSocketSession(Flux.just(toWebSocketMessage("{\"type\":\"connection_init\"}")));
handler.handle(session).block(TIMEOUT);
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, PING))
.consumeNextWith((message) -> assertMessageType(message, PING))
.consumeNextWith((message) -> assertMessageType(message, PING))
.thenCancel()
.verify(TIMEOUT);
}
@Test
void unauthorizedWithoutMessageType() {
TestWebSocketSession session = handle(Flux.just(
@@ -122,7 +145,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage("{\"id\":\"" + SUBSCRIPTION_ID + "\", \"payload\":" + BOOK_QUERY_PAYLOAD + "}")));
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -141,7 +164,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
TestWebSocketSession session = handle(input);
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -167,7 +190,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> {
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.resolvedType()).isEqualTo(GraphQlWebSocketMessageType.CONNECTION_ACK);
assertThat(actual.resolvedType()).isEqualTo(CONNECTION_ACK);
assertThat(actual.<Map<String, Object>>getPayload()).containsEntry("key", "A acknowledged");
})
.expectComplete()
@@ -181,8 +204,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage("{\"type\":\"ping\"}")));
StepVerifier.create(session.getOutput())
.consumeNextWith(message -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith(message -> assertMessageType(message, GraphQlWebSocketMessageType.PONG))
.consumeNextWith(message -> assertMessageType(message, CONNECTION_ACK))
.consumeNextWith(message -> assertMessageType(message, PONG))
.expectComplete()
.verify(TIMEOUT);
}
@@ -245,7 +268,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
toWebSocketMessage("{\"type\":\"connection_init\"}")));
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.expectComplete()
.verify(TIMEOUT);
@@ -288,7 +311,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
.verify(TIMEOUT);
assertThat(messages.size()).isEqualTo(2);
assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlWebSocketMessageType.CONNECTION_ACK);
assertThat(messages.get(0).resolvedType()).isEqualTo(CONNECTION_ACK);
assertThat(messages.get(1).resolvedType()).isEqualTo(GraphQlWebSocketMessageType.NEXT);
}
@@ -303,7 +326,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
String completeMessage = "{\"id\":\"" + SUBSCRIPTION_ID + "\",\"type\":\"complete\"}";
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, 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")
@@ -335,7 +358,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
.handle(session).block(TIMEOUT);
StepVerifier.create(session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, CONNECTION_ACK))
.consumeNextWith((message) -> {
GraphQlWebSocketMessage actual = decode(message);
assertThat(actual.getId()).isEqualTo(SUBSCRIPTION_ID);
@@ -399,7 +422,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlWebSocketMessageType messageType) {
GraphQlWebSocketMessage message = decode(webSocketMessage);
assertThat(message.resolvedType()).isEqualTo(messageType);
if (messageType != GraphQlWebSocketMessageType.CONNECTION_ACK && messageType != GraphQlWebSocketMessageType.PONG) {
if (!Set.of(CONNECTION_ACK, PING, PONG).contains(messageType)) {
assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID);
}
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2022 the original author or authors.
* Copyright 2002-2024 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.
@@ -24,6 +24,7 @@ import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.function.BiConsumer;
import java.util.function.Consumer;
@@ -62,6 +63,9 @@ import org.springframework.web.socket.WebSocketMessage;
import static org.assertj.core.api.Assertions.as;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.graphql.server.support.GraphQlWebSocketMessageType.CONNECTION_ACK;
import static org.springframework.graphql.server.support.GraphQlWebSocketMessageType.PING;
import static org.springframework.graphql.server.support.GraphQlWebSocketMessageType.PONG;
/**
* Unit tests for {@link GraphQlWebSocketHandler}.
@@ -125,6 +129,23 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
.verify(TIMEOUT);
}
@Test
void keepAlive() throws Exception {
GraphQlWebSocketHandler webSocketHandler =
new GraphQlWebSocketHandler(initHandler(), converter, Duration.ofSeconds(60), Duration.ofMillis(10));
handle(webSocketHandler, new TextMessage("{\"type\":\"connection_init\"}"));
StepVerifier.create(this.session.getOutput())
.consumeNextWith((message) -> assertMessageType(message, GraphQlWebSocketMessageType.CONNECTION_ACK))
.consumeNextWith((message) -> assertMessageType(message, PING))
.consumeNextWith((message) -> assertMessageType(message, PING))
.consumeNextWith((message) -> assertMessageType(message, PING))
.then(this.session::close)// Complete output Flux
.thenCancel()
.verify(TIMEOUT);
}
@Test
void unauthorizedWithoutMessageType() throws Exception {
handle(this.handler,
@@ -464,7 +485,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport {
private void assertMessageType(WebSocketMessage<?> webSocketMessage, GraphQlWebSocketMessageType messageType) {
GraphQlWebSocketMessage message = decode(webSocketMessage);
assertThat(message.resolvedType()).isEqualTo(messageType);
if (messageType != GraphQlWebSocketMessageType.CONNECTION_ACK && messageType != GraphQlWebSocketMessageType.PONG) {
if (!Set.of(CONNECTION_ACK, PING, PONG).contains(messageType)) {
assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID);
}
}