diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/WebGraphQlInterceptor.java b/spring-graphql/src/main/java/org/springframework/graphql/server/WebGraphQlInterceptor.java index 31333aa8..eb6c43ec 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/WebGraphQlInterceptor.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/WebGraphQlInterceptor.java @@ -22,6 +22,7 @@ import reactor.core.publisher.Mono; import org.springframework.beans.factory.ObjectProvider; import org.springframework.graphql.ExecutionGraphQlService; +import org.springframework.util.Assert; /** @@ -44,7 +45,8 @@ public interface WebGraphQlInterceptor { /** * Intercept a request and delegate to the rest of the chain including other * interceptors and a {@link ExecutionGraphQlService}. - * @param request the request to execute + * @param request the request which may be a {@link WebSocketGraphQlRequest} + * when intercepting a GraphQL request over WebSocket * @param chain the rest of the chain to execute the request * @return a {@link Mono} with the response */ @@ -57,7 +59,13 @@ public interface WebGraphQlInterceptor { * @return a new interceptor that chains the two */ default WebGraphQlInterceptor andThen(WebGraphQlInterceptor nextInterceptor) { - return (request, chain) -> intercept(request, nextRequest -> nextInterceptor.intercept(nextRequest, chain)); + return (request, chain) -> intercept(request, nextRequest -> { + if (request instanceof WebSocketGraphQlRequest) { + Assert.isTrue(nextRequest instanceof WebSocketGraphQlRequest, + "Expected WebSocketGraphQlRequest but was: " + nextRequest.getClass().getName()); + } + return nextInterceptor.intercept(nextRequest, chain); + }); } /** diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlInterceptor.java b/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlInterceptor.java index ea84e772..1c3107ec 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlInterceptor.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlInterceptor.java @@ -40,12 +40,12 @@ public interface WebSocketGraphQlInterceptor extends WebGraphQlInterceptor { * Handle the {@code "connection_init"} message at the start of a GraphQL over * WebSocket session and return an optional payload for the * {@code "connection_ack"} message to send back. - * @param sessionId the id of the WebSocket session + * @param sessionInfo information about the underlying WebSocket session * @param connectionInitPayload the payload from the {@code "connection_init"} message * @return the payload for the {@code "connection_ack"}, or empty */ default Mono handleConnectionInitialization( - String sessionId, Map connectionInitPayload) { + WebSocketSessionInfo sessionInfo, Map connectionInitPayload) { return Mono.empty(); } @@ -55,25 +55,25 @@ public interface WebSocketGraphQlInterceptor extends WebGraphQlInterceptor { * subscription stream. The underlying {@link org.reactivestreams.Publisher} * for the subscription is automatically cancelled. This callback is for any * additional, or more centralized handling across subscriptions. - * @param sessionId the id of the WebSocket session + * @param sessionInfo information about the underlying WebSocket session * @param subscriptionId the unique id for the subscription; correlates to the * {@link WebGraphQlRequest#getId() requestId} from the original {@code "subscribe"} * message that started the subscription * @return {@code Mono} for the completion of handling */ - default Mono handleCancelledSubscription(String sessionId, String subscriptionId) { + default Mono handleCancelledSubscription(WebSocketSessionInfo sessionInfo, String subscriptionId) { return Mono.empty(); } /** * Invoked when the WebSocket session is closed, from either side. - * @param sessionId the id of the WebSocket session + * @param sessionInfo information about the underlying WebSocket session * @param statusCode the WebSocket "close" status code * @param connectionInitPayload the payload from the {@code "connect_init"} * message received at the start of the connection */ default void handleConnectionClosed( - String sessionId, int statusCode, Map connectionInitPayload) { + WebSocketSessionInfo sessionInfo, int statusCode, Map connectionInitPayload) { } } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlRequest.java b/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlRequest.java new file mode 100644 index 00000000..ee9291a1 --- /dev/null +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketGraphQlRequest.java @@ -0,0 +1,67 @@ +/* + * Copyright 2020-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.server; + + +import java.net.URI; +import java.util.Locale; +import java.util.Map; + +import org.springframework.http.HttpHeaders; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; + + +/** + * {@link org.springframework.graphql.server.WebGraphQlRequest} extension for + * server handling of GraphQL over WebSocket requests. + * + * @author Rossen Stoyanchev + * @since 1.0.0 + */ +public class WebSocketGraphQlRequest extends WebGraphQlRequest { + + private final WebSocketSessionInfo sessionInfo; + + + /** + * Create an instance. + * @param uri the URL for the HTTP request or WebSocket handshake + * @param headers the HTTP request headers + * @param body the deserialized content of the GraphQL request + * @param id the id from the GraphQL over WebSocket {@code "subscribe"} message + * @param locale the locale from the HTTP request, if any + * @param sessionInfo the WebSocket session id + */ + public WebSocketGraphQlRequest( + URI uri, HttpHeaders headers, Map body, String id, @Nullable Locale locale, + WebSocketSessionInfo sessionInfo) { + + super(uri, headers, body, id, locale); + Assert.notNull(sessionInfo, "WebSocketSessionInfo is required"); + this.sessionInfo = sessionInfo; + } + + + /** + * Return information about the underlying WebSocket session. + */ + public WebSocketSessionInfo getSessionInfo() { + return this.sessionInfo; + } + +} diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketSessionInfo.java b/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketSessionInfo.java new file mode 100644 index 00000000..01eeda4e --- /dev/null +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/WebSocketSessionInfo.java @@ -0,0 +1,69 @@ +/* + * 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.server; + +import java.net.InetSocketAddress; +import java.net.URI; +import java.security.Principal; +import java.util.Map; + +import reactor.core.publisher.Mono; + +import org.springframework.http.HttpHeaders; +import org.springframework.lang.Nullable; + +/** + * Expose information about the underlying WebSocketSession including the + * session id, the attributes, and HTTP handshake request. + * + * @author Rossen Stoyanchev + * @since 1.0.0 + */ +public interface WebSocketSessionInfo { + + /** + * Return the id for the WebSocketSession. + */ + String getId(); + + /** + * Return the map with attributes associated with the WebSocket session. + */ + Map getAttributes(); + + /** + * Return the URL for the WebSocket endpoint. + */ + URI getUri(); + + /** + * Return the HTTP headers from the handshake request. + */ + HttpHeaders getHeaders(); + + /** + * Return the principal associated with the handshake request, if any. + */ + Mono getPrincipal(); + + /** + * For a server session this is the remote address where the handshake + * request came from. For a client session, it is {@code null}. + */ + @Nullable + InetSocketAddress getRemoteAddress(); + +} diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandler.java index ab2cf86f..03f935df 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandler.java @@ -16,6 +16,9 @@ package org.springframework.graphql.server.webflux; +import java.net.InetSocketAddress; +import java.net.URI; +import java.security.Principal; import java.time.Duration; import java.util.Arrays; import java.util.Collections; @@ -33,10 +36,12 @@ import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; 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.WebSocketGraphQlRequest; +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.util.Assert; import org.springframework.util.CollectionUtils; @@ -106,6 +111,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { } // Session state + WebSocketSessionInfo sessionInfo = new WebFluxSessionInfo(session); AtomicReference> connectionInitPayloadRef = new AtomicReference<>(); Map subscriptions = new ConcurrentHashMap<>(); @@ -123,7 +129,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { return; } int statusCode = (closeStatus != null ? closeStatus.getCode() : 1005); - this.webSocketInterceptor.handleConnectionClosed(session.getId(), statusCode, connectionInitPayload); + this.webSocketInterceptor.handleConnectionClosed(sessionInfo, statusCode, connectionInitPayload); }) .subscribe(); @@ -139,8 +145,8 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { if (id == null) { return GraphQlStatus.close(session, GraphQlStatus.INVALID_MESSAGE_STATUS); } - WebGraphQlRequest request = new WebGraphQlRequest( - handshakeInfo.getUri(), handshakeInfo.getHeaders(), payload, id, null); + WebSocketGraphQlRequest request = new WebSocketGraphQlRequest( + handshakeInfo.getUri(), handshakeInfo.getHeaders(), payload, id, null, sessionInfo); if (logger.isDebugEnabled()) { logger.debug("Executing: " + request); } @@ -155,7 +161,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { if (subscription != null) { subscription.cancel(); } - return this.webSocketInterceptor.handleCancelledSubscription(session.getId(), id) + return this.webSocketInterceptor.handleCancelledSubscription(sessionInfo, id) .thenMany(Flux.empty()); } return Flux.empty(); @@ -163,7 +169,7 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { if (!connectionInitPayloadRef.compareAndSet(null, payload)) { return GraphQlStatus.close(session, GraphQlStatus.TOO_MANY_INIT_REQUESTS_STATUS); } - return this.webSocketInterceptor.handleConnectionInitialization(session.getId(), payload) + return this.webSocketInterceptor.handleConnectionInitialization(sessionInfo, payload) .defaultIfEmpty(Collections.emptyMap()) .map(ackPayload -> this.codecDelegate.encodeConnectionAck(session, ackPayload)) .flux() @@ -232,6 +238,46 @@ public class GraphQlWebSocketHandler implements WebSocketHandler { } + private static class WebFluxSessionInfo implements WebSocketSessionInfo { + + private final WebSocketSession session; + + private WebFluxSessionInfo(WebSocketSession session) { + this.session = session; + } + + @Override + public String getId() { + return this.session.getId(); + } + + @Override + public Map getAttributes() { + return this.session.getAttributes(); + } + + @Override + public URI getUri() { + return this.session.getHandshakeInfo().getUri(); + } + + @Override + public HttpHeaders getHeaders() { + return this.session.getHandshakeInfo().getHeaders(); + } + + @Override + public Mono getPrincipal() { + return this.session.getHandshakeInfo().getPrincipal(); + } + + @Override + public InetSocketAddress getRemoteAddress() { + return this.session.getHandshakeInfo().getRemoteAddress(); + } + } + + @SuppressWarnings("serial") private static class SubscriptionExistsException extends RuntimeException { } diff --git a/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandler.java b/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandler.java index 588be8e5..aecf07b3 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandler.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandler.java @@ -22,7 +22,9 @@ import java.io.Closeable; import java.io.IOException; import java.io.InputStream; import java.io.OutputStream; +import java.net.InetSocketAddress; import java.net.URI; +import java.security.Principal; import java.time.Duration; import java.util.Arrays; import java.util.Collections; @@ -47,9 +49,10 @@ import reactor.core.scheduler.Schedulers; import org.springframework.graphql.execution.ThreadLocalAccessor; 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.WebSocketGraphQlRequest; +import org.springframework.graphql.server.WebSocketSessionInfo; import org.springframework.graphql.server.support.GraphQlWebSocketMessage; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpInputMessage; @@ -148,7 +151,7 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub return; } - SessionState sessionState = new SessionState(session.getId()); + SessionState sessionState = new SessionState(session.getId(), new WebMvcSessionInfo(session)); this.sessionInfoMap.put(session.getId(), sessionState); Mono.delay(this.initTimeoutDuration) @@ -173,10 +176,10 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub GraphQlWebSocketMessage message = decode(webSocketMessage); String id = message.getId(); Map payload = message.getPayload(); - SessionState sessionState = getSessionInfo(session); + SessionState state = getSessionInfo(session); switch (message.resolvedType()) { case SUBSCRIBE: - if (sessionState.getConnectionInitPayload() == null) { + if (state.getConnectionInitPayload() == null) { GraphQlStatus.closeSession(session, GraphQlStatus.UNAUTHORIZED_STATUS); return; } @@ -187,36 +190,37 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub URI uri = session.getUri(); Assert.notNull(uri, "Expected handshake url"); HttpHeaders headers = session.getHandshakeHeaders(); - WebGraphQlRequest request = new WebGraphQlRequest(uri, headers, payload, id, null); + WebSocketGraphQlRequest request = + new WebSocketGraphQlRequest(uri, headers, payload, id, null, state.getSessionInfo()); if (logger.isDebugEnabled()) { logger.debug("Executing: " + request); } this.graphQlHandler.handleRequest(request) .flatMapMany((response) -> handleResponse(session, request.getId(), response)) - .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread - .subscribe(new SendMessageSubscriber(id, session, sessionState)); + .publishOn(state.getScheduler()) // Serial blocking send via single thread + .subscribe(new SendMessageSubscriber(id, session, state)); return; case PING: session.sendMessage(encode(GraphQlWebSocketMessage.pong(null))); return; case COMPLETE: if (id != null) { - Subscription subscription = sessionState.getSubscriptions().remove(id); + Subscription subscription = state.getSubscriptions().remove(id); if (subscription != null) { subscription.cancel(); } - this.webSocketGraphQlInterceptor.handleCancelledSubscription(session.getId(), id) + this.webSocketGraphQlInterceptor.handleCancelledSubscription(state.getSessionInfo(), id) .block(Duration.ofSeconds(10)); } return; case CONNECTION_INIT: - if (!sessionState.setConnectionInitPayload(payload)) { + if (!state.setConnectionInitPayload(payload)) { GraphQlStatus.closeSession(session, GraphQlStatus.TOO_MANY_INIT_REQUESTS_STATUS); return; } - this.webSocketGraphQlInterceptor.handleConnectionInitialization(session.getId(), payload) + this.webSocketGraphQlInterceptor.handleConnectionInitialization(state.getSessionInfo(), payload) .defaultIfEmpty(Collections.emptyMap()) - .publishOn(sessionState.getScheduler()) // Serial blocking send via single thread + .publishOn(state.getScheduler()) // Serial blocking send via single thread .doOnNext(ackPayload -> { TextMessage outputMessage = encode(GraphQlWebSocketMessage.connectionAck(ackPayload)); try { @@ -311,12 +315,13 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub @Override public void afterConnectionClosed(WebSocketSession session, CloseStatus closeStatus) { String id = session.getId(); - SessionState info = this.sessionInfoMap.remove(id); - if (info != null) { - info.dispose(); - Map connectionInitPayload = info.getConnectionInitPayload(); + SessionState state = this.sessionInfoMap.remove(id); + if (state != null) { + state.dispose(); + Map connectionInitPayload = state.getConnectionInitPayload(); if (connectionInitPayload != null) { - this.webSocketGraphQlInterceptor.handleConnectionClosed(id, closeStatus.getCode(), connectionInitPayload); + this.webSocketGraphQlInterceptor.handleConnectionClosed( + state.getSessionInfo(), closeStatus.getCode(), connectionInitPayload); } } } @@ -434,14 +439,21 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub private static class SessionState { + private final WebSocketSessionInfo sessionInfo; + private final AtomicReference> connectionInitPayloadRef = new AtomicReference<>(); private final Map subscriptions = new ConcurrentHashMap<>(); private final Scheduler scheduler; - SessionState(String sessionId) { - this.scheduler = Schedulers.newSingle("GraphQL-WsSession-" + sessionId); + SessionState(String graphQlSessionId, WebSocketSessionInfo sessionInfo) { + this.sessionInfo = sessionInfo; + this.scheduler = Schedulers.newSingle("GraphQL-WsSession-" + graphQlSessionId); + } + + public WebSocketSessionInfo getSessionInfo() { + return this.sessionInfo; } @Nullable @@ -477,6 +489,48 @@ public class GraphQlWebSocketHandler extends TextWebSocketHandler implements Sub } + + private static class WebMvcSessionInfo implements WebSocketSessionInfo { + + private final WebSocketSession session; + + private WebMvcSessionInfo(WebSocketSession session) { + this.session = session; + } + + @Override + public String getId() { + return this.session.getId(); + } + + @Override + public Map getAttributes() { + return this.session.getAttributes(); + } + + @Override + public URI getUri() { + Assert.notNull(this.session.getUri(), "Expected URI"); + return this.session.getUri(); + } + + @Override + public HttpHeaders getHeaders() { + return this.session.getHandshakeHeaders(); + } + + @Override + public Mono getPrincipal() { + return Mono.justOrEmpty(this.session.getPrincipal()); + } + + @Override + public InetSocketAddress getRemoteAddress() { + return this.session.getRemoteAddress(); + } + } + + private static class SendMessageSubscriber extends BaseSubscriber { private final String subscriptionId; diff --git a/spring-graphql/src/main/java/org/springframework/graphql/support/DefaultExecutionGraphQlRequest.java b/spring-graphql/src/main/java/org/springframework/graphql/support/DefaultExecutionGraphQlRequest.java index cf23ee62..227bed73 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/support/DefaultExecutionGraphQlRequest.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/support/DefaultExecutionGraphQlRequest.java @@ -124,7 +124,7 @@ public class DefaultExecutionGraphQlRequest extends DefaultGraphQlRequest implem @Override public String toString() { - return super.toString() + (getLocale() != null ? ", Locale=" + getLocale() : ""); + return super.toString() + ", id=" + getId() + (getLocale() != null ? ", Locale=" + getLocale() : ""); } } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandlerTests.java index 916d0c1f..5b83f890 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/server/webflux/GraphQlWebSocketHandlerTests.java @@ -42,6 +42,7 @@ 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.WebSocketSessionInfo; import org.springframework.graphql.server.support.GraphQlWebSocketMessage; import org.springframework.graphql.server.support.GraphQlWebSocketMessageType; import org.springframework.http.codec.ServerCodecConfigurer; @@ -152,7 +153,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new WebSocketGraphQlInterceptor() { @Override - public Mono handleConnectionInitialization(String sessionId, Map payload) { + public Mono handleConnectionInitialization(WebSocketSessionInfo info, Map payload) { Object value = payload.get("key"); return Mono.just(Collections.singletonMap("key", value + " acknowledged")); } @@ -192,9 +193,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new WebSocketGraphQlInterceptor() { @Override - public void handleConnectionClosed(String sessionId, int status, Map payload) { + public void handleConnectionClosed(WebSocketSessionInfo info, int status, Map payload) { called.set(true); - assertThat(sessionId).isEqualTo("1"); + assertThat(info.getId()).isEqualTo("1"); assertThat(status).isEqualTo(closeStatus.getCode()); assertThat(payload).hasSize(1).containsEntry("key", "A"); } @@ -212,7 +213,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { new WebSocketGraphQlInterceptor() { @Override - public Mono handleConnectionInitialization(String sessionId, Map payload) { + public Mono handleConnectionInitialization(WebSocketSessionInfo info, Map payload) { return Mono.error(new IllegalStateException()); } }); diff --git a/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandlerTests.java b/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandlerTests.java index 7afad415..4a54d73b 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandlerTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/server/webmvc/GraphQlWebSocketHandlerTests.java @@ -42,6 +42,7 @@ import org.springframework.graphql.server.WebGraphQlHandler; import org.springframework.graphql.server.WebGraphQlInterceptor; import org.springframework.graphql.server.WebSocketGraphQlInterceptor; import org.springframework.graphql.server.WebSocketHandlerTestSupport; +import org.springframework.graphql.server.WebSocketSessionInfo; import org.springframework.graphql.server.support.GraphQlWebSocketMessage; import org.springframework.graphql.server.support.GraphQlWebSocketMessageType; import org.springframework.http.HttpHeaders; @@ -154,7 +155,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { WebSocketGraphQlInterceptor interceptor = new WebSocketGraphQlInterceptor() { @Override - public Mono handleConnectionInitialization(String sessionId, Map payload) { + public Mono handleConnectionInitialization(WebSocketSessionInfo info, Map payload) { Object value = payload.get("key"); return Mono.just(Collections.singletonMap("key", value + " acknowledged")); } @@ -198,9 +199,9 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { WebSocketGraphQlInterceptor interceptor = new WebSocketGraphQlInterceptor() { @Override - public void handleConnectionClosed(String sessionId, int status, Map payload) { + public void handleConnectionClosed(WebSocketSessionInfo info, int status, Map payload) { called.set(true); - assertThat(sessionId).isEqualTo("1"); + assertThat(info.getId()).isEqualTo("1"); assertThat(status).isEqualTo(closeStatus.getCode()); assertThat(payload).hasSize(1).containsEntry("key", "A"); } @@ -225,7 +226,7 @@ public class GraphQlWebSocketHandlerTests extends WebSocketHandlerTestSupport { WebSocketGraphQlInterceptor interceptor = new WebSocketGraphQlInterceptor() { @Override - public Mono handleConnectionInitialization(String sessionId, Map payload) { + public Mono handleConnectionInitialization(WebSocketSessionInfo info, Map payload) { return Mono.error(new IllegalStateException()); } };