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