Refactor StompRelayPubSubMessageHandler

This commit is contained in:
Rossen Stoyanchev
2013-06-17 15:14:26 +01:00
parent 3dabe21563
commit 426875b4f6
6 changed files with 206 additions and 114 deletions

View File

@@ -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.
*/

View File

@@ -91,7 +91,7 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler<Mes
for (String pattern : this.disallowedDestinations) {
if (this.pathMatcher.match(pattern, destination)) {
if (logger.isTraceEnabled()) {
logger.trace("Skip notification message id=" + message.getHeaders().getId());
logger.trace("Skip message id=" + message.getHeaders().getId());
}
return false;
}
@@ -105,7 +105,7 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler<Mes
}
}
if (logger.isTraceEnabled()) {
logger.trace("Skip notification message id=" + message.getHeaders().getId());
logger.trace("Skip message id=" + message.getHeaders().getId());
}
return false;
}
@@ -123,10 +123,6 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler<Mes
return;
}
if (logger.isTraceEnabled()) {
logger.trace("Handling message id=" + message.getHeaders().getId());
}
if (MessageType.MESSAGE.equals(messageType)) {
handlePublish(message);
}

View File

@@ -55,6 +55,7 @@ public enum StompCommand {
commandToMessageType.put(StompCommand.CONNECT, MessageType.CONNECT);
commandToMessageType.put(StompCommand.STOMP, MessageType.CONNECT);
commandToMessageType.put(StompCommand.SEND, MessageType.MESSAGE);
commandToMessageType.put(StompCommand.MESSAGE, MessageType.MESSAGE);
commandToMessageType.put(StompCommand.SUBSCRIBE, MessageType.SUBSCRIBE);
commandToMessageType.put(StompCommand.UNSUBSCRIBE, MessageType.UNSUBSCRIBE);
commandToMessageType.put(StompCommand.DISCONNECT, MessageType.DISCONNECT);

View File

@@ -17,16 +17,21 @@
package org.springframework.web.messaging.stomp.support;
import java.nio.charset.Charset;
import java.util.ArrayList;
import java.util.Collection;
import java.util.List;
import java.util.Map;
import java.util.concurrent.BlockingQueue;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.LinkedBlockingQueue;
import java.util.concurrent.atomic.AtomicBoolean;
import org.springframework.http.MediaType;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.support.MessageBuilder;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
import org.springframework.web.messaging.MessageType;
import org.springframework.web.messaging.PubSubChannelRegistry;
import org.springframework.web.messaging.PubSubHeaders;
@@ -60,8 +65,7 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler
private final TcpClient<String, String> tcpClient;
private final Map<String, TcpConnection<String, String>> connections =
new ConcurrentHashMap<String, TcpConnection<String, String>>();
private final Map<String, RelaySession> relaySessions = new ConcurrentHashMap<String, RelaySession>();
/**
@@ -92,99 +96,15 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler
}
@Override
public void handleConnect(final Message<?> message) {
final String sessionId = (String) message.getHeaders().get(PubSubHeaders.SESSION_ID);
Promise<TcpConnection<String, String>> promise = this.tcpClient.open();
promise.onSuccess(new Consumer<TcpConnection<String,String>>() {
@Override
public void accept(TcpConnection<String, String> connection) {
connections.put(sessionId, connection);
forwardMessage(message, StompCommand.CONNECT);
}
});
promise.consume(new Consumer<TcpConnection<String,String>>() {
@Override
public void accept(TcpConnection<String, String> connection) {
connection.in().consume(new Consumer<String>() {
@Override
public void accept(String stompFrame) {
if (stompFrame.isEmpty()) {
// TODO: why are we getting empty frames?
return;
}
Message<byte[]> message = stompMessageConverter.toMessage(stompFrame, sessionId);
clientChannel.send(message);
}
});
}
});
// TODO: ATM no way to detect closed socket
// StompHeaders stompHeaders = StompHeaders.create(StompCommand.ERROR);
// stompHeaders.setMessage("Socket closed, STOMP session=" + sessionId);
// stompHeaders.setSessionId(sessionId);
// Message<byte[]> errorMessage = new GenericMessage<byte[]>(new byte[0], stompHeaders.toMessageHeaders());
// getClientChannel().send(errorMessage);
}
private void forwardMessage(Message<?> message, StompCommand command) {
StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders());
String sessionId = headers.getSessionId();
byte[] bytesToWrite;
try {
headers.setStompCommandIfNotSet(StompCommand.SEND);
MediaType contentType = headers.getContentType();
byte[] payload = this.payloadConverter.convertToPayload(message.getPayload(), contentType);
Message<byte[]> byteMessage = MessageBuilder.fromPayloadAndHeaders(payload, headers.toMessageHeaders()).build();
bytesToWrite = this.stompMessageConverter.fromMessage(byteMessage);
}
catch (Throwable ex) {
logger.error("Failed to forward message " + message, ex);
public void handleConnect(Message<?> message) {
StompHeaders stompHeaders = StompHeaders.fromMessageHeaders(message.getHeaders());
String sessionId = stompHeaders.getSessionId();
if (sessionId == null) {
logger.error("No sessionId in message " + message);
return;
}
TcpConnection<String, String> 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<String, String> getConnection(String sessionId) {
TcpConnection<String, String> 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<TcpConnection<String, String>> promise;
private final AtomicBoolean isConnected = new AtomicBoolean(false);
private final BlockingQueue<Message<?>> messageQueue = new LinkedBlockingQueue<Message<?>>(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<TcpConnection<String,String>>() {
@Override
public void accept(TcpConnection<String, String> connection) {
connection.in().consume(new Consumer<String>() {
@Override
public void accept(String stompFrame) {
readStompFrame(stompFrame);
}
});
stompHeaders.setHeartbeat(0, 0); // TODO
forwardInternal(message, stompHeaders, connection);
}
});
this.promise.onError(new Consumer<Throwable>() {
@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<byte[]> 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<byte[]> 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<String, String> connection = this.promise.get();
if (this.messageQueue.isEmpty()) {
forwardInternal(message, headers, connection);
}
else {
this.messageQueue.add(message);
flushMessages(connection);
}
}
private void flushMessages(TcpConnection<String, String> connection) {
List<Message<?>> messages = new ArrayList<Message<?>>();
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<String, String> connection) {
try {
headers.setStompCommandIfNotSet(StompCommand.SEND);
MediaType contentType = headers.getContentType();
byte[] payload = payloadConverter.convertToPayload(message.getPayload(), contentType);
Message<byte[]> 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;
}
}
}

View File

@@ -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.

View File

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