diff --git a/build.gradle b/build.gradle index a900f37e60..39225d98f7 100644 --- a/build.gradle +++ b/build.gradle @@ -586,8 +586,9 @@ project('spring-integration-websocket') { compile ("org.springframework:spring-webmvc:$springVersion", optional) + testCompile project(":spring-integration-event") testCompile "org.apache.tomcat.embed:tomcat-embed-websocket:$tomcatVersion" - testCompile("org.apache.tomcat.embed:tomcat-embed-logging-juli:${tomcatVersion}") + testCompile("org.apache.tomcat.embed:tomcat-embed-logging-juli:$tomcatVersion") } } diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/event/ReceiptEvent.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/event/ReceiptEvent.java new file mode 100644 index 0000000000..584d95d1e8 --- /dev/null +++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/event/ReceiptEvent.java @@ -0,0 +1,36 @@ +/* + * Copyright 2015 the original author or authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.integration.websocket.event; + +import org.springframework.messaging.Message; +import org.springframework.web.socket.messaging.AbstractSubProtocolEvent; + +/** + * The {@link AbstractSubProtocolEvent} implementation, which is emitted + * for the WebSocket sub-protocol-specific {@code RECEIPT} frame on the client side. + * + * @author Artem Bilan + * @since 4.1.3 + */ +@SuppressWarnings("serial") +public class ReceiptEvent extends AbstractSubProtocolEvent { + + public ReceiptEvent(Object source, Message message) { + super(source, message); + } + +} diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapter.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapter.java index 1782a367e5..2478119372 100644 --- a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapter.java +++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapter.java @@ -1,5 +1,5 @@ /* - * Copyright 2014 the original author or authors. + * Copyright 2014-2015 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -23,6 +23,8 @@ import java.util.ListIterator; import java.util.Map; import java.util.concurrent.atomic.AtomicReference; +import org.springframework.context.ApplicationEventPublisher; +import org.springframework.context.ApplicationEventPublisherAware; import org.springframework.context.Lifecycle; import org.springframework.integration.channel.FixedSubscriberChannel; import org.springframework.integration.endpoint.MessageProducerSupport; @@ -30,6 +32,7 @@ import org.springframework.integration.support.json.JacksonJsonUtils; import org.springframework.integration.websocket.IntegrationWebSocketContainer; import org.springframework.integration.websocket.ServerWebSocketContainer; import org.springframework.integration.websocket.WebSocketListener; +import org.springframework.integration.websocket.event.ReceiptEvent; import org.springframework.integration.websocket.support.PassThruSubProtocolHandler; import org.springframework.integration.websocket.support.SubProtocolHandlerRegistry; import org.springframework.messaging.Message; @@ -48,6 +51,7 @@ import org.springframework.messaging.simp.SimpMessageType; import org.springframework.messaging.simp.broker.AbstractBrokerMessageHandler; import org.springframework.messaging.simp.broker.SimpleBrokerMessageHandler; import org.springframework.messaging.simp.stomp.StompBrokerRelayMessageHandler; +import org.springframework.messaging.simp.stomp.StompCommand; import org.springframework.messaging.support.MessageBuilder; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; @@ -55,17 +59,21 @@ import org.springframework.util.MimeTypeUtils; import org.springframework.web.socket.CloseStatus; import org.springframework.web.socket.WebSocketMessage; import org.springframework.web.socket.WebSocketSession; +import org.springframework.web.socket.messaging.SessionConnectedEvent; /** * @author Artem Bilan * @since 4.1 */ -public class WebSocketInboundChannelAdapter extends MessageProducerSupport implements WebSocketListener { +public class WebSocketInboundChannelAdapter extends MessageProducerSupport + implements WebSocketListener, ApplicationEventPublisherAware { private static final byte[] EMPTY_PAYLOAD = new byte[0]; private final List defaultConverters = new ArrayList(3); + private ApplicationEventPublisher eventPublisher; + { this.defaultConverters.add(new StringMessageConverter()); this.defaultConverters.add(new ByteArrayMessageConverter()); @@ -170,6 +178,11 @@ public class WebSocketInboundChannelAdapter extends MessageProducerSupport imple this.useBroker = useBroker; } + @Override + public void setApplicationEventPublisher(ApplicationEventPublisher applicationEventPublisher) { + this.eventPublisher = applicationEventPublisher; + } + @Override protected void onInit() { super.onInit(); @@ -258,11 +271,15 @@ public class WebSocketInboundChannelAdapter extends MessageProducerSupport imple return this.active; } + @SuppressWarnings("unchecked") private void handleMessageAndSend(Message message) throws Exception { SimpMessageHeaderAccessor headerAccessor = SimpMessageHeaderAccessor.wrap(message); + StompCommand stompCommand = (StompCommand) headerAccessor.getHeader("stompCommand"); SimpMessageType messageType = headerAccessor.getMessageType(); if ((messageType == null || SimpMessageType.MESSAGE.equals(messageType) - || (SimpMessageType.CONNECT.equals(messageType) && !this.useBroker)) + || (SimpMessageType.CONNECT.equals(messageType) && !this.useBroker) + || StompCommand.CONNECTED.equals(stompCommand) + || StompCommand.RECEIPT.equals(stompCommand)) && !checkDestinationPrefix(headerAccessor.getDestination())) { if (SimpMessageType.CONNECT.equals(messageType)) { String sessionId = headerAccessor.getSessionId(); @@ -273,6 +290,12 @@ public class WebSocketInboundChannelAdapter extends MessageProducerSupport imple WebSocketSession session = this.webSocketContainer.getSession(sessionId); this.subProtocolHandlerRegistry.findProtocolHandler(session).handleMessageToClient(session, ackMessage); } + else if (StompCommand.CONNECTED.equals(stompCommand)) { + this.eventPublisher.publishEvent(new SessionConnectedEvent(this, (Message) message)); + } + else if (StompCommand.RECEIPT.equals(stompCommand)) { + this.eventPublisher.publishEvent(new ReceiptEvent(this, (Message) message)); + } else { headerAccessor.removeHeader(SimpMessageHeaderAccessor.NATIVE_HEADERS); Object payload = this.messageConverter.fromMessage(message, this.payloadType.get()); diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/StompIntegrationTests.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/StompIntegrationTests.java index 16edbd1638..f74de2c4e3 100644 --- a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/StompIntegrationTests.java +++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/StompIntegrationTests.java @@ -16,8 +16,10 @@ package org.springframework.integration.websocket.client; +import static org.hamcrest.Matchers.instanceOf; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertThat; import static org.junit.Assert.assertTrue; import java.lang.annotation.ElementType; @@ -36,6 +38,8 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.beans.factory.annotation.Qualifier; import org.springframework.beans.factory.annotation.Value; import org.springframework.context.ApplicationContext; +import org.springframework.context.ApplicationEvent; +import org.springframework.context.ApplicationListener; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.ComponentScan; import org.springframework.context.annotation.Configuration; @@ -50,19 +54,23 @@ import org.springframework.integration.channel.DirectChannel; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.config.EnableIntegration; import org.springframework.integration.core.MessageProducer; +import org.springframework.integration.event.inbound.ApplicationEventListeningMessageProducer; import org.springframework.integration.test.util.TestUtils; import org.springframework.integration.transformer.ExpressionEvaluatingTransformer; import org.springframework.integration.websocket.ClientWebSocketContainer; import org.springframework.integration.websocket.IntegrationWebSocketContainer; import org.springframework.integration.websocket.TomcatWebSocketTestServer; +import org.springframework.integration.websocket.event.ReceiptEvent; import org.springframework.integration.websocket.inbound.WebSocketInboundChannelAdapter; import org.springframework.integration.websocket.outbound.WebSocketOutboundMessageHandler; import org.springframework.integration.websocket.support.SubProtocolHandlerRegistry; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; import org.springframework.messaging.MessageHandler; +import org.springframework.messaging.PollableChannel; import org.springframework.messaging.handler.annotation.MessageExceptionHandler; import org.springframework.messaging.handler.annotation.MessageMapping; +import org.springframework.messaging.simp.SimpMessagingTemplate; import org.springframework.messaging.simp.annotation.SendToUser; import org.springframework.messaging.simp.annotation.SubscribeMapping; import org.springframework.messaging.simp.broker.SimpleBrokerMessageHandler; @@ -70,6 +78,7 @@ import org.springframework.messaging.simp.broker.SubscriptionRegistry; import org.springframework.messaging.simp.config.MessageBrokerRegistry; import org.springframework.messaging.simp.stomp.StompCommand; import org.springframework.messaging.simp.stomp.StompHeaderAccessor; +import org.springframework.messaging.support.AbstractSubscribableChannel; import org.springframework.messaging.support.MessageBuilder; import org.springframework.stereotype.Controller; import org.springframework.test.annotation.DirtiesContext; @@ -79,6 +88,9 @@ import org.springframework.web.socket.client.standard.StandardWebSocketClient; import org.springframework.web.socket.config.annotation.AbstractWebSocketMessageBrokerConfigurer; import org.springframework.web.socket.config.annotation.EnableWebSocketMessageBroker; import org.springframework.web.socket.config.annotation.StompEndpointRegistry; +import org.springframework.web.socket.messaging.AbstractSubProtocolEvent; +import org.springframework.web.socket.messaging.SessionConnectedEvent; +import org.springframework.web.socket.messaging.SessionSubscribeEvent; import org.springframework.web.socket.messaging.StompSubProtocolHandler; import org.springframework.web.socket.messaging.SubProtocolHandler; import org.springframework.web.socket.server.standard.TomcatRequestUpgradeStrategy; @@ -93,6 +105,8 @@ import org.springframework.web.socket.server.support.DefaultHandshakeHandler; @DirtiesContext public class StompIntegrationTests { + private static final SpelExpressionParser PARSER = new SpelExpressionParser(); + @Value("#{server.serverContext}") private ApplicationContext serverContext; @@ -104,10 +118,24 @@ public class StompIntegrationTests { @Qualifier("webSocketInputChannel") private QueueChannel webSocketInputChannel; + @Autowired + @Qualifier("webSocketEvents") + private PollableChannel webSocketEvents; + @Test public void sendMessageToController() throws Exception { + StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.CONNECT); + this.webSocketOutputChannel.send(MessageBuilder.withPayload(new byte[0]).setHeaders(headers).build()); - StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SEND); + Message receive = this.webSocketEvents.receive(10000); + assertNotNull(receive); + Object event = receive.getPayload(); + assertThat(event, instanceOf(SessionConnectedEvent.class)); + Message connectedMessage = ((SessionConnectedEvent) event).getMessage(); + headers = StompHeaderAccessor.wrap(connectedMessage); + assertEquals(StompCommand.CONNECTED, headers.getCommand()); + + headers = StompHeaderAccessor.create(StompCommand.SEND); headers.setSubscriptionId("sub1"); headers.setDestination("/app/simple"); Message message = MessageBuilder.withPayload("foo").setHeaders(headers).build(); @@ -124,22 +152,32 @@ public class StompIntegrationTests { StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE); headers.setSubscriptionId("subs1"); headers.setDestination("/topic/increment"); + headers.setReceipt("myReceipt"); Message message = MessageBuilder.withPayload(ByteBuffer.allocate(0).array()) .setHeaders(headers) .build(); + this.webSocketOutputChannel.send(message); + + Message receive = this.webSocketEvents.receive(10000); + assertNotNull(receive); + Object event = receive.getPayload(); + assertThat(event, instanceOf(ReceiptEvent.class)); + Message receiptMessage = ((ReceiptEvent) event).getMessage(); + headers = StompHeaderAccessor.wrap(receiptMessage); + assertEquals(StompCommand.RECEIPT, headers.getCommand()); + assertEquals("myReceipt", headers.getReceiptId()); + + waitForSubscribe("increment"); + headers = StompHeaderAccessor.create(StompCommand.SEND); headers.setSubscriptionId("subs1"); headers.setDestination("/app/increment"); Message message2 = MessageBuilder.withPayload(5).setHeaders(headers).build(); - this.webSocketOutputChannel.send(message); - - waitForSubscribe("increment"); - this.webSocketOutputChannel.send(message2); - Message receive = webSocketInputChannel.receive(10000); + receive = webSocketInputChannel.receive(10000); assertNotNull(receive); assertEquals("6", receive.getPayload()); } @@ -330,6 +368,19 @@ public class StompIntegrationTests { new SubProtocolHandlerRegistry(stompSubProtocolHandler())); } + @Bean + public PollableChannel webSocketEvents() { + return new QueueChannel(); + } + + @Bean + @SuppressWarnings("unchecked") + public ApplicationListener webSocketEventListener() { + ApplicationEventListeningMessageProducer producer = new ApplicationEventListeningMessageProducer(); + producer.setEventTypes(AbstractSubProtocolEvent.class); + producer.setOutputChannel(webSocketEvents()); + return producer; + } } // WebSocket Server part @@ -427,6 +478,28 @@ public class StompIntegrationTests { configurer.enableSimpleBroker("/topic", "/queue"); } + //TODO SimpleBrokerMessageHandler doesn't support RECEIPT frame, hence we emulate it this way + @Bean + @SuppressWarnings("unchecked") + public ApplicationListener webSocketEventListener( + final AbstractSubscribableChannel clientOutboundChannel) { + return new ApplicationListener() { + + @Override + public void onApplicationEvent(SessionSubscribeEvent event) { + Message message = event.getMessage(); + StompHeaderAccessor stompHeaderAccessor = StompHeaderAccessor.wrap(message); + if (stompHeaderAccessor.getReceipt() != null) { + stompHeaderAccessor.setHeader("stompCommand", StompCommand.RECEIPT); + stompHeaderAccessor.setReceiptId(stompHeaderAccessor.getReceipt()); + clientOutboundChannel.send( + MessageBuilder.createMessage(new byte[0], stompHeaderAccessor.getMessageHeaders())); + } + } + + }; + } + } }