From 5025c304b81622f2e81882a46ea7150c07d6946b Mon Sep 17 00:00:00 2001 From: Andy Wilkinson Date: Fri, 27 Sep 2013 15:41:46 +0100 Subject: [PATCH] Introduce CONNECT_ACK message type Previously, handling of a STOMP CONNECT message and sending of a CONNECTED response was performed by StompProtocolHandler if it was backed by SimpleBrokerMessageHandler, or left up to the real message broker if it was backed by StompBrokerRelayMessageHandler. This wasn't ideal as it should be StompProtocolHandler's job to simply map messages to and from the STOMP protocol, not to do part of the broker's job and respond directly to CONNECT. This commit introduces a new message type, CONNECT_ACK. When it receives a CONNECT message, SimpleBrokerMessageHandler will now respond with a CONNECT_ACK message that StompProtocolHandler can map into a STOMP CONNECTED message. The CONNECT_ACK message contains the CONNECT message as a header so that StompProtocolHandler has access to its accept-version header. StompProtocolHandler has been simplified so that a CONNECT message is always passed to the output channel, irrespective of whether it's backed by a simple broker or a real broker. The handleConnect flag, and the code that would set it correctly depending on the app's configuration, has been removed. --- .../simp/SimpMessageHeaderAccessor.java | 2 + .../messaging/simp/SimpMessageType.java | 2 + .../config/ServletStompEndpointRegistry.java | 4 +- ...cketMessageBrokerConfigurationSupport.java | 3 +- .../handler/SimpleBrokerMessageHandler.java | 13 ++- .../simp/stomp/StompProtocolHandler.java | 96 ++++++------------- .../ServletStompEndpointRegistryTests.java | 2 +- .../SimpleBrokerMessageHandlerTests.java | 27 ++++++ .../simp/stomp/StompProtocolHandlerTests.java | 23 +++-- 9 files changed, 90 insertions(+), 82 deletions(-) diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageHeaderAccessor.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageHeaderAccessor.java index 7ec3c279e0..40dced0f75 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageHeaderAccessor.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageHeaderAccessor.java @@ -41,6 +41,8 @@ import org.springframework.util.Assert; */ public class SimpMessageHeaderAccessor extends NativeMessageHeaderAccessor { + public static final String CONNECT_MESSAGE_HEADER = "connectMessage"; + public static final String DESTINATION_HEADER = "destination"; public static final String MESSAGE_TYPE_HEADER = "messageType"; diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageType.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageType.java index 3351dd5b7e..b61917714d 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageType.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/SimpMessageType.java @@ -28,6 +28,8 @@ public enum SimpMessageType { CONNECT, + CONNECT_ACK, + MESSAGE, SUBSCRIBE, diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistry.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistry.java index 0f54c67fd8..3e1050c920 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistry.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistry.java @@ -54,8 +54,7 @@ public class ServletStompEndpointRegistry implements StompEndpointRegistry { public ServletStompEndpointRegistry(WebSocketHandler webSocketHandler, - MutableUserQueueSuffixResolver userQueueSuffixResolver, TaskScheduler defaultSockJsTaskScheduler, - boolean handleConnect) { + MutableUserQueueSuffixResolver userQueueSuffixResolver, TaskScheduler defaultSockJsTaskScheduler) { Assert.notNull(webSocketHandler); Assert.notNull(userQueueSuffixResolver); @@ -64,7 +63,6 @@ public class ServletStompEndpointRegistry implements StompEndpointRegistry { this.subProtocolWebSocketHandler = findSubProtocolWebSocketHandler(webSocketHandler); this.stompHandler = new StompProtocolHandler(); this.stompHandler.setUserQueueSuffixResolver(userQueueSuffixResolver); - this.stompHandler.setHandleConnect(handleConnect); this.sockJsScheduler = defaultSockJsTaskScheduler; } diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/config/WebSocketMessageBrokerConfigurationSupport.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/config/WebSocketMessageBrokerConfigurationSupport.java index ba5fb13baa..427f1979a7 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/config/WebSocketMessageBrokerConfigurationSupport.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/config/WebSocketMessageBrokerConfigurationSupport.java @@ -57,9 +57,8 @@ public abstract class WebSocketMessageBrokerConfigurationSupport { @Bean public HandlerMapping brokerWebSocketHandlerMapping() { - boolean brokerRelayConfigured = getMessageBrokerConfigurer().getStompBrokerRelay() != null; ServletStompEndpointRegistry registry = new ServletStompEndpointRegistry(subProtocolWebSocketHandler(), - userQueueSuffixResolver(), brokerDefaultSockJsTaskScheduler(), !brokerRelayConfigured); + userQueueSuffixResolver(), brokerDefaultSockJsTaskScheduler()); registerStompEndpoints(registry); AbstractHandlerMapping hm = registry.getHandlerMapping(); diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandler.java index 2e7638429d..a117f3ceec 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandler.java @@ -33,6 +33,8 @@ import org.springframework.util.MultiValueMap; */ public class SimpleBrokerMessageHandler extends AbstractBrokerMessageHandler { + private static final byte[] EMPTY_PAYLOAD = new byte[0]; + private final MessageChannel messageChannel; private SubscriptionRegistry subscriptionRegistry = new DefaultSubscriptionRegistry(); @@ -96,8 +98,17 @@ public class SimpleBrokerMessageHandler extends AbstractBrokerMessageHandler { sendMessageToSubscribers(headers.getDestination(), message); } else if (SimpMessageType.DISCONNECT.equals(messageType)) { - String sessionId = SimpMessageHeaderAccessor.wrap(message).getSessionId(); + String sessionId = headers.getSessionId(); this.subscriptionRegistry.unregisterAllSubscriptions(sessionId); + } else if (SimpMessageType.CONNECT.equals(messageType)) { + String sessionId = headers.getSessionId(); + SimpMessageHeaderAccessor connectAckHeaders = + SimpMessageHeaderAccessor.create(SimpMessageType.CONNECT_ACK); + connectAckHeaders.setSessionId(sessionId); + connectAckHeaders.setHeader(SimpMessageHeaderAccessor.CONNECT_MESSAGE_HEADER, message); + Message connectAck = + MessageBuilder.withPayloadAndHeaders(EMPTY_PAYLOAD, connectAckHeaders).build(); + this.messageChannel.send(connectAck); } } diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/StompProtocolHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/StompProtocolHandler.java index 6bbd655e91..8fdfd0da0f 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/StompProtocolHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/StompProtocolHandler.java @@ -70,8 +70,6 @@ public class StompProtocolHandler implements SubProtocolHandler { private MutableUserQueueSuffixResolver queueSuffixResolver = new SimpleUserQueueSuffixResolver(); - private volatile boolean handleConnect = false; - /** * Configure a resolver to use to maintain queue suffixes for user * @see {@link org.springframework.messaging.simp.handler.UserDestinationMessageHandler} @@ -87,29 +85,6 @@ public class StompProtocolHandler implements SubProtocolHandler { return this.queueSuffixResolver; } - /** - * Configures the handling of CONNECT frames. When {@code true}, CONNECT - * frames will be handled by this handler, and a CONNECTED response will be - * sent. When {@code false}, CONNECT frames will be forwarded for - * handling by another component. - * - * @param handleConnect {@code true} if connect frames should be handled - * by this handler, {@code false} otherwise. - */ - public void setHandleConnect(boolean handleConnect) { - this.handleConnect = handleConnect; - } - - /** - * Returns whether or not this handler will handle CONNECT frames. - * - * @return Returns {@code true} if this handler will handle CONNECT frames, - * otherwise {@code false}. - */ - public boolean willHandleConnect() { - return this.handleConnect; - } - @Override public List getSupportedProtocols() { return Arrays.asList("v10.stomp", "v11.stomp", "v12.stomp"); @@ -144,13 +119,7 @@ public class StompProtocolHandler implements SubProtocolHandler { headers.setUser(session.getPrincipal()); message = MessageBuilder.withPayloadAndHeaders(message.getPayload(), headers).build(); - - if (this.handleConnect && SimpMessageType.CONNECT.equals(headers.getMessageType())) { - handleConnect(session, message); - } - else { - outputChannel.send(message); - } + outputChannel.send(message); } catch (Throwable t) { logger.error("Terminating STOMP session due to failure to send message: ", t); @@ -170,13 +139,15 @@ public class StompProtocolHandler implements SubProtocolHandler { headers.setCommandIfNotSet(StompCommand.MESSAGE); } + if (headers.getMessageType() == SimpMessageType.CONNECT_ACK) { + StompHeaderAccessor connectedHeaders = StompHeaderAccessor.create(StompCommand.CONNECTED); + connectedHeaders.setVersion(getVersion(headers)); + connectedHeaders.setHeartbeat(0, 0); + headers = connectedHeaders; + } + if (headers.getCommand() == StompCommand.CONNECTED) { - if (this.handleConnect) { - // Ignore since we already sent it - return; - } else { - augmentConnectedHeaders(headers, session); - } + augmentConnectedHeaders(headers, session); } if (StompCommand.MESSAGE.equals(headers.getCommand()) && (headers.getSubscriptionId() == null)) { @@ -208,35 +179,6 @@ public class StompProtocolHandler implements SubProtocolHandler { } } - protected void handleConnect(WebSocketSession session, Message message) throws IOException { - - StompHeaderAccessor connectHeaders = StompHeaderAccessor.wrap(message); - StompHeaderAccessor connectedHeaders = StompHeaderAccessor.create(StompCommand.CONNECTED); - - Set acceptVersions = connectHeaders.getAcceptVersion(); - if (acceptVersions.contains("1.2")) { - connectedHeaders.setVersion("1.2"); - } - else if (acceptVersions.contains("1.1")) { - connectedHeaders.setVersion("1.1"); - } - else if (acceptVersions.isEmpty()) { - // 1.0 - } - else { - throw new StompConversionException("Unsupported version '" + acceptVersions + "'"); - } - connectedHeaders.setHeartbeat(0,0); - - augmentConnectedHeaders(connectedHeaders, session); - - // TODO: security - - Message connectedMessage = MessageBuilder.withPayloadAndHeaders(new byte[0], connectedHeaders).build(); - String payload = new String(this.stompEncoder.encode(connectedMessage), Charset.forName("UTF-8")); - session.sendMessage(new TextMessage(payload)); - } - private void augmentConnectedHeaders(StompHeaderAccessor headers, WebSocketSession session) { Principal principal = session.getPrincipal(); if (principal != null) { @@ -287,4 +229,24 @@ public class StompProtocolHandler implements SubProtocolHandler { outputChannel.send(message); } + private String getVersion(StompHeaderAccessor connectAckHeaders) { + Message connectMessage = + (Message) connectAckHeaders.getHeader(StompHeaderAccessor.CONNECT_MESSAGE_HEADER); + StompHeaderAccessor connectHeaders = StompHeaderAccessor.wrap(connectMessage); + + Set acceptVersions = connectHeaders.getAcceptVersion(); + if (acceptVersions.contains("1.2")) { + return "1.2"; + } + else if (acceptVersions.contains("1.1")) { + return "1.1"; + } + else if (acceptVersions.isEmpty()) { + return null; + } + else { + throw new StompConversionException("Unsupported version '" + acceptVersions + "'"); + } + } + } diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistryTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistryTests.java index f4c190ba39..7531e8ed20 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistryTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/config/ServletStompEndpointRegistryTests.java @@ -53,7 +53,7 @@ public class ServletStompEndpointRegistryTests { this.webSocketHandler = new SubProtocolWebSocketHandler(channel); this.queueSuffixResolver = new SimpleUserQueueSuffixResolver(); TaskScheduler taskScheduler = Mockito.mock(TaskScheduler.class); - this.registry = new ServletStompEndpointRegistry(webSocketHandler, queueSuffixResolver, taskScheduler, false); + this.registry = new ServletStompEndpointRegistry(webSocketHandler, queueSuffixResolver, taskScheduler); } diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandlerTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandlerTests.java index b7d72c67d7..f61bf7516b 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandlerTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/handler/SimpleBrokerMessageHandlerTests.java @@ -30,6 +30,8 @@ import org.springframework.messaging.simp.SimpMessageHeaderAccessor; import org.springframework.messaging.simp.SimpMessageType; import org.springframework.messaging.support.MessageBuilder; +import static org.junit.Assert.*; + import static org.mockito.Mockito.*; @@ -111,6 +113,24 @@ public class SimpleBrokerMessageHandlerTests { assertCapturedMessage(sess2, "sub3", "/bar"); } + @Test + public void connect() { + + String sess1 = "sess1"; + + this.messageHandler.start(); + + Message connectMessage = createConnectMessage(sess1); + this.messageHandler.handleMessage(connectMessage); + + verify(this.clientChannel, times(1)).send(this.messageCaptor.capture()); + Message connectAckMessage = this.messageCaptor.getValue(); + + SimpMessageHeaderAccessor connectAckHeaders = SimpMessageHeaderAccessor.wrap(connectAckMessage); + assertEquals(connectMessage, connectAckHeaders.getHeader(SimpMessageHeaderAccessor.CONNECT_MESSAGE_HEADER)); + assertEquals(sess1, connectAckHeaders.getSessionId()); + } + protected Message createSubscriptionMessage(String sessionId, String subcriptionId, String destination) { @@ -122,6 +142,13 @@ public class SimpleBrokerMessageHandlerTests { return MessageBuilder.withPayload("").copyHeaders(headers.toMap()).build(); } + protected Message createConnectMessage(String sessionId) { + SimpMessageHeaderAccessor headers = SimpMessageHeaderAccessor.create(SimpMessageType.CONNECT); + headers.setSessionId(sessionId); + + return MessageBuilder.withPayloadAndHeaders("", headers).build(); + } + protected Message createMessage(String destination, String payload) { SimpMessageHeaderAccessor headers = SimpMessageHeaderAccessor.create(SimpMessageType.MESSAGE); diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/StompProtocolHandlerTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/StompProtocolHandlerTests.java index 78311b36d8..c44914b1df 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/StompProtocolHandlerTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/StompProtocolHandlerTests.java @@ -26,6 +26,9 @@ import org.mockito.ArgumentCaptor; import org.mockito.Mockito; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; +import org.springframework.messaging.simp.SimpMessageHeaderAccessor; +import org.springframework.messaging.simp.SimpMessageType; +import org.springframework.messaging.support.MessageBuilder; import org.springframework.web.socket.TextMessage; import org.springframework.web.socket.support.TestPrincipal; import org.springframework.web.socket.support.TestWebSocketSession; @@ -61,20 +64,25 @@ public class StompProtocolHandlerTests { } @Test - public void connectedResponseIsSentWhenHandlingConnect() { - this.stompHandler.setHandleConnect(true); + public void connectedResponseIsSentWhenConnectAckIsToBeSentToClient() { + StompHeaderAccessor connectHeaders = StompHeaderAccessor.create(StompCommand.CONNECT); + connectHeaders.setHeartbeat(10000, 10000); + connectHeaders.setNativeHeader(StompHeaderAccessor.STOMP_ACCEPT_VERSION_HEADER, "1.0,1.1"); - TextMessage textMessage = StompTextMessageBuilder.create(StompCommand.CONNECT).headers( - "login:guest", "passcode:guest", "accept-version:1.1,1.0", "heart-beat:10000,10000").build(); + Message connectMessage = MessageBuilder.withPayloadAndHeaders(new byte[0], connectHeaders).build(); - this.stompHandler.handleMessageFromClient(this.session, textMessage, this.channel); + SimpMessageHeaderAccessor connectAckHeaders = SimpMessageHeaderAccessor.create(SimpMessageType.CONNECT_ACK); + connectAckHeaders.setHeader(SimpMessageHeaderAccessor.CONNECT_MESSAGE_HEADER, connectMessage); + + Message connectAck = MessageBuilder.withPayloadAndHeaders(new byte[0], connectAckHeaders).build(); + this.stompHandler.handleMessageToClient(this.session, connectAck); verifyNoMoreInteractions(this.channel); // Check CONNECTED reply assertEquals(1, this.session.getSentMessages().size()); - textMessage = (TextMessage) this.session.getSentMessages().get(0); + TextMessage textMessage = (TextMessage) this.session.getSentMessages().get(0); Message message = new StompDecoder().decode(ByteBuffer.wrap(textMessage.getPayload().getBytes())); StompHeaderAccessor replyHeaders = StompHeaderAccessor.wrap(message); @@ -86,8 +94,7 @@ public class StompProtocolHandlerTests { } @Test - public void connectIsForwardedWhenNotHandlingConnect() { - this.stompHandler.setHandleConnect(false); + public void messagesAreAugmentedAndForwarded() { TextMessage textMessage = StompTextMessageBuilder.create(StompCommand.CONNECT).headers( "login:guest", "passcode:guest", "accept-version:1.1,1.0", "heart-beat:10000,10000").build();