diff --git a/samples/webflux-security/src/test/java/io/spring/sample/graphql/WebFluxSecuritySampleTests.java b/samples/webflux-security/src/test/java/io/spring/sample/graphql/WebFluxSecuritySampleTests.java index cffee115..43b91673 100644 --- a/samples/webflux-security/src/test/java/io/spring/sample/graphql/WebFluxSecuritySampleTests.java +++ b/samples/webflux-security/src/test/java/io/spring/sample/graphql/WebFluxSecuritySampleTests.java @@ -22,10 +22,9 @@ import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; -import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.boot.test.context.SpringBootTest.WebEnvironment; -import org.springframework.boot.web.server.LocalServerPort; +import org.springframework.boot.test.web.server.LocalServerPort; import org.springframework.graphql.execution.ErrorType; import org.springframework.graphql.test.tester.WebGraphQlTester; import org.springframework.graphql.test.tester.WebSocketGraphQlTester; diff --git a/samples/webflux-websocket/src/test/java/io/spring/sample/graphql/WebFluxWebSocketSampleIntegrationTests.java b/samples/webflux-websocket/src/test/java/io/spring/sample/graphql/WebFluxWebSocketSampleIntegrationTests.java index 3f60f381..fca21379 100644 --- a/samples/webflux-websocket/src/test/java/io/spring/sample/graphql/WebFluxWebSocketSampleIntegrationTests.java +++ b/samples/webflux-websocket/src/test/java/io/spring/sample/graphql/WebFluxWebSocketSampleIntegrationTests.java @@ -24,7 +24,7 @@ import reactor.test.StepVerifier; import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.test.context.SpringBootTest; -import org.springframework.boot.web.server.LocalServerPort; +import org.springframework.boot.test.web.server.LocalServerPort; import org.springframework.graphql.test.tester.GraphQlTester; import org.springframework.graphql.test.tester.WebSocketGraphQlTester; import org.springframework.web.reactive.socket.client.ReactorNettyWebSocketClient; 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 6e1bd8e4..f89a0041 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 @@ -27,7 +27,9 @@ import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import reactor.core.Scannable; import reactor.core.publisher.Flux; +import reactor.core.publisher.FluxSink; import reactor.core.publisher.Mono; +import reactor.core.publisher.MonoSink; import reactor.core.publisher.Sinks; import org.springframework.graphql.GraphQlRequest; @@ -378,11 +380,9 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { private final AtomicLong requestIndex = new AtomicLong(); - private final Sinks.Many requestSink = Sinks.many().unicast().onBackpressureBuffer(); + private final RequestSink requestSink = new RequestSink(); - private final Map responseMap = new ConcurrentHashMap<>(); - - private final Map subscriptionMap = new ConcurrentHashMap<>(); + private final Map requestStateMap = new ConcurrentHashMap<>(); GraphQlSession(WebSocketSession webSocketSession) { @@ -394,62 +394,49 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { * Return the {@code Flux} of GraphQL requests to send as WebSocket messages. */ public Flux getRequestFlux() { - return this.requestSink.asFlux(); + return this.requestSink.getRequestFlux(); } + + // Outbound messages + public Mono execute(GraphQlRequest request) { String id = String.valueOf(this.requestIndex.incrementAndGet()); - try { - GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request); - ResponseState state = new ResponseState(request); - this.responseMap.put(id, state); - trySend(message); - return state.sink().asMono().doOnCancel(() -> this.responseMap.remove(id)); - } - catch (Exception ex) { - this.responseMap.remove(id); - return Mono.error(ex); - } + return Mono.create(sink -> { + SingleResponseRequestState state = new SingleResponseRequestState(request, sink); + this.requestStateMap.put(id, state); + try { + GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request); + this.requestSink.sendRequest(message); + } + catch (Exception ex) { + this.requestStateMap.remove(id); + sink.error(ex); + } + }).doOnCancel(() -> this.requestStateMap.remove(id)); } public Flux executeSubscription(GraphQlRequest request) { String id = String.valueOf(this.requestIndex.incrementAndGet()); - try { - GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request); - SubscriptionState state = new SubscriptionState(request); - this.subscriptionMap.put(id, state); - trySend(message); - return state.sink().asFlux().doOnCancel(() -> stopSubscription(id)); - } - catch (Exception ex) { - this.subscriptionMap.remove(id); - return Flux.error(ex); - } - } - - public void sendPong(@Nullable Map payload) { - GraphQlWebSocketMessage message = GraphQlWebSocketMessage.pong(payload); - trySend(message); - } - - // TODO: queue to serialize sending? - - private void trySend(GraphQlWebSocketMessage message) { - Sinks.EmitResult emitResult = null; - for (int i = 0; i < 100; i++) { - emitResult = this.requestSink.tryEmitNext(message); - if (emitResult != Sinks.EmitResult.FAIL_NON_SERIALIZED) { - break; + return Flux.create(sink -> { + SubscriptionRequestState state = new SubscriptionRequestState(request, sink); + this.requestStateMap.put(id, state); + try { + GraphQlWebSocketMessage message = GraphQlWebSocketMessage.subscribe(id, request); + this.requestSink.sendRequest(message); } - } - Assert.state(emitResult.isSuccess(), "Failed to send request: " + emitResult); + catch (Exception ex) { + this.requestStateMap.remove(id); + sink.error(ex); + } + }).doOnCancel(() -> stopSubscription(id)); } private void stopSubscription(String id) { - SubscriptionState state = this.subscriptionMap.remove(id); + RequestState state = this.requestStateMap.remove(id); if (state != null) { try { - trySend(GraphQlWebSocketMessage.complete(id)); + this.requestSink.sendRequest(GraphQlWebSocketMessage.complete(id)); } catch (Exception ex) { if (logger.isErrorEnabled()) { @@ -462,34 +449,34 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { } } + public void sendPong(@Nullable Map payload) { + GraphQlWebSocketMessage message = GraphQlWebSocketMessage.pong(payload); + this.requestSink.sendRequest(message); + } + + + // Inbound messages + /** * Handle a "next" message and route to its recipient. */ public void handleNext(GraphQlWebSocketMessage message) { String id = message.getId(); - ResponseState responseState = this.responseMap.remove(id); - SubscriptionState subscriptionState = this.subscriptionMap.get(id); - - if (responseState == null && subscriptionState == null) { + RequestState requestState = this.requestStateMap.get(id); + if (requestState == null) { if (logger.isDebugEnabled()) { - logger.debug("No receiver for message: " + message); + logger.debug("No receiver for: " + message); } return; } - Map responseMap = message.getPayload(); - GraphQlResponse graphQlResponse = new ResponseMapGraphQlResponse(responseMap); - - Sinks.EmitResult emitResult = (responseState != null ? - responseState.sink().tryEmitValue(graphQlResponse) : - subscriptionState.sink().tryEmitNext(graphQlResponse)); - - if (emitResult.isFailure()) { - // Just log: cannot overflow, is serialized, and cancel is handled in doOnCancel - if (logger.isDebugEnabled()) { - logger.debug("Message: " + message + " could not be emitted: " + emitResult); - } + if (requestState instanceof SingleResponseRequestState) { + this.requestStateMap.remove(id); } + + Map payload = message.getPayload(); + GraphQlResponse graphQlResponse = new ResponseMapGraphQlResponse(payload); + requestState.handleResponse(graphQlResponse); } /** @@ -498,12 +485,10 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { */ public void handleError(GraphQlWebSocketMessage message) { String id = message.getId(); - ResponseState responseState = this.responseMap.remove(id); - SubscriptionState subscriptionState = this.subscriptionMap.remove(id); - - if (responseState == null && subscriptionState == null) { + RequestState requestState = this.requestStateMap.remove(id); + if (requestState == null) { if (logger.isDebugEnabled()) { - logger.debug("No receiver for message: " + message); + logger.debug("No receiver for: " + message); } return; } @@ -511,18 +496,13 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { List> errorList = message.getPayload(); GraphQlResponse response = new ResponseMapGraphQlResponse(Collections.singletonMap("errors", errorList)); - Sinks.EmitResult emitResult; - if (responseState != null) { - emitResult = responseState.sink().tryEmitValue(response); + if (requestState instanceof SingleResponseRequestState) { + requestState.handleResponse(response); } else { List errors = response.getErrors(); - Exception ex = new SubscriptionErrorException(subscriptionState.request(), errors); - emitResult = subscriptionState.sink().tryEmitError(ex); - } - - if (emitResult.isFailure() && logger.isDebugEnabled()) { - logger.debug("Error: " + message + " could not be emitted: " + emitResult); + Exception ex = new SubscriptionErrorException(requestState.getRequest(), errors); + requestState.handlerError(ex); } } @@ -530,15 +510,15 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { * Handle a "complete" message. */ public void handleComplete(GraphQlWebSocketMessage message) { - ResponseState responseState = this.responseMap.remove(message.getId()); - SubscriptionState subscriptionState = this.subscriptionMap.remove(message.getId()); - - if (responseState != null) { - responseState.sink().tryEmitEmpty(); - } - else if (subscriptionState != null) { - subscriptionState.sink().tryEmitComplete(); + String id = message.getId(); + RequestState requestState = this.requestStateMap.remove(id); + if (requestState == null) { + if (logger.isDebugEnabled()) { + logger.debug("No receiver for': " + message); + } + return; } + requestState.handleCompletion(); } /** @@ -560,10 +540,8 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { * Terminate and clean all in-progress requests with the given error. */ public void terminateRequests(String message, CloseStatus status) { - this.responseMap.values().forEach(info -> info.emitDisconnectError(message, status)); - this.subscriptionMap.values().forEach(info -> info.emitDisconnectError(message, status) ); - this.responseMap.clear(); - this.subscriptionMap.clear(); + this.requestStateMap.values().forEach(info -> info.emitDisconnectError(message, status)); + this.requestStateMap.clear(); } @Override @@ -611,26 +589,49 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { } + /** + * Holds the request {@code Flux} and associated {@link FluxSink}. + */ + private static class RequestSink { + + @Nullable + private FluxSink requestSink; + + private final Flux requestFlux = Flux.create(sink -> { + Assert.state(this.requestSink == null, "Expected single subscriber only for outbound messages"); + this.requestSink = sink; + }); + + public Flux getRequestFlux() { + return this.requestFlux; + } + + public void sendRequest(GraphQlWebSocketMessage message) { + Assert.state(this.requestSink != null, "Unexpected request before Flux is subscribed to"); + this.requestSink.next(message); + } + + } + + /** * Base class, state container for any request type. */ - private abstract static class AbstractRequestState { + private interface RequestState { - private final GraphQlRequest request; + GraphQlRequest getRequest(); - public AbstractRequestState(GraphQlRequest request) { - this.request = request; + void handleResponse(GraphQlResponse response); + + void handlerError(Throwable ex); + + void handleCompletion(); + + default void emitDisconnectError(String message, CloseStatus closeStatus) { + emitDisconnectError(new WebSocketDisconnectedException(message, getRequest(), closeStatus)); } - public GraphQlRequest request() { - return this.request; - } - - public void emitDisconnectError(String message, CloseStatus closeStatus) { - emitDisconnectError(new WebSocketDisconnectedException(message, this.request, closeStatus)); - } - - protected abstract void emitDisconnectError(WebSocketDisconnectedException ex); + void emitDisconnectError(WebSocketDisconnectedException ex); } @@ -638,21 +639,40 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { /** * State container for a request that emits a single response. */ - private static class ResponseState extends AbstractRequestState { + private static class SingleResponseRequestState implements RequestState { - private final Sinks.One sink = Sinks.one(); + private final GraphQlRequest request; - ResponseState(GraphQlRequest request) { - super(request); - } + private final MonoSink responseSink; - public Sinks.One sink() { - return this.sink; + SingleResponseRequestState(GraphQlRequest request, MonoSink responseSink) { + this.request = request; + this.responseSink = responseSink; } @Override - protected void emitDisconnectError(WebSocketDisconnectedException ex) { - this.sink.tryEmitError(ex); + public GraphQlRequest getRequest() { + return this.request; + } + + @Override + public void handleResponse(GraphQlResponse response) { + this.responseSink.success(response); + } + + @Override + public void handlerError(Throwable ex) { + this.responseSink.error(ex); + } + + @Override + public void handleCompletion() { + this.responseSink.success(); + } + + @Override + public void emitDisconnectError(WebSocketDisconnectedException ex) { + handlerError(ex); } } @@ -661,21 +681,40 @@ final class WebSocketGraphQlTransport implements GraphQlTransport { /** * State container for a subscription request that emits a stream of responses. */ - private static class SubscriptionState extends AbstractRequestState { + private static class SubscriptionRequestState implements RequestState { - private final Sinks.Many sink = Sinks.many().unicast().onBackpressureBuffer(); + private final GraphQlRequest request; - SubscriptionState(GraphQlRequest request) { - super(request); - } + private final FluxSink responseSink; - public Sinks.Many sink() { - return this.sink; + SubscriptionRequestState(GraphQlRequest request, FluxSink responseSink) { + this.request = request; + this.responseSink = responseSink; } @Override - protected void emitDisconnectError(WebSocketDisconnectedException ex) { - this.sink.tryEmitError(ex); + public GraphQlRequest getRequest() { + return request; + } + + @Override + public void handleResponse(GraphQlResponse response) { + this.responseSink.next(response); + } + + @Override + public void handlerError(Throwable ex) { + this.responseSink.error(ex); + } + + @Override + public void handleCompletion() { + this.responseSink.complete(); + } + + @Override + public void emitDisconnectError(WebSocketDisconnectedException ex) { + handlerError(ex); } }