diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubHeaders.java b/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubHeaders.java index 80c12c3994..dfbfda9b32 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubHeaders.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubHeaders.java @@ -126,6 +126,13 @@ public class PubSubHeaders { return new PubSubHeaders(MessageType.MESSAGE, null, null); } + /** + * Create {@link PubSubHeaders} for a new {@link Message} of a specific type. + */ + public static PubSubHeaders create(MessageType messageType) { + return new PubSubHeaders(messageType, null, null); + } + /** * Create {@link PubSubHeaders} from existing message headers. */ diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java index 4e8a0eb761..a6565b1e1f 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java @@ -91,7 +91,7 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler message) { + StompHeaders stompHeaders = StompHeaders.fromMessageHeaders(message.getHeaders()); + String sessionId = stompHeaders.getSessionId(); + if (sessionId == null) { + logger.error("No sessionId in message " + message); return; } - - TcpConnection connection = getConnection(sessionId); - Assert.notNull(connection, "TCP connection to message broker not found, sessionId=" + sessionId); - try { - if (logger.isTraceEnabled()) { - logger.trace("Forwarding STOMP " + headers.getStompCommand() + " message"); - } - connection.out().accept(new String(bytesToWrite, Charset.forName("UTF-8"))); - } - catch (Throwable ex) { - logger.error("Could not get TCP connection " + sessionId, ex); - try { - if (connection != null) { - connection.close(); - } - } - catch (Throwable t) { - // ignore - } - } - } - - private TcpConnection getConnection(String sessionId) { - TcpConnection connection = this.connections.get(sessionId); - if (connection == null) { - try { - Thread.sleep(1000); - } - catch (InterruptedException e) { - return null; - } - } - connection = this.connections.get(sessionId); - return connection; + RelaySession relaySession = new RelaySession(message, stompHeaders); + this.relaySessions.put(sessionId, relaySession); } @Override @@ -204,7 +124,15 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler @Override public void handleDisconnect(Message message) { - forwardMessage(message, StompCommand.DISCONNECT); + StompHeaders stompHeaders = StompHeaders.fromMessageHeaders(message.getHeaders()); + if (stompHeaders.getStompCommand() != null) { + forwardMessage(message, StompCommand.DISCONNECT); + } + String sessionId = stompHeaders.getSessionId(); + if (sessionId == null) { + logger.error("No sessionId in message " + message); + return; + } } @Override @@ -214,15 +142,164 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler forwardMessage(message, command); } - // TODO: + private void forwardMessage(Message message, StompCommand command) { -/* @Override - public void handleClientConnectionClosed(String sessionId) { - if (logger.isDebugEnabled()) { - logger.debug("Client connection closed for STOMP session=" + sessionId + ". Clearing relay session."); + StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders()); + headers.setStompCommandIfNotSet(command); + + String sessionId = headers.getSessionId(); + if (sessionId == null) { + logger.error("No sessionId in message " + message); + return; } - clearRelaySession(sessionId); - } -*/ + RelaySession session = this.relaySessions.get(sessionId); + if (session == null) { + // TODO: default (non-user) session for sending messages? + logger.warn("No relay session for " + sessionId + ". Message '" + message + "' cannot be forwarded"); + return; + } + + session.forward(message, headers); + } + + + private final class RelaySession { + + private final String sessionId; + + private final Promise> promise; + + private final AtomicBoolean isConnected = new AtomicBoolean(false); + + private final BlockingQueue> messageQueue = new LinkedBlockingQueue>(50); + + + public RelaySession(final Message message, final StompHeaders stompHeaders) { + + Assert.notNull(message, "message is required"); + Assert.notNull(stompHeaders, "stompHeaders is required"); + + this.sessionId = stompHeaders.getSessionId(); + this.promise = tcpClient.open(); + + this.promise.consume(new Consumer>() { + @Override + public void accept(TcpConnection connection) { + connection.in().consume(new Consumer() { + @Override + public void accept(String stompFrame) { + readStompFrame(stompFrame); + } + }); + stompHeaders.setHeartbeat(0, 0); // TODO + forwardInternal(message, stompHeaders, connection); + } + }); + + this.promise.onError(new Consumer() { + @Override + public void accept(Throwable ex) { + relaySessions.remove(sessionId); + logger.error("Failed to connect to broker", ex); + sendError(sessionId, "Failed to connect to message broker " + ex.toString()); + } + }); + + // TODO: ATM no way to detect closed socket + } + + private void readStompFrame(String stompFrame) { + + if (StringUtils.isEmpty(stompFrame)) { + // heartbeat? + return; + } + + Message message = stompMessageConverter.toMessage(stompFrame, this.sessionId); + if (logger.isTraceEnabled()) { + logger.trace("Reading message " + message); + } + + StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders()); + if (StompCommand.CONNECTED == headers.getStompCommand()) { + this.isConnected.set(true); + flushMessages(promise.get()); + return; + } + if (StompCommand.ERROR == headers.getStompCommand()) { + if (logger.isDebugEnabled()) { + logger.warn("STOMP ERROR: " + headers.getMessage() + ". Removing session: " + this.sessionId); + } + relaySessions.remove(this.sessionId); + } + clientChannel.send(message); + } + + private void sendError(String sessionId, String errorText) { + StompHeaders stompHeaders = StompHeaders.create(StompCommand.ERROR); + stompHeaders.setSessionId(sessionId); + stompHeaders.setMessage(errorText); + Message errorMessage = MessageBuilder.fromPayloadAndHeaders( + new byte[0], stompHeaders.toMessageHeaders()).build(); + clientChannel.send(errorMessage); + } + + public void forward(Message message, StompHeaders headers) { + + if (!this.isConnected.get()) { + message = MessageBuilder.fromPayloadAndHeaders(message.getPayload(), headers.toMessageHeaders()).build(); + if (logger.isTraceEnabled()) { + logger.trace("Adding to queue message " + message + ", queue size=" + this.messageQueue.size()); + } + this.messageQueue.add(message); + return; + } + + TcpConnection connection = this.promise.get(); + + if (this.messageQueue.isEmpty()) { + forwardInternal(message, headers, connection); + } + else { + this.messageQueue.add(message); + flushMessages(connection); + } + } + + private void flushMessages(TcpConnection connection) { + List> messages = new ArrayList>(); + this.messageQueue.drainTo(messages); + for (Message message : messages) { + StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders()); + if (!forwardInternal(message, headers, connection)) { + return; + } + } + } + + private boolean forwardInternal(Message message, StompHeaders headers, TcpConnection connection) { + try { + headers.setStompCommandIfNotSet(StompCommand.SEND); + + MediaType contentType = headers.getContentType(); + byte[] payload = payloadConverter.convertToPayload(message.getPayload(), contentType); + Message byteMessage = MessageBuilder.fromPayloadAndHeaders(payload, headers.toMessageHeaders()).build(); + + if (logger.isTraceEnabled()) { + logger.trace("Forwarding message " + byteMessage); + } + + byte[] bytesToWrite = stompMessageConverter.fromMessage(byteMessage); + connection.send(new String(bytesToWrite, Charset.forName("UTF-8"))); + } + catch (Throwable ex) { + logger.error("Failed to forward message " + message, ex); + connection.close(); + sendError(this.sessionId, "Failed to forward message " + message + ": " + ex.getMessage()); + return false; + } + return true; + } + } } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java index 2cd76d0f20..002384a74c 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java @@ -31,6 +31,7 @@ import org.springframework.messaging.MessageHandler; import org.springframework.messaging.support.MessageBuilder; import org.springframework.web.messaging.MessageType; import org.springframework.web.messaging.PubSubChannelRegistry; +import org.springframework.web.messaging.PubSubHeaders; import org.springframework.web.messaging.converter.CompositeMessageConverter; import org.springframework.web.messaging.converter.MessageConverter; import org.springframework.web.messaging.stomp.StompCommand; @@ -198,13 +199,15 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement } } - /* + @SuppressWarnings("unchecked") @Override public void afterConnectionClosed(WebSocketSession session, CloseStatus status) throws Exception { this.sessions.remove(session.getId()); - eventBus.send(AbstractMessageService.CLIENT_CONNECTION_CLOSED_KEY, session.getId()); + PubSubHeaders headers = PubSubHeaders.create(MessageType.DISCONNECT); + headers.setSessionId(session.getId()); + Message message = MessageBuilder.fromPayloadAndHeaders(new byte[0], headers.toMessageHeaders()).build(); + this.outputChannel.send(message); } - */ /** * Handle STOMP messages going back out to WebSocket clients. diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorPubSubChannelRegistry.java b/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorPubSubChannelRegistry.java index cf67eae796..688d0d44f5 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorPubSubChannelRegistry.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorPubSubChannelRegistry.java @@ -32,9 +32,17 @@ public class ReactorPubSubChannelRegistry extends AbstractPubSubChannelRegistry Assert.notNull(reactor, "reactor is required"); - setClientInputChannel(new ReactorMessageChannel(reactor)); - setClientOutputChannel(new ReactorMessageChannel(reactor)); - setMessageBrokerChannel(new ReactorMessageChannel(reactor)); + ReactorMessageChannel channel = new ReactorMessageChannel(reactor); + channel.setName("clientInputChannel"); + setClientInputChannel(channel); + + channel = new ReactorMessageChannel(reactor); + channel.setName("clientOutputChannel"); + setClientOutputChannel(channel); + + channel = new ReactorMessageChannel(reactor); + channel.setName("messageBrokerChannel"); + setMessageBrokerChannel(channel); } }