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