From 12f80be1f0738109a91471a2dca56f7d02614863 Mon Sep 17 00:00:00 2001 From: Rossen Stoyanchev Date: Fri, 23 Dec 2016 18:04:53 -0500 Subject: [PATCH] AbstractListenerWebSocketSession handles Mono The HandlerSubcriber from each listener session implementation is now consolidated into AbstractListenerWebSocketSession since the handling of onComplete or onError in any case is about delegating to the session. This also allows for the UndertowWebSocketHandlerAdapter to become simply an (Undertow) AbstractReceiveListener. Issue: SPR-14527 --- .../AbstractListenerWebSocketSession.java | 59 +++++++- .../adapter/JettyWebSocketHandlerAdapter.java | 51 +------ .../socket/adapter/JettyWebSocketSession.java | 11 +- .../StandardWebSocketHandlerAdapter.java | 51 +------ .../adapter/StandardWebSocketSession.java | 11 +- .../UndertowWebSocketHandlerAdapter.java | 141 +++++------------- .../adapter/UndertowWebSocketSession.java | 11 +- .../socket/client/JettyWebSocketClient.java | 12 +- .../client/StandardWebSocketClient.java | 10 +- .../client/UndertowWebSocketClient.java | 11 +- .../upgrade/JettyRequestUpgradeStrategy.java | 2 +- .../UndertowRequestUpgradeStrategy.java | 45 +++--- 12 files changed, 172 insertions(+), 243 deletions(-) diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/AbstractListenerWebSocketSession.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/AbstractListenerWebSocketSession.java index 4784b5bbd8..b8a1d4a282 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/AbstractListenerWebSocketSession.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/AbstractListenerWebSocketSession.java @@ -20,8 +20,11 @@ import java.io.IOException; import java.util.concurrent.atomic.AtomicBoolean; import org.reactivestreams.Publisher; +import org.reactivestreams.Subscriber; +import org.reactivestreams.Subscription; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import reactor.core.publisher.MonoProcessor; import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.http.server.reactive.AbstractListenerReadPublisher; @@ -37,11 +40,15 @@ import org.springframework.web.reactive.socket.WebSocketSession; * event-listener WebSocket APIs (e.g. Java WebSocket API JSR-356, Jetty, * Undertow) and Reactive Streams. * + *

