diff --git a/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java b/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java index 8baf517e..89251664 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/client/WebSocketGraphQlTransport.java @@ -265,6 +265,9 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { case NEXT: graphQlSession.handleNext(message); break; + case PING: + graphQlSession.sendPong(null); + break; case ERROR: graphQlSession.handleError(message); break; @@ -416,6 +419,11 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { } } + public void sendPong(@Nullable Map payload) { + GraphQlMessage message = GraphQlMessage.pong(payload); + trySend(message); + } + // TODO: queue to serialize sending? private void trySend(GraphQlMessage message) { diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java index 12909254..74cc28cb 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandler.java @@ -147,6 +147,8 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { return this.graphQlHandler.handleRequest(input) .flatMapMany((output) -> handleWebOutput(session, id, subscriptions, output)) .doOnTerminate(() -> subscriptions.remove(id)); + case PING: + return Flux.just(this.codecDelegate.encode(session, GraphQlMessage.pong(null))); case COMPLETE: if (id != null) { Subscription subscription = subscriptions.remove(id); diff --git a/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java index b9a9cc11..50fb2079 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandler.java @@ -166,6 +166,9 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread .subscribe(new SendMessageSubscriber(id, session, sessionState)); return; + case PING: + session.sendMessage(encode(GraphQlMessage.pong(null))); + return; case COMPLETE: if (id != null) { Subscription subscription = sessionState.getSubscriptions().remove(id); diff --git a/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java b/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java index d458d5e5..e25b6096 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/client/MockWebSocketGraphQlTransportTests.java @@ -168,6 +168,23 @@ public class MockWebSocketGraphQlTransportTests { GraphQlMessage.complete("1")); } + @Test + void pingHandling() { + + TestWebSocketClient client = new TestWebSocketClient(new PingResponseHandler(this.result1)); + WebSocketGraphQlTransport transport = createTransport(client); + + StepVerifier.create(transport.execute(new GraphQlRequest("{Query1}"))) + .expectNext(this.result1) + .expectComplete() + .verify(TIMEOUT); + + assertActualClientMessages(client.getConnection(0), + GraphQlMessage.connectionInit(null), + GraphQlMessage.pong(null), + GraphQlMessage.subscribe("1", new GraphQlRequest("{Query1}"))); + } + @Test void start() { MockGraphQlWebSocketServer handler = new MockGraphQlWebSocketServer(); @@ -310,6 +327,42 @@ public class MockWebSocketGraphQlTransportTests { } + /** + * Server handler that inserts a "ping" after the "connection_ack". + */ + private static class PingResponseHandler implements WebSocketHandler { + + private final ExecutionResult result; + + private final CodecDelegate codecDelegate = new CodecDelegate(); + + private PingResponseHandler(ExecutionResult result) { + this.result = result; + } + + @Override + public Mono handle(WebSocketSession session) { + return session.send(session.receive() + .flatMap(webSocketMessage -> { + GraphQlMessage message = this.codecDelegate.decode(webSocketMessage); + switch (message.resolvedType()) { + case CONNECTION_INIT: + return Flux.just(GraphQlMessage.connectionAck(null), GraphQlMessage.ping(null)); + case SUBSCRIBE: + return Flux.just(GraphQlMessage.next("1", this.result)); + case PONG: + return Flux.empty(); + default: + return Flux.error(new IllegalStateException("Unexpected message: " + message)); + } + }) + .map(graphQlMessage -> this.codecDelegate.encode(session, graphQlMessage)) + ); + } + + } + + /** * Server handler that returns an unexpected (client) message. */ diff --git a/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java index 1a36b3f9..d2da508b 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/web/webflux/GraphQlWebSocketHandlerTests.java @@ -59,6 +59,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { private static final Jackson2JsonDecoder decoder = new Jackson2JsonDecoder(); + private static final Duration TIMEOUT = Duration.ofSeconds(5); + + @Test void query() { TestWebSocketSession session = handle(Flux.just( @@ -77,7 +80,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .containsEntry("name", "Nineteen Eighty-Four"); }) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -101,7 +105,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "1")) .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5")) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -112,11 +117,13 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4400, "Invalid message")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -129,11 +136,13 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4400, "Invalid message")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -155,7 +164,21 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { assertThat(actual.resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); assertThat(actual.>getPayload()).containsEntry("key", "A acknowledged"); }) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); + } + + @Test + void pingHandling() { + TestWebSocketSession session = handle(Flux.just( + toWebSocketMessage("{\"type\":\"connection_init\",\"payload\":{\"key\":\"A\"}}"), + toWebSocketMessage("{\"type\":\"ping\"}"))); + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) + .consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.PONG)) + .expectComplete() + .verify(TIMEOUT); } @Test @@ -197,14 +220,15 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()).verifyComplete(); StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4401, "Unauthorized")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test void unauthorizedWithoutConnectionInit() { TestWebSocketSession session = handle(Flux.just(toWebSocketMessage(BOOK_SUBSCRIPTION))); - StepVerifier.create(session.getOutput()).verifyComplete(); + StepVerifier.create(session.getOutput()).expectComplete().verify(TIMEOUT); StepVerifier.create(session.closeStatus()).expectNext(new CloseStatus(4401, "Unauthorized")).verifyComplete(); } @@ -216,11 +240,13 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4429, "Too many initialisation requests")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -229,11 +255,12 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { initHandler(), ServerCodecConfigurer.create(), Duration.ofMillis(50)); TestWebSocketSession session = new TestWebSocketSession(Flux.empty()); - handler.handle(session).block(); + handler.handle(session).block(TIMEOUT); StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4408, "Connection initialisation timeout")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -251,7 +278,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4409, "Subscriber for " + SUBSCRIPTION_ID + " already exists")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); assertThat(messages.size()).isEqualTo(2); assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); @@ -305,7 +333,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { TestWebSocketSession session = new TestWebSocketSession(Flux.just( toWebSocketMessage("{\"type\":\"connection_init\"}"), toWebSocketMessage(GREETING_QUERY))); - handler.handle(session).block(); + handler.handle(session).block(TIMEOUT); StepVerifier.create(session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) @@ -331,7 +359,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .extractingByKey("extensions", as(InstanceOfAssertFactories.map(String.class, Object.class))) .containsEntry("classification", "DataFetchingException")); }) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } private TestWebSocketSession handle(Flux input, WebInterceptor... interceptors) { @@ -341,7 +370,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { Duration.ofSeconds(60)); TestWebSocketSession session = new TestWebSocketSession(input); - handler.handle(session).block(); + handler.handle(session).block(TIMEOUT); return session; } @@ -359,7 +388,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlMessageType messageType) { GraphQlMessage message = decode(webSocketMessage); assertThat(message.resolvedType()).isEqualTo(messageType); - if (messageType != GraphQlMessageType.CONNECTION_ACK) { + if (messageType != GraphQlMessageType.CONNECTION_ACK && messageType != GraphQlMessageType.PONG) { assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID); } } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java index 1c9afac4..f5d714d7 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/web/webmvc/GraphQlWebSocketHandlerTests.java @@ -61,10 +61,14 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { private static final HttpMessageConverter converter = new MappingJackson2HttpMessageConverter(); + private static final Duration TIMEOUT = Duration.ofSeconds(5); + + private final TestWebSocketSession session = new TestWebSocketSession(); private final GraphQlWebSocketHandler handler = initWebSocketHandler(); + @Test void query() throws Exception { handle(this.handler, @@ -84,7 +88,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { }) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) .then(this.session::close) // Complete output Flux - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -107,7 +112,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .consumeNextWith((message) -> bookPayloadAssertion.accept(message, "5")) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.COMPLETE)) .then(this.session::close)// Complete output Flux - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -119,7 +125,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(this.session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); assertThat(this.session.getCloseStatus()).isEqualTo(new CloseStatus(4400, "Invalid message")); } @@ -132,7 +139,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(this.session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); assertThat(this.session.getCloseStatus()).isEqualTo(new CloseStatus(4400, "Invalid message")); } @@ -159,7 +167,23 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { assertThat(message.>getPayload()).containsEntry("key", "A acknowledged"); }) .then(this.session::close) // Complete output Flux - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); + } + + @Test + void pingHandling() throws Exception { + + handle(initWebSocketHandler(), + new TextMessage("{\"type\":\"connection_init\",\"payload\":{\"key\":\"A\"}}"), + new TextMessage("{\"type\":\"ping\"}")); + + StepVerifier.create(session.getOutput()) + .consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) + .consumeNextWith(message -> assertMessageType(message, GraphQlMessageType.PONG)) + .then(this.session::close) // Complete output Flux + .expectComplete() + .verify(TIMEOUT); } @Test @@ -185,7 +209,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.getOutput()) .expectNextCount(1) .then(this.session::close) // Complete output Flux - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); handler.afterConnectionClosed(this.session, closeStatus); assertThat(called).isTrue(); @@ -206,7 +231,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(session.closeStatus()) .expectNext(new CloseStatus(4401, "Unauthorized")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -225,7 +251,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(this.session.getOutput()) .consumeNextWith((message) -> assertMessageType(message, GraphQlMessageType.CONNECTION_ACK)) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); assertThat(this.session.getCloseStatus()).isEqualTo(new CloseStatus(4429, "Too many initialisation requests")); } @@ -237,7 +264,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(this.session.closeStatus()) .expectNext(new CloseStatus(4408, "Connection initialisation timeout")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } @Test @@ -253,7 +281,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { StepVerifier.create(this.session.closeStatus()) .expectNext(new CloseStatus(4409, "Subscriber for " + SUBSCRIPTION_ID + " already exists")) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); assertThat(messages.size()).isEqualTo(2); assertThat(messages.get(0).resolvedType()).isEqualTo(GraphQlMessageType.CONNECTION_ACK); @@ -338,7 +367,8 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { .containsEntry("classification", "DataFetchingException")); }) .then(this.session::close) - .verifyComplete(); + .expectComplete() + .verify(TIMEOUT); } private void handle(GraphQlWebSocketHandler handler, TextMessage... textMessages) throws Exception { @@ -372,7 +402,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { private void assertMessageType(WebSocketMessage webSocketMessage, GraphQlMessageType messageType) { GraphQlMessage message = decode(webSocketMessage); assertThat(message.resolvedType()).isEqualTo(messageType); - if (messageType != GraphQlMessageType.CONNECTION_ACK) { + if (messageType != GraphQlMessageType.CONNECTION_ACK && messageType != GraphQlMessageType.PONG) { assertThat(message.getId()).isEqualTo(SUBSCRIPTION_ID); } }