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:
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user