Also an implementation of {@link Subscriber} so it can be used as + * the completion subscriber for session handling + * * @author Violeta Georgieva * @author Rossen Stoyanchev * @since 5.0 */ -public abstract class AbstractListenerWebSocketSession extends AbstractWebSocketSession { +public abstract class AbstractListenerWebSocketSession extends AbstractWebSocketSession + implements Subscriber { /** * The "back-pressure" buffer size to use if the underlying WebSocket API @@ -50,6 +57,8 @@ public abstract class AbstractListenerWebSocketSession extends AbstractWebSoc private static final int RECEIVE_BUFFER_SIZE = 8192; + private final MonoProcessor completionMono; + private final WebSocketReceivePublisher receivePublisher = new WebSocketReceivePublisher(); private volatile WebSocketSendProcessor sendProcessor; @@ -57,10 +66,28 @@ public abstract class AbstractListenerWebSocketSession extends AbstractWebSoc private final AtomicBoolean sendCalled = new AtomicBoolean(); + /** + * Base constructor. + * @param delegate the native WebSocket session, channel, or connection + * @param id the session id + * @param handshakeInfo the handshake info + * @param bufferFactory the DataBuffer factor for the current connection + */ public AbstractListenerWebSocketSession(T delegate, String id, HandshakeInfo handshakeInfo, DataBufferFactory bufferFactory) { + this(delegate, id, handshakeInfo, bufferFactory, null); + } + + /** + * Alternative constructor with completion {@link Mono} to propagate + * the session completion (success or error) (for client-side use). + */ + public AbstractListenerWebSocketSession(T delegate, String id, HandshakeInfo handshakeInfo, + DataBufferFactory bufferFactory, MonoProcessor completionMono) { + super(delegate, id, handshakeInfo, bufferFactory); + this.completionMono = completionMono; } @@ -145,6 +172,36 @@ public abstract class AbstractListenerWebSocketSession extends AbstractWebSoc } + // Subscriber implementation + + @Override + public void onSubscribe(Subscription subscription) { + subscription.request(Long.MAX_VALUE); + } + + @Override + public void onNext(Void aVoid) { + // no op + } + + @Override + public void onError(Throwable ex) { + if (this.completionMono != null) { + this.completionMono.onError(ex); + } + int code = CloseStatus.SERVER_ERROR.getCode(); + close(new CloseStatus(code, ex.getMessage())); + } + + @Override + public void onComplete() { + if (this.completionMono != null) { + this.completionMono.onComplete(); + } + close(); + } + + private final class WebSocketReceivePublisher extends AbstractListenerReadPublisher { private volatile WebSocketMessage webSocketMessage; diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketHandlerAdapter.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketHandlerAdapter.java index 29abdd3a9e..f27c056ceb 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketHandlerAdapter.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketHandlerAdapter.java @@ -29,9 +29,6 @@ import org.eclipse.jetty.websocket.api.annotations.OnWebSocketMessage; import org.eclipse.jetty.websocket.api.annotations.WebSocket; import org.eclipse.jetty.websocket.api.extensions.Frame; import org.eclipse.jetty.websocket.common.OpCode; -import org.reactivestreams.Subscriber; -import org.reactivestreams.Subscription; -import reactor.core.publisher.MonoProcessor; import org.springframework.core.io.buffer.DataBuffer; import org.springframework.util.Assert; @@ -57,8 +54,6 @@ public class JettyWebSocketHandlerAdapter { private final WebSocketHandler delegateHandler; - private final MonoProcessor completionMono; - private final Function sessionFactory; private JettyWebSocketSession delegateSession; @@ -67,24 +62,17 @@ public class JettyWebSocketHandlerAdapter { public JettyWebSocketHandlerAdapter(WebSocketHandler handler, Function sessionFactory) { - this(handler, null, sessionFactory); - } - - public JettyWebSocketHandlerAdapter(WebSocketHandler handler, MonoProcessor completionMono, - Function sessionFactory) { - Assert.notNull("WebSocketHandler is required"); Assert.notNull("'sessionFactory' is required"); this.delegateHandler = handler; - this.completionMono = completionMono; this.sessionFactory = sessionFactory; } + @OnWebSocketConnect public void onWebSocketConnect(Session session) { this.delegateSession = sessionFactory.apply(session); - HandlerResultSubscriber subscriber = new HandlerResultSubscriber(); - this.delegateHandler.handle(this.delegateSession).subscribe(subscriber); + this.delegateHandler.handle(this.delegateSession).subscribe(this.delegateSession); } @OnWebSocketMessage @@ -150,39 +138,4 @@ public class JettyWebSocketHandlerAdapter { } } - - private final class HandlerResultSubscriber implements Subscriber { - - @Override - public void onSubscribe(Subscription subscription) { - subscription.request(Long.MAX_VALUE); - } - - @Override - public void onNext(Void aVoid) { - // no op - } - - @Override - public void onError(Throwable ex) { - if (completionMono != null) { - completionMono.onError(ex); - } - if (delegateSession != null) { - int code = CloseStatus.SERVER_ERROR.getCode(); - delegateSession.close(new CloseStatus(code, ex.getMessage())); - } - } - - @Override - public void onComplete() { - if (completionMono != null) { - completionMono.onComplete(); - } - if (delegateSession != null) { - delegateSession.close(); - } - } - } - } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketSession.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketSession.java index 6c7ddabd6d..97549e1a12 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketSession.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/JettyWebSocketSession.java @@ -23,6 +23,7 @@ import java.nio.charset.StandardCharsets; import org.eclipse.jetty.websocket.api.Session; import org.eclipse.jetty.websocket.api.WriteCallback; import reactor.core.publisher.Mono; +import reactor.core.publisher.MonoProcessor; import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.util.ObjectUtils; @@ -42,8 +43,14 @@ import org.springframework.web.reactive.socket.WebSocketSession; public class JettyWebSocketSession extends AbstractListenerWebSocketSession { - public JettyWebSocketSession(Session session, HandshakeInfo info, DataBufferFactory bufferFactory) { - super(session, ObjectUtils.getIdentityHexString(session), info, bufferFactory); + public JettyWebSocketSession(Session session, HandshakeInfo info, DataBufferFactory factory) { + this(session, info, factory, null); + } + + public JettyWebSocketSession(Session session, HandshakeInfo info, DataBufferFactory factory, + MonoProcessor completionMono) { + + super(session, ObjectUtils.getIdentityHexString(session), info, factory, completionMono); } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketHandlerAdapter.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketHandlerAdapter.java index 351fd88de6..d0c9aa613d 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketHandlerAdapter.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketHandlerAdapter.java @@ -25,10 +25,6 @@ import javax.websocket.EndpointConfig; import javax.websocket.PongMessage; import javax.websocket.Session; -import org.reactivestreams.Subscriber; -import org.reactivestreams.Subscription; -import reactor.core.publisher.MonoProcessor; - import org.springframework.core.io.buffer.DataBuffer; import org.springframework.util.Assert; import org.springframework.web.reactive.socket.CloseStatus; @@ -49,8 +45,6 @@ public class StandardWebSocketHandlerAdapter extends Endpoint { private final WebSocketHandler delegateHandler; - private final MonoProcessor completionMono; - private Function sessionFactory; private StandardWebSocketSession delegateSession; @@ -59,16 +53,9 @@ public class StandardWebSocketHandlerAdapter extends Endpoint { public StandardWebSocketHandlerAdapter(WebSocketHandler handler, Function sessionFactory) { - this(handler, null, sessionFactory); - } - - public StandardWebSocketHandlerAdapter(WebSocketHandler handler, MonoProcessor completionMono, - Function sessionFactory) { - Assert.notNull("WebSocketHandler is required"); Assert.notNull("'sessionFactory' is required"); this.delegateHandler = handler; - this.completionMono = completionMono; this.sessionFactory = sessionFactory; } @@ -91,8 +78,7 @@ public class StandardWebSocketHandlerAdapter extends Endpoint { this.delegateSession.handleMessage(webSocketMessage.getType(), webSocketMessage); }); - HandlerResultSubscriber resultSubscriber = new HandlerResultSubscriber(); - this.delegateHandler.handle(this.delegateSession).subscribe(resultSubscriber); + this.delegateHandler.handle(this.delegateSession).subscribe(this.delegateSession); } private WebSocketMessage toMessage(T message) { @@ -130,39 +116,4 @@ public class StandardWebSocketHandlerAdapter extends Endpoint { } } - - private final class HandlerResultSubscriber implements Subscriber { - - @Override - public void onSubscribe(Subscription subscription) { - subscription.request(Long.MAX_VALUE); - } - - @Override - public void onNext(Void aVoid) { - // no op - } - - @Override - public void onError(Throwable ex) { - if (completionMono != null) { - completionMono.onError(ex); - } - if (delegateSession != null) { - int code = CloseStatus.SERVER_ERROR.getCode(); - delegateSession.close(new CloseStatus(code, ex.getMessage())); - } - } - - @Override - public void onComplete() { - if (completionMono != null) { - completionMono.onComplete(); - } - if (delegateSession != null) { - delegateSession.close(); - } - } - } - } \ No newline at end of file diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketSession.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketSession.java index a0c5ffdc59..72cd740509 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketSession.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/StandardWebSocketSession.java @@ -26,6 +26,7 @@ import javax.websocket.SendResult; import javax.websocket.Session; import reactor.core.publisher.Mono; +import reactor.core.publisher.MonoProcessor; import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.web.reactive.socket.CloseStatus; @@ -44,8 +45,14 @@ import org.springframework.web.reactive.socket.WebSocketSession; public class StandardWebSocketSession extends AbstractListenerWebSocketSession { - public StandardWebSocketSession(Session session, HandshakeInfo info, DataBufferFactory bufferFactory) { - super(session, session.getId(), info, bufferFactory); + public StandardWebSocketSession(Session session, HandshakeInfo info, DataBufferFactory factory) { + this(session, info, factory, null); + } + + public StandardWebSocketSession(Session session, HandshakeInfo info, DataBufferFactory factory, + MonoProcessor completionMono) { + + super(session, session.getId(), info, factory, completionMono); } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketHandlerAdapter.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketHandlerAdapter.java index 0b5322e223..313bc0997d 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketHandlerAdapter.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketHandlerAdapter.java @@ -25,9 +25,6 @@ import io.undertow.websockets.core.BufferedBinaryMessage; import io.undertow.websockets.core.BufferedTextMessage; import io.undertow.websockets.core.CloseMessage; import io.undertow.websockets.core.WebSocketChannel; -import org.reactivestreams.Subscriber; -import org.reactivestreams.Subscription; -import reactor.core.publisher.MonoProcessor; import org.springframework.core.io.buffer.DataBuffer; import org.springframework.util.Assert; @@ -35,7 +32,6 @@ import org.springframework.web.reactive.socket.CloseStatus; import org.springframework.web.reactive.socket.WebSocketHandler; import org.springframework.web.reactive.socket.WebSocketMessage; import org.springframework.web.reactive.socket.WebSocketMessage.Type; -import org.springframework.web.reactive.socket.WebSocketSession; /** * Undertow {@link WebSocketConnectionCallback} implementation that adapts and @@ -45,118 +41,61 @@ import org.springframework.web.reactive.socket.WebSocketSession; * @author Rossen Stoyanchev * @since 5.0 */ -public class UndertowWebSocketHandlerAdapter { +public class UndertowWebSocketHandlerAdapter extends AbstractReceiveListener { - private final WebSocketHandler delegateHandler; - - private final MonoProcessor completionMono; - - private UndertowWebSocketSession delegateSession; + private final UndertowWebSocketSession session; - public UndertowWebSocketHandlerAdapter(WebSocketHandler handler) { - this(handler, null); - } - - public UndertowWebSocketHandlerAdapter(WebSocketHandler handler, MonoProcessor completionMono) { - - Assert.notNull("WebSocketHandler is required"); - Assert.notNull("'sessionFactory' is required"); - this.delegateHandler = handler; - this.completionMono = completionMono; + public UndertowWebSocketHandlerAdapter(UndertowWebSocketSession session) { + Assert.notNull("UndertowWebSocketSession is required"); + this.session = session; } - public void handle(UndertowWebSocketSession webSocketSession) { - this.delegateSession = webSocketSession; - webSocketSession.getDelegate().getReceiveSetter().set(new UndertowReceiveListener()); - webSocketSession.getDelegate().resumeReceives(); - - HandlerResultSubscriber resultSubscriber = new HandlerResultSubscriber(); - this.delegateHandler.handle(this.delegateSession).subscribe(resultSubscriber); + @Override + protected void onFullTextMessage(WebSocketChannel channel, BufferedTextMessage message) { + this.session.handleMessage(Type.TEXT, toMessage(Type.TEXT, message.getData())); } - - private final class UndertowReceiveListener extends AbstractReceiveListener { - - @Override - protected void onFullTextMessage(WebSocketChannel channel, BufferedTextMessage message) { - delegateSession.handleMessage(Type.TEXT, toMessage(Type.TEXT, message.getData())); - } - - @Override - protected void onFullBinaryMessage(WebSocketChannel channel, BufferedBinaryMessage message) { - delegateSession.handleMessage(Type.BINARY, toMessage(Type.BINARY, message.getData().getResource())); - message.getData().free(); - } - - @Override - protected void onFullPongMessage(WebSocketChannel channel, BufferedBinaryMessage message) { - delegateSession.handleMessage(Type.PONG, toMessage(Type.PONG, message.getData().getResource())); - message.getData().free(); - } - - @Override - protected void onFullCloseMessage(WebSocketChannel channel, BufferedBinaryMessage message) { - CloseMessage closeMessage = new CloseMessage(message.getData().getResource()); - delegateSession.handleClose(new CloseStatus(closeMessage.getCode(), closeMessage.getReason())); - message.getData().free(); - } - - @Override - protected void onError(WebSocketChannel channel, Throwable error) { - delegateSession.handleError(error); - } - - private WebSocketMessage toMessage(Type type, T message) { - WebSocketSession session = delegateSession; - Assert.state(session != null, "Cannot create message without a session"); - if (Type.TEXT.equals(type)) { - byte[] bytes = ((String) message).getBytes(StandardCharsets.UTF_8); - return new WebSocketMessage(Type.TEXT, session.bufferFactory().wrap(bytes)); - } - else if (Type.BINARY.equals(type)) { - DataBuffer buffer = session.bufferFactory().allocateBuffer().write((ByteBuffer[]) message); - return new WebSocketMessage(Type.BINARY, buffer); - } - else if (Type.PONG.equals(type)) { - DataBuffer buffer = session.bufferFactory().allocateBuffer().write((ByteBuffer[]) message); - return new WebSocketMessage(Type.PONG, buffer); - } - else { - throw new IllegalArgumentException("Unexpected message type: " + message); - } - } + @Override + protected void onFullBinaryMessage(WebSocketChannel channel, BufferedBinaryMessage message) { + this.session.handleMessage(Type.BINARY, toMessage(Type.BINARY, message.getData().getResource())); + message.getData().free(); } + @Override + protected void onFullPongMessage(WebSocketChannel channel, BufferedBinaryMessage message) { + this.session.handleMessage(Type.PONG, toMessage(Type.PONG, message.getData().getResource())); + message.getData().free(); + } - private final class HandlerResultSubscriber implements Subscriber { + @Override + protected void onFullCloseMessage(WebSocketChannel channel, BufferedBinaryMessage message) { + CloseMessage closeMessage = new CloseMessage(message.getData().getResource()); + this.session.handleClose(new CloseStatus(closeMessage.getCode(), closeMessage.getReason())); + message.getData().free(); + } - @Override - public void onSubscribe(Subscription subscription) { - subscription.request(Long.MAX_VALUE); + @Override + protected void onError(WebSocketChannel channel, Throwable error) { + this.session.handleError(error); + } + + private WebSocketMessage toMessage(Type type, T message) { + if (Type.TEXT.equals(type)) { + byte[] bytes = ((String) message).getBytes(StandardCharsets.UTF_8); + return new WebSocketMessage(Type.TEXT, session.bufferFactory().wrap(bytes)); } - - @Override - public void onNext(Void aVoid) { - // no op + else if (Type.BINARY.equals(type)) { + DataBuffer buffer = session.bufferFactory().allocateBuffer().write((ByteBuffer[]) message); + return new WebSocketMessage(Type.BINARY, buffer); } - - @Override - public void onError(Throwable ex) { - if (completionMono != null) { - completionMono.onError(ex); - } - int code = CloseStatus.SERVER_ERROR.getCode(); - delegateSession.close(new CloseStatus(code, ex.getMessage())); + else if (Type.PONG.equals(type)) { + DataBuffer buffer = session.bufferFactory().allocateBuffer().write((ByteBuffer[]) message); + return new WebSocketMessage(Type.PONG, buffer); } - - @Override - public void onComplete() { - if (completionMono != null) { - completionMono.onComplete(); - } - delegateSession.close(); + else { + throw new IllegalArgumentException("Unexpected message type: " + message); } } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketSession.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketSession.java index 9606818f22..90a0ea6e5a 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketSession.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/adapter/UndertowWebSocketSession.java @@ -25,6 +25,7 @@ import io.undertow.websockets.core.WebSocketCallback; import io.undertow.websockets.core.WebSocketChannel; import io.undertow.websockets.core.WebSockets; import reactor.core.publisher.Mono; +import reactor.core.publisher.MonoProcessor; import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.util.ObjectUtils; @@ -44,10 +45,14 @@ import org.springframework.web.reactive.socket.WebSocketSession; public class UndertowWebSocketSession extends AbstractListenerWebSocketSession { - public UndertowWebSocketSession(WebSocketChannel channel, HandshakeInfo handshakeInfo, - DataBufferFactory bufferFactory) { + public UndertowWebSocketSession(WebSocketChannel channel, HandshakeInfo info, DataBufferFactory factory) { + this(channel, info, factory, null); + } - super(channel, ObjectUtils.getIdentityHexString(channel), handshakeInfo, bufferFactory); + public UndertowWebSocketSession(WebSocketChannel channel, HandshakeInfo info, + DataBufferFactory factory, MonoProcessor completionMono) { + + super(channel, ObjectUtils.getIdentityHexString(channel), info, factory, completionMono); } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/JettyWebSocketClient.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/JettyWebSocketClient.java index 6341206718..397590742c 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/JettyWebSocketClient.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/JettyWebSocketClient.java @@ -82,7 +82,7 @@ public class JettyWebSocketClient extends WebSocketClientSupport implements WebS return Mono.fromCallable( () -> { String[] protocols = beforeHandshake(url, headers, handler); - ClientUpgradeRequest upgradeRequest = createRequest(url, headers, protocols); + ClientUpgradeRequest upgradeRequest = createRequest(headers, protocols); Object jettyHandler = createJettyHandler(url, handler, completionMono); return this.wsClient.connect(jettyHandler, url, upgradeRequest); }) @@ -91,18 +91,20 @@ public class JettyWebSocketClient extends WebSocketClientSupport implements WebS private Object createJettyHandler(URI url, WebSocketHandler handler, MonoProcessor completion) { return new JettyWebSocketHandlerAdapter( - handler, completion, session -> createJettySession(url, session)); + handler, session -> createJettySession(url, completion, session)); } - private JettyWebSocketSession createJettySession(URI url, Session session) { + private JettyWebSocketSession createJettySession(URI url, MonoProcessor completion, + Session session) { + UpgradeResponse response = session.getUpgradeResponse(); HttpHeaders responseHeaders = new HttpHeaders(); response.getHeaders().forEach(responseHeaders::put); HandshakeInfo info = afterHandshake(url, responseHeaders); - return new JettyWebSocketSession(session, info, bufferFactory); + return new JettyWebSocketSession(session, info, this.bufferFactory, completion); } - private ClientUpgradeRequest createRequest(URI url, HttpHeaders headers, String[] protocols) { + private ClientUpgradeRequest createRequest(HttpHeaders headers, String[] protocols) { ClientUpgradeRequest request = new ClientUpgradeRequest(); request.setSubProtocols(protocols); headers.forEach(request::setHeader); diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/StandardWebSocketClient.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/StandardWebSocketClient.java index 53be682970..98ff388630 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/StandardWebSocketClient.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/StandardWebSocketClient.java @@ -105,13 +105,15 @@ public class StandardWebSocketClient extends WebSocketClientSupport implements W private StandardWebSocketHandlerAdapter createEndpoint(URI url, WebSocketHandler handler, MonoProcessor completion, DefaultConfigurator configurator) { - return new StandardWebSocketHandlerAdapter(handler, completion, - session -> createSession(url, configurator.getResponseHeaders(), session)); + return new StandardWebSocketHandlerAdapter(handler, + session -> createSession(url, configurator.getResponseHeaders(), completion, session)); } - private StandardWebSocketSession createSession(URI url, HttpHeaders responseHeaders, Session session) { + private StandardWebSocketSession createSession(URI url, HttpHeaders responseHeaders, + MonoProcessor completion, Session session) { + HandshakeInfo info = afterHandshake(url, responseHeaders); - return new StandardWebSocketSession(session, info, this.bufferFactory); + return new StandardWebSocketSession(session, info, this.bufferFactory, completion); } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/UndertowWebSocketClient.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/UndertowWebSocketClient.java index fe5403e13b..4a2ac6314f 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/UndertowWebSocketClient.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/client/UndertowWebSocketClient.java @@ -169,12 +169,17 @@ public class UndertowWebSocketClient extends WebSocketClientSupport implements W .then(completionMono); } - private void handleWebSocket(URI url, WebSocketHandler handler, MonoProcessor completionMono, + private void handleWebSocket(URI url, WebSocketHandler handler, MonoProcessor completion, DefaultNegotiation negotiation, WebSocketChannel channel) { HandshakeInfo info = afterHandshake(url, negotiation.getResponseHeaders()); - UndertowWebSocketSession session = new UndertowWebSocketSession(channel, info, this.bufferFactory); - new UndertowWebSocketHandlerAdapter(handler, completionMono).handle(session); + UndertowWebSocketSession session = new UndertowWebSocketSession(channel, info, bufferFactory, completion); + UndertowWebSocketHandlerAdapter adapter = new UndertowWebSocketHandlerAdapter(session); + + channel.getReceiveSetter().set(adapter); + channel.resumeReceives(); + + handler.handle(session).subscribe(session); } diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/JettyRequestUpgradeStrategy.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/JettyRequestUpgradeStrategy.java index c265a9c287..5773342478 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/JettyRequestUpgradeStrategy.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/JettyRequestUpgradeStrategy.java @@ -120,7 +120,7 @@ public class JettyRequestUpgradeStrategy implements RequestUpgradeStrategy, Life HttpServletResponse servletResponse = getHttpServletResponse(response); JettyWebSocketHandlerAdapter adapter = new JettyWebSocketHandlerAdapter(handler, - null, session -> { + session -> { HandshakeInfo info = getHandshakeInfo(exchange, subProtocol); DataBufferFactory factory = response.bufferFactory(); return new JettyWebSocketSession(session, info, factory); diff --git a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/UndertowRequestUpgradeStrategy.java b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/UndertowRequestUpgradeStrategy.java index b3aec4ed3c..a0d5e77bb2 100644 --- a/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/UndertowRequestUpgradeStrategy.java +++ b/spring-web-reactive/src/main/java/org/springframework/web/reactive/socket/server/upgrade/UndertowRequestUpgradeStrategy.java @@ -66,8 +66,14 @@ public class UndertowRequestUpgradeStrategy implements RequestUpgradeStrategy { Hybi13Handshake handshake = new Hybi13Handshake(protocols, false); List handshakes = Collections.singletonList(handshake); + URI url = request.getURI(); + HttpHeaders headers = request.getHeaders(); + Mono principal = exchange.getPrincipal(); + HandshakeInfo info = new HandshakeInfo(url, headers, principal, subProtocol); + DataBufferFactory bufferFactory = exchange.getResponse().bufferFactory(); + try { - DefaultCallback callback = new DefaultCallback(exchange, handler, subProtocol); + DefaultCallback callback = new DefaultCallback(info, handler, bufferFactory); new WebSocketProtocolHandshakeHandler(handshakes, callback).handleRequest(httpExchange); } catch (Exception ex) { @@ -80,40 +86,35 @@ public class UndertowRequestUpgradeStrategy implements RequestUpgradeStrategy { private class DefaultCallback implements WebSocketConnectionCallback { - private final ServerWebExchange exchange; + private final HandshakeInfo handshakeInfo; private final WebSocketHandler handler; - private final Optional subProtocol; + private final DataBufferFactory bufferFactory; - public DefaultCallback(ServerWebExchange exchange, WebSocketHandler handler, - Optional subProtocol) { + public DefaultCallback(HandshakeInfo handshakeInfo, WebSocketHandler handler, + DataBufferFactory bufferFactory) { - this.exchange = exchange; + this.handshakeInfo = handshakeInfo; this.handler = handler; - this.subProtocol = subProtocol; + this.bufferFactory = bufferFactory; } @Override public void onConnect(WebSocketHttpExchange httpExchange, WebSocketChannel channel) { - UndertowWebSocketHandlerAdapter adapter = new UndertowWebSocketHandlerAdapter(this.handler); - UndertowWebSocketSession session = createWebSocketSession(channel); - adapter.handle(session); + + UndertowWebSocketSession session = createSession(channel); + UndertowWebSocketHandlerAdapter adapter = new UndertowWebSocketHandlerAdapter(session); + + channel.getReceiveSetter().set(adapter); + channel.resumeReceives(); + + this.handler.handle(session).subscribe(session); } - private UndertowWebSocketSession createWebSocketSession(WebSocketChannel channel) { - HandshakeInfo info = getHandshakeInfo(); - DataBufferFactory bufferFactory = this.exchange.getResponse().bufferFactory(); - return new UndertowWebSocketSession(channel, info, bufferFactory); - } - - private HandshakeInfo getHandshakeInfo() { - ServerHttpRequest request = this.exchange.getRequest(); - URI url = request.getURI(); - HttpHeaders headers = request.getHeaders(); - Mono principal = this.exchange.getPrincipal(); - return new HandshakeInfo(url, headers, principal, this.subProtocol); + private UndertowWebSocketSession createSession(WebSocketChannel channel) { + return new UndertowWebSocketSession(channel, this.handshakeInfo, this.bufferFactory); } }