Refactor use of sinks in WebSocketGraphQlTransport

Replace use of Sinks which does not deal with concurrent sending and
use MonoSink and FluxSink instead.

Closes gh-388
This commit is contained in:
rstoyanchev
2022-05-13 13:59:39 +01:00
parent 57e9ccbfd8
commit c43a444c78
3 changed files with 161 additions and 123 deletions

View File

@@ -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;

View File

@@ -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;

View File

@@ -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<GraphQlWebSocketMessage> requestSink = Sinks.many().unicast().onBackpressureBuffer();
private final RequestSink requestSink = new RequestSink();
private final Map<String, ResponseState> responseMap = new ConcurrentHashMap<>();
private final Map<String, SubscriptionState> subscriptionMap = new ConcurrentHashMap<>();
private final Map<String, RequestState> 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<GraphQlWebSocketMessage> getRequestFlux() {
return this.requestSink.asFlux();
return this.requestSink.getRequestFlux();
}
// Outbound messages
public Mono<GraphQlResponse> 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.<GraphQlResponse>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<GraphQlResponse> 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<String, Object> 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.<GraphQlResponse>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<String, Object> 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<String, Object> 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<String, Object> 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<Map<String, Object>> 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<ResponseError> 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<GraphQlWebSocketMessage> requestSink;
private final Flux<GraphQlWebSocketMessage> requestFlux = Flux.create(sink -> {
Assert.state(this.requestSink == null, "Expected single subscriber only for outbound messages");
this.requestSink = sink;
});
public Flux<GraphQlWebSocketMessage> 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<GraphQlResponse> sink = Sinks.one();
private final GraphQlRequest request;
ResponseState(GraphQlRequest request) {
super(request);
}
private final MonoSink<GraphQlResponse> responseSink;
public Sinks.One<GraphQlResponse> sink() {
return this.sink;
SingleResponseRequestState(GraphQlRequest request, MonoSink<GraphQlResponse> 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<GraphQlResponse> sink = Sinks.many().unicast().onBackpressureBuffer();
private final GraphQlRequest request;
SubscriptionState(GraphQlRequest request) {
super(request);
}
private final FluxSink<GraphQlResponse> responseSink;
public Sinks.Many<GraphQlResponse> sink() {
return this.sink;
SubscriptionRequestState(GraphQlRequest request, FluxSink<GraphQlResponse> 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);
}
}