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); } }