diff --git a/build.gradle b/build.gradle
index 409eb79138..9741434238 100644
--- a/build.gradle
+++ b/build.gradle
@@ -88,6 +88,7 @@ subprojects { subproject ->
javaxActivationVersion = '1.1.1'
javaxMailVersion = '1.4.7'
jedisVersion = '2.4.2'
+ jettyVersion = '9.2.1.v20140609'
jmsApiVersion = '1.1-rev-1'
jpaApiVersion = '2.0.0'
jrubyVersion = '1.7.12'
@@ -118,7 +119,7 @@ subprojects { subproject ->
springSecurityVersion = '3.2.4.RELEASE'
springSocialTwitterVersion = '1.1.0.RELEASE'
springRetryVersion = '1.1.0.RELEASE'
- springVersion = project.hasProperty('springVersion') ? project.springVersion : '4.0.6.RELEASE'
+ springVersion = project.hasProperty('springVersion') ? project.springVersion : '4.1.0.BUILD-SNAPSHOT'
springWsVersion = '2.2.0.RELEASE'
xmlUnitVersion = '1.5'
xstreamVersion = '1.4.7'
@@ -595,6 +596,17 @@ project('spring-integration-websocket') {
dependencies {
compile project(":spring-integration-core")
compile "org.springframework:spring-websocket:$springVersion"
+
+ testCompile "org.springframework:spring-webmvc:$springVersion"
+ testCompile("org.eclipse.jetty:jetty-webapp:$jettyVersion") {
+ exclude group: "javax.servlet", module: "javax.servlet"
+ }
+ testCompile("org.eclipse.jetty.websocket:websocket-server:$jettyVersion") {
+ exclude group: "javax.servlet", module: "javax.servlet"
+ }
+ testCompile "org.eclipse.jetty.websocket:websocket-client:$jettyVersion"
+ testCompile"org.eclipse.jetty:jetty-client:$jettyVersion"
+ testCompile "org.slf4j:slf4j-jcl:$slf4jVersion"
}
}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/ClientWebSocketContainer.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/ClientWebSocketContainer.java
new file mode 100644
index 0000000000..df1346186f
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/ClientWebSocketContainer.java
@@ -0,0 +1,222 @@
+/*
+ * Copyright 2014 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;
+
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.TimeUnit;
+
+import org.springframework.context.Lifecycle;
+import org.springframework.context.SmartLifecycle;
+import org.springframework.http.HttpHeaders;
+import org.springframework.util.Assert;
+import org.springframework.util.concurrent.ListenableFuture;
+import org.springframework.util.concurrent.ListenableFutureCallback;
+import org.springframework.web.socket.CloseStatus;
+import org.springframework.web.socket.WebSocketHttpHeaders;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.client.ConnectionManagerSupport;
+import org.springframework.web.socket.client.WebSocketClient;
+
+/**
+ * The {@link IntegrationWebSocketContainer} implementation for the {@code client}
+ * Web-Socket connection.
+ *
+ * Represent the composition over an internal {@link ConnectionManagerSupport}
+ * implementation.
+ *
+ * Accepts the {@link #clientSession} {@link WebSocketSession} on
+ * {@link ClientWebSocketContainer.IntegrationWebSocketConnectionManager#openConnection()}
+ * event, which can be accessed from this container using {@link #getSession(String)}.
+ *
+ * @author Artem Bilan
+ * @since 4.1
+ */
+public final class ClientWebSocketContainer extends IntegrationWebSocketContainer implements SmartLifecycle {
+
+ private final WebSocketHttpHeaders headers = new WebSocketHttpHeaders();
+
+ private final ConnectionManagerSupport connectionManager;
+
+ private volatile CountDownLatch connectionLatch;
+
+ private WebSocketSession clientSession;
+
+ private volatile Throwable openConnectionException;
+
+ public ClientWebSocketContainer(WebSocketClient client, String uriTemplate, Object... uriVariables) {
+ Assert.notNull(client, "'client' must not be null");
+ this.connectionManager = new IntegrationWebSocketConnectionManager(client, uriTemplate, uriVariables);
+ }
+
+ public void setOrigin(String origin) {
+ this.headers.setOrigin(origin);
+ }
+
+ public void setHeaders(HttpHeaders headers) {
+ this.headers.clear();
+ this.headers.putAll(headers);
+ }
+
+ /**
+ * Return the {@link #clientSession} {@link WebSocketSession}.
+ * Independently of provided argument, this method always returns only the
+ * established {@link #clientSession}
+ * @param sessionId the {@code sessionId}. Can be {@code null}.
+ * @return the {@link #clientSession}, if established.
+ */
+ @Override
+ public WebSocketSession getSession(String sessionId) throws Exception {
+ if (this.isRunning()) {
+ try {
+ this.connectionLatch.await(10, TimeUnit.SECONDS);
+ }
+ catch (InterruptedException e) {
+ logger.error("'clientSession' has not been established during 'openConnection'");
+ }
+ }
+ if (this.openConnectionException != null) {
+ throw new IllegalStateException(this.openConnectionException);
+ }
+ Assert.state(this.clientSession != null,
+ "'clientSession' has not been established. Consider to 'start' this container.");
+ return this.clientSession;
+ }
+
+ public void setAutoStartup(boolean autoStartup) {
+ this.connectionManager.setAutoStartup(autoStartup);
+ }
+
+ public void setPhase(int phase) {
+ this.connectionManager.setPhase(phase);
+ }
+
+ @Override
+ public boolean isAutoStartup() {
+ return this.connectionManager.isAutoStartup();
+ }
+
+ @Override
+ public int getPhase() {
+ return this.connectionManager.getPhase();
+ }
+
+ @Override
+ public boolean isRunning() {
+ return this.connectionManager.isRunning();
+ }
+
+ @Override
+ public void start() {
+ this.connectionManager.start();
+ this.connectionLatch = new CountDownLatch(1);
+ }
+
+ @Override
+ public void stop() {
+ this.connectionManager.stop();
+ }
+
+ @Override
+ public void stop(Runnable callback) {
+ this.connectionManager.stop(callback);
+ }
+
+ /**
+ * The {@link ConnectionManagerSupport} implementation to provide open/close operations
+ * for an external Web-Socket service, based on provided {@link WebSocketClient} and {@code uriTemplate}.
+ *
+ * Opened {@link WebSocketSession} is populated to the wrapping {@link ClientWebSocketContainer}.
+ *
+ * The {@link #webSocketHandler} is used to handle {@link WebSocketSession} events.
+ */
+ private class IntegrationWebSocketConnectionManager extends ConnectionManagerSupport {
+
+ private final WebSocketClient client;
+
+ private final boolean syncClientLifecycle;
+
+ public IntegrationWebSocketConnectionManager(WebSocketClient client, String uriTemplate, Object... uriVariables) {
+ super(uriTemplate, uriVariables);
+ this.client = client;
+ this.syncClientLifecycle = ((client instanceof Lifecycle) && !((Lifecycle) client).isRunning());
+ }
+
+ @Override
+ public void startInternal() {
+ if (this.syncClientLifecycle) {
+ ((Lifecycle) this.client).start();
+ }
+ super.startInternal();
+ }
+
+ @Override
+ public void stopInternal() throws Exception {
+ if (this.syncClientLifecycle) {
+ ((Lifecycle) this.client).stop();
+ }
+ try {
+ super.stopInternal();
+ }
+ finally {
+ ClientWebSocketContainer.this.clientSession = null;
+ }
+ }
+
+ @Override
+ protected void openConnection() {
+
+ logger.info("Connecting to WebSocket at " + getUri());
+ ClientWebSocketContainer.this.headers.setSecWebSocketProtocol(ClientWebSocketContainer.this.getSubProtocols());
+ ListenableFuture future =
+ this.client.doHandshake(ClientWebSocketContainer.this.webSocketHandler,
+ ClientWebSocketContainer.this.headers, getUri());
+
+ future.addCallback(new ListenableFutureCallback() {
+
+ @Override
+ public void onSuccess(WebSocketSession session) {
+ ClientWebSocketContainer.this.clientSession = session;
+ logger.info("Successfully connected");
+ ClientWebSocketContainer.this.connectionLatch.countDown();
+ }
+
+ @Override
+ public void onFailure(Throwable t) {
+ logger.error("Failed to connect", t);
+ ClientWebSocketContainer.this.openConnectionException = t;
+ ClientWebSocketContainer.this.connectionLatch.countDown();
+ }
+ });
+ }
+
+ @Override
+ protected void closeConnection() throws Exception {
+ if (ClientWebSocketContainer.this.clientSession != null) {
+ ClientWebSocketContainer.this.closeSession(ClientWebSocketContainer.this.clientSession,
+ CloseStatus.NORMAL);
+ }
+ }
+
+ @Override
+ protected boolean isConnected() {
+ return ((ClientWebSocketContainer.this.clientSession != null)
+ && (ClientWebSocketContainer.this.clientSession.isOpen()));
+ }
+
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/IntegrationWebSocketContainer.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/IntegrationWebSocketContainer.java
new file mode 100644
index 0000000000..9a6d6f0954
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/IntegrationWebSocketContainer.java
@@ -0,0 +1,226 @@
+/*
+ * Copyright 2014 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;
+
+import java.util.ArrayList;
+import java.util.Collections;
+import java.util.List;
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
+
+import org.apache.commons.logging.Log;
+import org.apache.commons.logging.LogFactory;
+
+import org.springframework.beans.factory.DisposableBean;
+import org.springframework.context.ApplicationEvent;
+import org.springframework.context.ApplicationEventPublisher;
+import org.springframework.context.ApplicationEventPublisherAware;
+import org.springframework.util.Assert;
+import org.springframework.web.socket.CloseStatus;
+import org.springframework.web.socket.SubProtocolCapable;
+import org.springframework.web.socket.WebSocketHandler;
+import org.springframework.web.socket.WebSocketMessage;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.handler.ConcurrentWebSocketSessionDecorator;
+import org.springframework.web.socket.messaging.SessionDisconnectEvent;
+
+/**
+ * The high-level 'connection factory pattern' contract over low-level Web-Socket
+ * configuration.
+ *
+ * Provides the composition for the internal {@link WebSocketHandler}
+ * implementation, which is used with native Web-Socket containers.
+ *
+ * Collects established {@link WebSocketSession}s, which can be accessed using
+ * {@link #getSession(String)}.
+ *
+ * Can accept the {@link WebSocketListener} to delegate {@link WebSocketSession} events
+ * from the internal {@link IntegrationWebSocketContainer.IntegrationWebSocketHandler}.
+ *
+ * Supported sub-protocols can be configured, but {@link WebSocketListener#getSubProtocols()}
+ * have a precedent.
+ *
+ * @author Artem Bilan
+ * @since 4.1
+ * @see org.springframework.integration.websocket.inbound.WebSocketInboundChannelAdapter
+ * @see org.springframework.integration.websocket.outbound.WebSocketOutboundMessageHandler
+ */
+public abstract class IntegrationWebSocketContainer implements ApplicationEventPublisherAware, DisposableBean {
+
+ protected final Log logger = LogFactory.getLog(this.getClass());
+
+ protected final WebSocketHandler webSocketHandler = new IntegrationWebSocketHandler();
+
+ protected final Map sessions = new ConcurrentHashMap();
+
+ private final List supportedProtocols = new ArrayList();
+
+ private volatile WebSocketListener messageListener;
+
+ private volatile int sendTimeLimit = 10 * 1000;
+
+ private volatile int sendBufferSizeLimit = 512 * 1024;
+
+ private ApplicationEventPublisher eventPublisher;
+
+ public void setSendTimeLimit(int sendTimeLimit) {
+ this.sendTimeLimit = sendTimeLimit;
+ }
+
+ public void setSendBufferSizeLimit(int sendBufferSizeLimit) {
+ this.sendBufferSizeLimit = sendBufferSizeLimit;
+ }
+
+ @Override
+ public void setApplicationEventPublisher(ApplicationEventPublisher applicationEventPublisher) {
+ this.eventPublisher = applicationEventPublisher;
+ }
+
+ public void setMessageListener(WebSocketListener messageListener) {
+ Assert.state(this.messageListener == null || this.messageListener == messageListener,
+ "'messageListener' is already configured");
+ this.messageListener = messageListener;
+ }
+
+ public void setSupportedProtocols(String... protocols) {
+ this.supportedProtocols.clear();
+ addSupportedProtocols(protocols);
+ }
+
+ public void addSupportedProtocols(String... protocols) {
+ for (String protocol : protocols) {
+ this.supportedProtocols.add(protocol.toLowerCase());
+ }
+ }
+
+ public List getSubProtocols() {
+ List protocols = new ArrayList();
+ if (this.messageListener != null) {
+ protocols.addAll(this.messageListener.getSubProtocols());
+ }
+ protocols.addAll(this.supportedProtocols);
+ return Collections.unmodifiableList(protocols);
+ }
+
+ public WebSocketSession getSession(String sessionId) throws Exception {
+ WebSocketSession session = this.sessions.get(sessionId);
+ Assert.notNull(session, "Session not found for id '" + sessionId + "'");
+ return session;
+ }
+
+ public void closeSession(WebSocketSession session, CloseStatus closeStatus) throws Exception {
+ // Session may be unresponsive so clear first
+ session.close(closeStatus);
+ this.webSocketHandler.afterConnectionClosed(session, closeStatus);
+ }
+
+ @Override
+ public void destroy() throws Exception {
+ // Notify sessions to stop flushing messages
+ for (WebSocketSession session : this.sessions.values()) {
+ try {
+ session.close(CloseStatus.GOING_AWAY);
+ }
+ catch (Throwable t) {
+ logger.error("Failed to close session id '" + session.getId() + "': " + t.getMessage());
+ }
+ }
+ this.sessions.clear();
+ }
+
+ private void publishEvent(ApplicationEvent event) {
+ try {
+ this.eventPublisher.publishEvent(event);
+ }
+ catch (Throwable ex) {
+ logger.error("Error while publishing " + event, ex);
+ }
+ }
+
+ /**
+ * An internal {@link WebSocketHandler} implementation to be used with native
+ * Web-Socket containers.
+ *
+ * Delegates all operations to the wrapping {@link IntegrationWebSocketContainer}
+ * and its {@link WebSocketListener}.
+ */
+ private class IntegrationWebSocketHandler implements WebSocketHandler, SubProtocolCapable {
+
+ @Override
+ public List getSubProtocols() {
+ return IntegrationWebSocketContainer.this.getSubProtocols();
+ }
+
+ @Override
+ public void afterConnectionEstablished(WebSocketSession session) throws Exception {
+ session = new ConcurrentWebSocketSessionDecorator(session,
+ IntegrationWebSocketContainer.this.sendTimeLimit,
+ IntegrationWebSocketContainer.this.sendBufferSizeLimit);
+
+ IntegrationWebSocketContainer.this.sessions.put(session.getId(), session);
+ if (logger.isDebugEnabled()) {
+ logger.debug("Started WebSocket session = " + session.getId() + ", number of sessions = "
+ + IntegrationWebSocketContainer.this.sessions.size());
+ }
+ if (IntegrationWebSocketContainer.this.messageListener != null) {
+ IntegrationWebSocketContainer.this.messageListener.afterSessionStarted(session);
+ }
+ }
+
+ @Override
+ public void afterConnectionClosed(WebSocketSession session, CloseStatus closeStatus) throws Exception {
+ WebSocketSession removed = IntegrationWebSocketContainer.this.sessions.remove(session.getId());
+ if (removed != null) {
+ if (IntegrationWebSocketContainer.this.messageListener != null) {
+ IntegrationWebSocketContainer.this.messageListener.afterSessionEnded(session, closeStatus);
+ }
+ else if (IntegrationWebSocketContainer.this.eventPublisher != null) {
+ publishEvent(new SessionDisconnectEvent(this, session.getId(), closeStatus));
+ }
+ }
+ }
+
+ @Override
+ public void handleTransportError(WebSocketSession session, Throwable exception) throws Exception {
+ WebSocketSession removed = IntegrationWebSocketContainer.this.sessions.remove(session.getId());
+ if (removed != null) {
+ IntegrationWebSocketContainer.this.sessions.remove(session.getId());
+ if (IntegrationWebSocketContainer.this.eventPublisher != null) {
+ publishEvent(new SessionErrorEvent(this, session.getId(), exception));
+ }
+ }
+ }
+
+ @Override
+ public void handleMessage(WebSocketSession session, WebSocketMessage> message) throws Exception {
+ if (IntegrationWebSocketContainer.this.messageListener != null) {
+ IntegrationWebSocketContainer.this.messageListener.onMessage(session, message);
+ }
+ else if (logger.isInfoEnabled()) {
+ logger.info("This 'WebSocketHandlerContainer' isn't configured with 'WebSocketMessageListener'."
+ + " Received messages are ignored. Current message is: " + message);
+ }
+ }
+
+ @Override
+ public boolean supportsPartialMessages() {
+ return false;
+ }
+
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/SessionErrorEvent.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/SessionErrorEvent.java
new file mode 100644
index 0000000000..1fd6a5b329
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/SessionErrorEvent.java
@@ -0,0 +1,71 @@
+/*
+ * Copyright 2014 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;
+
+import org.springframework.context.ApplicationEvent;
+import org.springframework.util.Assert;
+
+/**
+ * The {@link ApplicationEvent} implementation to represent the
+ * {@link org.springframework.web.socket.WebSocketSession} errors.
+ *
+ * @author Artem Bilan
+ * @since 4.1
+ */
+@SuppressWarnings("serial")
+public class SessionErrorEvent extends ApplicationEvent {
+
+ private final String sessionId;
+
+ private final Throwable exception;
+
+
+ /**
+ * Create a new {@link ApplicationEvent} represented the error on the session.
+ * @param source the component that published the event (never {@code null})
+ * @param sessionId the id of the session
+ * @param exception the exception on the session
+ */
+ public SessionErrorEvent(Object source, String sessionId, Throwable exception) {
+ super(source);
+ Assert.notNull(sessionId, "'sessionId' must not be null");
+ this.sessionId = sessionId;
+ this.exception = exception;
+ }
+
+ /**
+ * Return the session id.
+ * @return the sessionId
+ */
+ public String getSessionId() {
+ return this.sessionId;
+ }
+
+ /**
+ * Return the exception for the session id.
+ * @return the exception
+ */
+ public Throwable getException() {
+ return exception;
+ }
+
+ @Override
+ public String toString() {
+ return "SessionErrorEvent: sessionId=" + this.sessionId;
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/WebSocketListener.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/WebSocketListener.java
new file mode 100644
index 0000000000..488759d554
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/WebSocketListener.java
@@ -0,0 +1,62 @@
+/*
+ * Copyright 2014 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;
+
+import org.springframework.web.socket.CloseStatus;
+import org.springframework.web.socket.SubProtocolCapable;
+import org.springframework.web.socket.WebSocketMessage;
+import org.springframework.web.socket.WebSocketSession;
+
+/**
+ * A contract for handling incoming {@link WebSocketMessage}s messages as part of a higher
+ * level protocol, referred to as "sub-protocol" in the WebSocket RFC specification.
+ *
+ * Implementations of this interface can be configured on a
+ * {@link IntegrationWebSocketContainer} which delegates messages and
+ * {@link WebSocketSession} events to this implementation.
+ *
+ * @author Andy Wilkinson
+ * @author Artem Bilan
+ * @since 4.1
+ * @see org.springframework.integration.websocket.inbound.WebSocketInboundChannelAdapter
+ */
+public interface WebSocketListener extends SubProtocolCapable {
+
+ /**
+ * Handle the received {@link WebSocketMessage}.
+ * @param session the WebSocket session
+ * @param message the WebSocket message
+ * @throws Exception the 'onMessage' Exception
+ */
+ void onMessage(WebSocketSession session, WebSocketMessage> message) throws Exception;
+
+ /**
+ * Invoked after a {@link WebSocketSession} has started.
+ * @param session the WebSocket session
+ * @throws Exception the 'afterSessionStarted' Exception
+ */
+ void afterSessionStarted(WebSocketSession session) throws Exception;
+
+ /**
+ * Invoked after a {@link WebSocketSession} has ended.
+ * @param session the WebSocket session
+ * @param closeStatus the reason why the session was closed
+ * @throws Exception the 'afterSessionEnded' Exception
+ */
+ void afterSessionEnded(WebSocketSession session, CloseStatus closeStatus) throws Exception;
+
+}
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
new file mode 100644
index 0000000000..627001a8e6
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapter.java
@@ -0,0 +1,225 @@
+/*
+ * Copyright 2014 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.inbound;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.ListIterator;
+import java.util.concurrent.atomic.AtomicReference;
+
+import org.springframework.context.Lifecycle;
+import org.springframework.integration.channel.FixedSubscriberChannel;
+import org.springframework.integration.endpoint.MessageProducerSupport;
+import org.springframework.integration.support.MessageBuilder;
+import org.springframework.integration.support.json.JacksonJsonUtils;
+import org.springframework.integration.websocket.IntegrationWebSocketContainer;
+import org.springframework.integration.websocket.WebSocketListener;
+import org.springframework.integration.websocket.support.PassThruSubProtocolHandler;
+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.MessagingException;
+import org.springframework.messaging.converter.ByteArrayMessageConverter;
+import org.springframework.messaging.converter.CompositeMessageConverter;
+import org.springframework.messaging.converter.DefaultContentTypeResolver;
+import org.springframework.messaging.converter.MappingJackson2MessageConverter;
+import org.springframework.messaging.converter.MessageConverter;
+import org.springframework.messaging.converter.StringMessageConverter;
+import org.springframework.messaging.simp.SimpMessageHeaderAccessor;
+import org.springframework.messaging.simp.SimpMessageType;
+import org.springframework.util.Assert;
+import org.springframework.util.CollectionUtils;
+import org.springframework.util.MimeTypeUtils;
+import org.springframework.web.socket.CloseStatus;
+import org.springframework.web.socket.WebSocketMessage;
+import org.springframework.web.socket.WebSocketSession;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+public class WebSocketInboundChannelAdapter extends MessageProducerSupport implements WebSocketListener {
+
+ private final List defaultConverters = new ArrayList(3);
+
+ {
+ this.defaultConverters.add(new StringMessageConverter());
+ this.defaultConverters.add(new ByteArrayMessageConverter());
+ if (JacksonJsonUtils.isJackson2Present()) {
+ DefaultContentTypeResolver resolver = new DefaultContentTypeResolver();
+ resolver.setDefaultMimeType(MimeTypeUtils.APPLICATION_JSON);
+ MappingJackson2MessageConverter converter = new MappingJackson2MessageConverter();
+ converter.setContentTypeResolver(resolver);
+ this.defaultConverters.add(converter);
+ }
+ }
+
+ private final CompositeMessageConverter messageConverter = new CompositeMessageConverter(this.defaultConverters);
+
+ private final IntegrationWebSocketContainer webSocketContainer;
+
+ private final SubProtocolHandlerRegistry protocolHandlerContainer;
+
+ private final MessageChannel subProtocolHandlerChannel;
+
+ private final AtomicReference> payloadType = new AtomicReference>(String.class);
+
+ private volatile List messageConverters;
+
+ private volatile boolean mergeWithDefaultConverters = false;
+
+ private volatile boolean active;
+
+ public WebSocketInboundChannelAdapter(IntegrationWebSocketContainer webSocketContainer) {
+ this(webSocketContainer, new SubProtocolHandlerRegistry(new PassThruSubProtocolHandler()));
+ }
+
+ public WebSocketInboundChannelAdapter(IntegrationWebSocketContainer webSocketContainer,
+ SubProtocolHandlerRegistry protocolHandlerRegistry) {
+ Assert.notNull(webSocketContainer, "'webSocketContainer' must not be null");
+ Assert.notNull(protocolHandlerRegistry, "'protocolHandlerRegistry' must not be null");
+ this.webSocketContainer = webSocketContainer;
+ this.protocolHandlerContainer = protocolHandlerRegistry;
+ this.subProtocolHandlerChannel = new FixedSubscriberChannel(new MessageHandler() {
+
+ @Override
+ public void handleMessage(Message> message) throws MessagingException {
+ Object payload = WebSocketInboundChannelAdapter.this.messageConverter.fromMessage(message,
+ WebSocketInboundChannelAdapter.this.payloadType.get());
+ SimpMessageHeaderAccessor headerAccessor = SimpMessageHeaderAccessor.wrap(message);
+ SimpMessageType messageType = headerAccessor.getMessageType();
+ if (messageType == null || SimpMessageType.MESSAGE.equals(messageType)) {
+ headerAccessor.removeHeader(SimpMessageHeaderAccessor.NATIVE_HEADERS);
+ sendMessage(MessageBuilder.withPayload(payload).copyHeaders(headerAccessor.toMap()).build());
+ }
+ else {
+ if (logger.isDebugEnabled()) {
+ logger.debug("Messages with non 'SimpMessageType.MESSAGE' type are ignored for sending to the " +
+ "'outputChannel'. They have to be emitted as 'ApplicationEvent's " +
+ "from the 'SubProtocolHandler'. Received message: " + message);
+ }
+ }
+ }
+
+ });
+ }
+
+ /**
+ * Set the message converters to use. These converters are used to convert the message to send for appropriate
+ * internal subProtocols type.
+ * @param messageConverters The message converters.
+ */
+ public void setMessageConverters(List messageConverters) {
+ Assert.noNullElements(messageConverters.toArray(), "'messageConverters' must not contain null entries");
+ this.messageConverters = new ArrayList(messageConverters);
+ }
+
+
+ /**
+ * Flag which determines if the default converters should be available after
+ * custom converters.
+ * @param mergeWithDefaultConverters true to merge, false to replace.
+ */
+ public void setMergeWithDefaultConverters(boolean mergeWithDefaultConverters) {
+ this.mergeWithDefaultConverters = mergeWithDefaultConverters;
+ }
+
+ /**
+ * Set the type for target message payload to convert the WebSocket message body to.
+ * @param payloadType to convert inbound WebSocket message body
+ * @see CompositeMessageConverter
+ */
+ public void setPayloadType(Class> payloadType) {
+ Assert.notNull(payloadType, "'payloadType' must not be null");
+ this.payloadType.set(payloadType);
+ }
+
+ @Override
+ protected void onInit() {
+ super.onInit();
+ this.webSocketContainer.setMessageListener(this);
+ if (!CollectionUtils.isEmpty(this.messageConverters)) {
+ List converters = this.messageConverter.getConverters();
+ if (this.mergeWithDefaultConverters) {
+ for (ListIterator iterator = this.messageConverters.listIterator(); iterator.hasPrevious(); ) {
+ MessageConverter converter = iterator.previous();
+ converters.add(0, converter);
+ }
+ }
+ else {
+ converters.clear();
+ converters.addAll(this.messageConverters);
+ }
+ }
+ }
+
+ @Override
+ public List getSubProtocols() {
+ return this.protocolHandlerContainer.getSubProtocols();
+ }
+
+ @Override
+ public void afterSessionStarted(WebSocketSession session) throws Exception {
+ if (isActive()) {
+ this.protocolHandlerContainer.findProtocolHandler(session)
+ .afterSessionStarted(session, this.subProtocolHandlerChannel);
+ }
+ }
+
+ @Override
+ public void afterSessionEnded(WebSocketSession session, CloseStatus closeStatus) throws Exception {
+ if (isActive()) {
+ this.protocolHandlerContainer.findProtocolHandler(session)
+ .afterSessionEnded(session, closeStatus, this.subProtocolHandlerChannel);
+ }
+ }
+
+ @Override
+ public void onMessage(WebSocketSession session, WebSocketMessage> webSocketMessage) throws Exception {
+ if (isActive()) {
+ this.protocolHandlerContainer.findProtocolHandler(session)
+ .handleMessageFromClient(session, webSocketMessage, this.subProtocolHandlerChannel);
+ }
+ }
+
+ @Override
+ public String getComponentType() {
+ return "websocket:inbound-channel-adapter";
+ }
+
+ @Override
+ protected void doStart() {
+ this.active = true;
+ if (this.webSocketContainer instanceof Lifecycle) {
+ ((Lifecycle) this.webSocketContainer).start();
+ }
+ }
+
+ @Override
+ protected void doStop() {
+ this.active = false;
+ }
+
+ private boolean isActive() {
+ if (!this.active) {
+ logger.warn("MessageProducer '" + this + "'isn't started to accept WebSocket events");
+ }
+ return this.active;
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/package-info.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/package-info.java
new file mode 100644
index 0000000000..ada6bf5018
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/inbound/package-info.java
@@ -0,0 +1,4 @@
+/**
+ * Provides classes which represent inbound WebSocket components.
+ */
+package org.springframework.integration.websocket.inbound;
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/outbound/WebSocketOutboundMessageHandler.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/outbound/WebSocketOutboundMessageHandler.java
new file mode 100644
index 0000000000..10df329fc2
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/outbound/WebSocketOutboundMessageHandler.java
@@ -0,0 +1,162 @@
+/*
+ * Copyright 2014 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.outbound;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.ListIterator;
+
+import org.springframework.integration.handler.AbstractMessageHandler;
+import org.springframework.integration.support.json.JacksonJsonUtils;
+import org.springframework.integration.websocket.ClientWebSocketContainer;
+import org.springframework.integration.websocket.IntegrationWebSocketContainer;
+import org.springframework.integration.websocket.support.PassThruSubProtocolHandler;
+import org.springframework.integration.websocket.support.SubProtocolHandlerRegistry;
+import org.springframework.messaging.Message;
+import org.springframework.messaging.converter.ByteArrayMessageConverter;
+import org.springframework.messaging.converter.CompositeMessageConverter;
+import org.springframework.messaging.converter.DefaultContentTypeResolver;
+import org.springframework.messaging.converter.MappingJackson2MessageConverter;
+import org.springframework.messaging.converter.MessageConverter;
+import org.springframework.messaging.converter.StringMessageConverter;
+import org.springframework.messaging.simp.SimpMessageHeaderAccessor;
+import org.springframework.messaging.simp.SimpMessageType;
+import org.springframework.util.Assert;
+import org.springframework.util.CollectionUtils;
+import org.springframework.util.MimeTypeUtils;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.handler.SessionLimitExceededException;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+public class WebSocketOutboundMessageHandler extends AbstractMessageHandler {
+
+ private final List defaultConverters = new ArrayList(3);
+
+ {
+ this.defaultConverters.add(new StringMessageConverter());
+ this.defaultConverters.add(new ByteArrayMessageConverter());
+ if (JacksonJsonUtils.isJackson2Present()) {
+ DefaultContentTypeResolver resolver = new DefaultContentTypeResolver();
+ resolver.setDefaultMimeType(MimeTypeUtils.APPLICATION_JSON);
+ MappingJackson2MessageConverter converter = new MappingJackson2MessageConverter();
+ converter.setContentTypeResolver(resolver);
+ this.defaultConverters.add(converter);
+ }
+ }
+
+ private final CompositeMessageConverter messageConverter = new CompositeMessageConverter(this.defaultConverters);
+
+ private final IntegrationWebSocketContainer webSocketContainer;
+
+ private final SubProtocolHandlerRegistry protocolHandlerContainer;
+
+ private final boolean client;
+
+ private volatile List messageConverters;
+
+ private volatile boolean mergeWithDefaultConverters = false;
+
+ public WebSocketOutboundMessageHandler(IntegrationWebSocketContainer webSocketContainer) {
+ this(webSocketContainer, new SubProtocolHandlerRegistry(new PassThruSubProtocolHandler()));
+ }
+
+ public WebSocketOutboundMessageHandler(IntegrationWebSocketContainer webSocketContainer,
+ SubProtocolHandlerRegistry protocolHandlerRegistry) {
+ Assert.notNull(webSocketContainer, "'webSocketContainer' must not be null");
+ Assert.notNull(protocolHandlerRegistry, "'protocolHandlerRegistry' must not be null");
+ this.webSocketContainer = webSocketContainer;
+ this.client = webSocketContainer instanceof ClientWebSocketContainer;
+ this.protocolHandlerContainer = protocolHandlerRegistry;
+ List subProtocols = protocolHandlerRegistry.getSubProtocols();
+ this.webSocketContainer.addSupportedProtocols(subProtocols.toArray(new String[subProtocols.size()]));
+ }
+
+ /**
+ * Set the message converters to use. These converters are used to convert the message to send for appropriate
+ * internal subProtocols type.
+ * @param messageConverters The message converters.
+ */
+ public void setMessageConverters(List messageConverters) {
+ Assert.noNullElements(messageConverters.toArray(), "'messageConverters' must not contain null entries");
+ this.messageConverters = new ArrayList(messageConverters);
+ }
+
+
+ /**
+ * Flag which determines if the default converters should be available after
+ * custom converters.
+ * @param mergeWithDefaultConverters true to merge, false to replace.
+ */
+ public void setMergeWithDefaultConverters(boolean mergeWithDefaultConverters) {
+ this.mergeWithDefaultConverters = mergeWithDefaultConverters;
+ }
+
+ @Override
+ public String getComponentType() {
+ return "websocket:outbound-channel-adapter";
+ }
+
+ @Override
+ protected void onInit() throws Exception {
+ super.onInit();
+ if (!CollectionUtils.isEmpty(this.messageConverters)) {
+ List converters = this.messageConverter.getConverters();
+ if (this.mergeWithDefaultConverters) {
+ for (ListIterator iterator = this.messageConverters.listIterator(); iterator.hasPrevious(); ) {
+ MessageConverter converter = iterator.previous();
+ converters.add(0, converter);
+ }
+ }
+ else {
+ converters.clear();
+ converters.addAll(this.messageConverters);
+ }
+ }
+ }
+
+ @Override
+ protected void handleMessageInternal(Message> message) throws Exception {
+ String sessionId = null;
+ if (!this.client) {
+ sessionId = this.protocolHandlerContainer.resolveSessionId(message);
+ if (sessionId == null) {
+ throw new IllegalArgumentException("The WebSocket 'sessionId' is required in the MessageHeaders");
+ }
+ }
+ WebSocketSession session = this.webSocketContainer.getSession(sessionId);
+ try {
+ SimpMessageHeaderAccessor headers = SimpMessageHeaderAccessor.wrap(message);
+ headers.setLeaveMutable(true);
+ headers.setMessageTypeIfNotSet(SimpMessageType.MESSAGE);
+ Message> messageToSend = this.messageConverter.toMessage(message.getPayload(), headers.getMessageHeaders());
+ this.protocolHandlerContainer.findProtocolHandler(session).handleMessageToClient(session, messageToSend);
+ }
+ catch (SessionLimitExceededException ex) {
+ try {
+ logger.error("Terminating session id '" + sessionId + "'", ex);
+ this.webSocketContainer.closeSession(session, ex.getStatus());
+ }
+ catch (Exception secondException) {
+ logger.error("Exception terminating session id '" + sessionId + "'", secondException);
+ }
+ }
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/outbound/package-info.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/outbound/package-info.java
new file mode 100644
index 0000000000..cb94e510e2
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/outbound/package-info.java
@@ -0,0 +1,4 @@
+/**
+ * Provides classes which represent outbound WebSocket components.
+ */
+package org.springframework.integration.websocket.outbound;
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/package-info.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/package-info.java
new file mode 100644
index 0000000000..d6fcd36907
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/package-info.java
@@ -0,0 +1,4 @@
+/**
+ * Provides classes used across all WebSocket components.
+ */
+package org.springframework.integration.websocket;
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/PassThruSubProtocolHandler.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/PassThruSubProtocolHandler.java
new file mode 100644
index 0000000000..54072082cb
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/PassThruSubProtocolHandler.java
@@ -0,0 +1,114 @@
+/*
+ * Copyright 2014 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.support;
+
+import java.nio.ByteBuffer;
+import java.util.ArrayList;
+import java.util.Arrays;
+import java.util.List;
+
+import org.springframework.messaging.Message;
+import org.springframework.messaging.MessageChannel;
+import org.springframework.messaging.simp.SimpAttributesContextHolder;
+import org.springframework.messaging.simp.SimpMessageHeaderAccessor;
+import org.springframework.messaging.simp.SimpMessageType;
+import org.springframework.messaging.support.MessageBuilder;
+import org.springframework.util.Assert;
+import org.springframework.web.socket.CloseStatus;
+import org.springframework.web.socket.TextMessage;
+import org.springframework.web.socket.WebSocketMessage;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+
+/**
+ * The simple 'pass thru' {@link SubProtocolHandler}, when there is no interests in the
+ * WebSocket sub-protocols.
+ * This class just convert {@link Message} to the {@link WebSocketMessage}
+ * on 'send' part and vise versa - on 'receive' part.
+ *
+ * @author Artem Bilan
+ * @since 4.1
+ */
+public class PassThruSubProtocolHandler implements SubProtocolHandler {
+
+ final List supportedProtocols = new ArrayList();
+
+ public void setSupportedProtocols(String... supportedProtocols) {
+ Assert.noNullElements(supportedProtocols, "'supportedProtocols' must not be empty");
+ this.supportedProtocols.addAll(Arrays.asList(supportedProtocols));
+ }
+
+ @Override
+ public List getSupportedProtocols() {
+ return supportedProtocols;
+ }
+
+ @Override
+ public void handleMessageFromClient(WebSocketSession session, WebSocketMessage> webSocketMessage,
+ MessageChannel outputChannel) throws Exception {
+ SimpMessageHeaderAccessor headerAccessor = SimpMessageHeaderAccessor.create(SimpMessageType.MESSAGE);
+ headerAccessor.setSessionId(session.getId());
+ headerAccessor.setSessionAttributes(session.getAttributes());
+ headerAccessor.setUser(session.getPrincipal());
+ headerAccessor.setHeader("content-length", webSocketMessage.getPayloadLength());
+ headerAccessor.setLeaveMutable(true);
+ Message> message =
+ MessageBuilder.createMessage(webSocketMessage.getPayload(), headerAccessor.getMessageHeaders());
+ try {
+ SimpAttributesContextHolder.setAttributesFromMessage(message);
+ outputChannel.send(message);
+ }
+ finally {
+ SimpAttributesContextHolder.resetAttributes();
+ }
+ }
+
+ @Override
+ public void handleMessageToClient(WebSocketSession session, Message> message) throws Exception {
+ Object payload = message.getPayload();
+ if (payload instanceof String) {
+ session.sendMessage(new TextMessage((String) payload));
+ }
+ else if (payload instanceof byte[]) {
+ session.sendMessage(new TextMessage((byte[]) payload));
+ }
+ else if (payload instanceof ByteBuffer) {
+ session.sendMessage(new TextMessage(((ByteBuffer) payload).array()));
+ }
+ else {
+ throw new IllegalArgumentException("Unsupported payload type: " + payload.getClass()
+ + ". Can be one of: " + Arrays.>asList(String.class, byte[].class, ByteBuffer.class));
+ }
+ }
+
+ @Override
+ public String resolveSessionId(Message> message) {
+ return SimpMessageHeaderAccessor.getSessionId(message.getHeaders());
+ }
+
+ @Override
+ public void afterSessionStarted(WebSocketSession session, MessageChannel outputChannel) throws Exception {
+ // Subclasses might implement this method
+ }
+
+ @Override
+ public void afterSessionEnded(WebSocketSession session, CloseStatus closeStatus, MessageChannel outputChannel)
+ throws Exception {
+ // Subclasses might implement this method
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/SubProtocolHandlerRegistry.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/SubProtocolHandlerRegistry.java
new file mode 100644
index 0000000000..aa265ea0d0
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/SubProtocolHandlerRegistry.java
@@ -0,0 +1,155 @@
+/*
+ * Copyright 2014 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.support;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.Map;
+import java.util.TreeMap;
+
+import org.apache.commons.logging.Log;
+import org.apache.commons.logging.LogFactory;
+
+import org.springframework.messaging.Message;
+import org.springframework.util.Assert;
+import org.springframework.util.CollectionUtils;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+
+/**
+ * The utility class to encapsulate search algorithms for a set of provided {@link SubProtocolHandler}s.
+ *
+ * For internal use only.
+ *
+ * @author Andy Wilkinson
+ * @author Artem Bilan
+ * @since 4.1
+ * @see org.springframework.integration.websocket.inbound.WebSocketInboundChannelAdapter
+ * @see org.springframework.integration.websocket.outbound.WebSocketOutboundMessageHandler
+ */
+public final class SubProtocolHandlerRegistry {
+
+ private final static Log logger = LogFactory.getLog(SubProtocolHandlerRegistry.class);
+
+ private final Map protocolHandlers =
+ new TreeMap(String.CASE_INSENSITIVE_ORDER);
+
+ private final SubProtocolHandler defaultProtocolHandler;
+
+ public SubProtocolHandlerRegistry(List protocolHandlers) {
+ this(protocolHandlers, null);
+ }
+
+ public SubProtocolHandlerRegistry(SubProtocolHandler defaultProtocolHandler) {
+ this(null, defaultProtocolHandler);
+ }
+
+ public SubProtocolHandlerRegistry(List protocolHandlers,
+ SubProtocolHandler defaultProtocolHandler) {
+ Assert.state(!CollectionUtils.isEmpty(protocolHandlers) || defaultProtocolHandler != null,
+ "One of 'protocolHandlers' or 'defaultProtocolHandler' must be provided");
+
+ if (!CollectionUtils.isEmpty(protocolHandlers)) {
+ for (SubProtocolHandler handler : protocolHandlers) {
+ List protocols = handler.getSupportedProtocols();
+ if (CollectionUtils.isEmpty(protocols)) {
+ logger.warn("No sub-protocols, ignoring handler " + handler);
+ continue;
+ }
+ for (String protocol : protocols) {
+ SubProtocolHandler replaced = this.protocolHandlers.put(protocol, handler);
+ if (replaced != null) {
+ throw new IllegalStateException("Failed to map handler " + handler
+ + " to protocol '" + protocol + "', it is already mapped to handler " + replaced);
+ }
+ }
+ }
+ }
+
+ if (this.protocolHandlers.size() == 1 && defaultProtocolHandler == null) {
+ this.defaultProtocolHandler = this.protocolHandlers.values().iterator().next();
+ }
+ else {
+ this.defaultProtocolHandler = defaultProtocolHandler;
+ if (this.protocolHandlers.isEmpty()) {
+ List protocols = this.defaultProtocolHandler.getSupportedProtocols();
+ for (String protocol : protocols) {
+ SubProtocolHandler replaced = this.protocolHandlers.put(protocol, this.defaultProtocolHandler);
+ if (replaced != null) {
+ throw new IllegalStateException("Failed to map handler " + this.defaultProtocolHandler
+ + " to protocol '" + protocol + "', it is already mapped to handler " + replaced);
+ }
+ }
+ }
+ }
+ }
+
+ /**
+ * Resolves the {@link SubProtocolHandler} for the given {@code session} using
+ * its {@link WebSocketSession#getAcceptedProtocol() accepted sub-protocol}.
+ * @param session The session to resolve the sub-protocol handler for
+ * @return The sub-protocol handler
+ * @throws IllegalStateException if a protocol handler cannot be resolved
+ */
+ public SubProtocolHandler findProtocolHandler(WebSocketSession session) {
+ SubProtocolHandler handler;
+ String protocol = session.getAcceptedProtocol();
+ if (protocol != null) {
+ handler = this.protocolHandlers.get(protocol);
+ Assert.state(handler != null,
+ "No handler for sub-protocol '" + protocol + "', handlers = " + this.protocolHandlers);
+ }
+ else {
+ handler = this.defaultProtocolHandler;
+ Assert.state(handler != null,
+ "No sub-protocol was requested and a default sub-protocol handler was not configured");
+ }
+ return handler;
+ }
+
+ /**
+ * Resolves the {@code sessionId} for the given {@code message} using
+ * the {@link SubProtocolHandler#resolveSessionId} algorithm.
+ * @param message The message to resolve the {@code sessionId} from.
+ * @return The sessionId or {@code null}, if no one {@link SubProtocolHandler}
+ * can't resolve it against provided {@code message}.
+ */
+ public String resolveSessionId(Message> message) {
+ for (SubProtocolHandler handler : this.protocolHandlers.values()) {
+ String sessionId = handler.resolveSessionId(message);
+ if (sessionId != null) {
+ return sessionId;
+ }
+ }
+ if (this.defaultProtocolHandler != null) {
+ String sessionId = this.defaultProtocolHandler.resolveSessionId(message);
+ if (sessionId != null) {
+ return sessionId;
+ }
+ }
+ return null;
+ }
+
+ /**
+ * Return the {@link List} of sub-protocols from provided {@link SubProtocolHandler}.
+ * @return The the {@link List} of supported sub-protocols.
+ */
+ public List getSubProtocols() {
+ return new ArrayList(this.protocolHandlers.keySet());
+ }
+
+}
diff --git a/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/package-info.java b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/package-info.java
new file mode 100644
index 0000000000..01d1778645
--- /dev/null
+++ b/spring-integration-websocket/src/main/java/org/springframework/integration/websocket/support/package-info.java
@@ -0,0 +1,4 @@
+/**
+ * Provides support classes used from WebSocket components.
+ */
+package org.springframework.integration.websocket.support;
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/ClientWebSocketContainerTests.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/ClientWebSocketContainerTests.java
new file mode 100644
index 0000000000..08ed4037c0
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/ClientWebSocketContainerTests.java
@@ -0,0 +1,131 @@
+/*
+ * Copyright 2014 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;
+
+
+import static org.hamcrest.Matchers.instanceOf;
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertFalse;
+import static org.junit.Assert.assertNotNull;
+import static org.junit.Assert.assertThat;
+import static org.junit.Assert.assertTrue;
+import static org.junit.Assert.fail;
+
+import java.nio.ByteBuffer;
+import java.util.Collections;
+import java.util.List;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.TimeUnit;
+
+import org.junit.AfterClass;
+import org.junit.BeforeClass;
+import org.junit.Test;
+
+import org.springframework.web.socket.CloseStatus;
+import org.springframework.web.socket.PingMessage;
+import org.springframework.web.socket.PongMessage;
+import org.springframework.web.socket.WebSocketMessage;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.client.jetty.JettyWebSocketClient;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+public class ClientWebSocketContainerTests {
+
+ private final static JettyWebSocketTestServer server = new JettyWebSocketTestServer(TestServerConfig.class);
+
+ @BeforeClass
+ public static void setup() throws Exception {
+ server.afterPropertiesSet();
+ }
+
+ @AfterClass
+ public static void tearDown() throws Exception {
+ server.destroy();
+ }
+
+ @Test
+ public void testClientWebSocketContainer() throws Exception {
+ ClientWebSocketContainer container =
+ new ClientWebSocketContainer(new JettyWebSocketClient(), server.getWsBaseUrl() + "/ws/websocket");
+
+ TestWebSocketListener messageListener = new TestWebSocketListener();
+ container.setMessageListener(messageListener);
+
+ container.start();
+
+ WebSocketSession session = container.getSession(null);
+ assertNotNull(session);
+ assertTrue(session.isOpen());
+ assertEquals("v10.stomp", session.getAcceptedProtocol());
+
+ //TODO Jetty Server treats empty ByteBuffer as 'null' for PongMessage
+ session.sendMessage(new PingMessage(ByteBuffer.wrap("ping".getBytes())));
+
+ assertTrue(messageListener.messageLatch.await(10, TimeUnit.SECONDS));
+
+ container.stop();
+ try {
+ container.getSession(null);
+ fail("IllegalStateException expected");
+ }
+ catch (Exception e) {
+ assertThat(e, instanceOf(IllegalStateException.class));
+ assertEquals(e.getMessage(), "'clientSession' has not been established. Consider to 'start' this container.");
+ }
+
+ assertTrue(messageListener.sessionEndedLatch.await(10, TimeUnit.SECONDS));
+ assertFalse(session.isOpen());
+ assertTrue(messageListener.started);
+ assertThat(messageListener.message, instanceOf(PongMessage.class));
+ }
+
+ private class TestWebSocketListener implements WebSocketListener {
+
+ public boolean started;
+
+ public final CountDownLatch messageLatch = new CountDownLatch(1);
+
+ public WebSocketMessage> message;
+
+ public final CountDownLatch sessionEndedLatch = new CountDownLatch(1);
+
+ @Override
+ public void onMessage(WebSocketSession session, WebSocketMessage> message) throws Exception {
+ this.message = message;
+ this.messageLatch.countDown();
+ }
+
+ @Override
+ public void afterSessionStarted(WebSocketSession session) throws Exception {
+ this.started = true;
+ }
+
+ @Override
+ public void afterSessionEnded(WebSocketSession session, CloseStatus closeStatus) throws Exception {
+ sessionEndedLatch.countDown();
+ }
+
+ @Override
+ public List getSubProtocols() {
+ return Collections.singletonList("v10.stomp");
+ }
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/JettyWebSocketTestServer.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/JettyWebSocketTestServer.java
new file mode 100644
index 0000000000..5e3e2c684e
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/JettyWebSocketTestServer.java
@@ -0,0 +1,75 @@
+/*
+ * Copyright 2014 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;
+
+import org.eclipse.jetty.server.Server;
+import org.eclipse.jetty.servlet.ServletContextHandler;
+import org.eclipse.jetty.servlet.ServletHolder;
+
+import org.springframework.beans.factory.DisposableBean;
+import org.springframework.beans.factory.InitializingBean;
+import org.springframework.util.SocketUtils;
+import org.springframework.web.context.support.AnnotationConfigWebApplicationContext;
+import org.springframework.web.servlet.DispatcherServlet;
+
+/**
+ * @author Rossen Stoyanchev
+ * @since 4.1
+ */
+public class JettyWebSocketTestServer implements InitializingBean, DisposableBean {
+
+ private final Server jettyServer;
+
+ private final int port;
+
+ private final AnnotationConfigWebApplicationContext serverContext;
+
+ public JettyWebSocketTestServer(Class>... serverConfigs) {
+ this.port = SocketUtils.findAvailableTcpPort();
+ this.jettyServer = new Server(this.port);
+ this.serverContext = new AnnotationConfigWebApplicationContext();
+ this.serverContext.register(serverConfigs);
+ this.serverContext.refresh();
+
+ ServletContextHandler contextHandler = new ServletContextHandler();
+ ServletHolder servletHolder = new ServletHolder(new DispatcherServlet(this.serverContext));
+ contextHandler.addServlet(servletHolder, "/");
+ this.jettyServer.setHandler(contextHandler);
+ }
+
+ public AnnotationConfigWebApplicationContext getServerContext() {
+ return serverContext;
+ }
+
+ public String getWsBaseUrl() {
+ return "ws://localhost:" + this.port;
+ }
+
+ @Override
+ public void afterPropertiesSet() throws Exception {
+ this.jettyServer.start();
+ }
+
+ @Override
+ public void destroy() throws Exception {
+ if (this.jettyServer.isRunning()) {
+ this.jettyServer.setStopTimeout(0);
+ this.jettyServer.stop();
+ }
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/TestServerConfig.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/TestServerConfig.java
new file mode 100644
index 0000000000..7d5cba95dd
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/TestServerConfig.java
@@ -0,0 +1,84 @@
+/*
+ * Copyright 2014 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;
+
+import org.springframework.context.annotation.Bean;
+import org.springframework.context.annotation.Configuration;
+import org.springframework.integration.channel.AbstractSubscribableChannel;
+import org.springframework.integration.channel.DirectChannel;
+import org.springframework.integration.channel.QueueChannel;
+import org.springframework.messaging.Message;
+import org.springframework.messaging.MessageChannel;
+import org.springframework.messaging.simp.SimpMessageHeaderAccessor;
+import org.springframework.messaging.support.ChannelInterceptorAdapter;
+import org.springframework.messaging.support.MessageBuilder;
+import org.springframework.web.socket.WebSocketHandler;
+import org.springframework.web.socket.config.annotation.EnableWebSocket;
+import org.springframework.web.socket.config.annotation.WebSocketConfigurer;
+import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry;
+import org.springframework.web.socket.messaging.StompSubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolWebSocketHandler;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+@Configuration
+@EnableWebSocket
+public class TestServerConfig implements WebSocketConfigurer {
+
+ @Bean
+ public MessageChannel clientInboundChannel() {
+ return new QueueChannel();
+ }
+
+ @Bean
+ public AbstractSubscribableChannel clientOutboundChannel() {
+ DirectChannel directChannel = new DirectChannel();
+ directChannel.addInterceptor(new ChannelInterceptorAdapter() {
+ @Override
+ public Message> preSend(Message> message, MessageChannel channel) {
+ SimpMessageHeaderAccessor headers = SimpMessageHeaderAccessor.wrap(message);
+ headers.setLeaveMutable(true);
+ return MessageBuilder.createMessage(message.getPayload(), headers.getMessageHeaders());
+ }
+ });
+ return directChannel;
+ }
+
+ @Bean
+ public SubProtocolHandler stompSubProtocolHandler() {
+ return new StompSubProtocolHandler();
+ }
+
+ @Bean
+ public WebSocketHandler subProtocolWebSocketHandler() {
+ SubProtocolWebSocketHandler webSocketHandler =
+ new SubProtocolWebSocketHandler(clientInboundChannel(), clientOutboundChannel());
+ webSocketHandler.setDefaultProtocolHandler(stompSubProtocolHandler());
+ return webSocketHandler;
+ }
+
+ @Override
+ public void registerWebSocketHandlers(WebSocketHandlerRegistry registry) {
+ registry.addHandler(subProtocolWebSocketHandler(), "/ws")
+ .withSockJS();
+ }
+
+}
+
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
new file mode 100644
index 0000000000..a3441d42f3
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/StompIntegrationTests.java
@@ -0,0 +1,389 @@
+/*
+ * Copyright 2014 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.client;
+
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertNotNull;
+import static org.junit.Assert.assertTrue;
+
+import java.lang.annotation.ElementType;
+import java.lang.annotation.Retention;
+import java.lang.annotation.RetentionPolicy;
+import java.lang.annotation.Target;
+import java.nio.ByteBuffer;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.TimeUnit;
+
+import org.junit.Test;
+import org.junit.runner.RunWith;
+
+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.annotation.Bean;
+import org.springframework.context.annotation.ComponentScan;
+import org.springframework.context.annotation.Configuration;
+import org.springframework.expression.ExpressionParser;
+import org.springframework.expression.spel.standard.SpelExpressionParser;
+import org.springframework.integration.annotation.Gateway;
+import org.springframework.integration.annotation.IntegrationComponentScan;
+import org.springframework.integration.annotation.MessagingGateway;
+import org.springframework.integration.annotation.ServiceActivator;
+import org.springframework.integration.annotation.Transformer;
+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.transformer.ExpressionEvaluatingTransformer;
+import org.springframework.integration.websocket.ClientWebSocketContainer;
+import org.springframework.integration.websocket.IntegrationWebSocketContainer;
+import org.springframework.integration.websocket.JettyWebSocketTestServer;
+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.handler.annotation.MessageExceptionHandler;
+import org.springframework.messaging.handler.annotation.MessageMapping;
+import org.springframework.messaging.simp.annotation.SendToUser;
+import org.springframework.messaging.simp.annotation.SubscribeMapping;
+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.MessageBuilder;
+import org.springframework.stereotype.Controller;
+import org.springframework.test.annotation.DirtiesContext;
+import org.springframework.test.context.ContextConfiguration;
+import org.springframework.test.context.junit4.SpringJUnit4ClassRunner;
+import org.springframework.web.socket.client.jetty.JettyWebSocketClient;
+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.StompSubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+import org.springframework.web.socket.server.jetty.JettyRequestUpgradeStrategy;
+import org.springframework.web.socket.server.support.DefaultHandshakeHandler;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+@ContextConfiguration
+@RunWith(SpringJUnit4ClassRunner.class)
+@DirtiesContext
+public class StompIntegrationTests {
+
+ @Value("#{server.serverContext}")
+ private ApplicationContext serverContext;
+
+ @Autowired
+ @Qualifier("webSocketOutputChannel")
+ private MessageChannel webSocketOutputChannel;
+
+ @Autowired
+ @Qualifier("webSocketInputChannel")
+ private QueueChannel webSocketInputChannel;
+
+ @Test
+ public void sendMessageToController() throws Exception {
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SEND);
+ headers.setSubscriptionId("sub1");
+ headers.setDestination("/app/simple");
+ Message message = MessageBuilder.withPayload("foo").setHeaders(headers).build();
+
+ this.webSocketOutputChannel.send(message);
+
+ SimpleController controller = this.serverContext.getBean(SimpleController.class);
+ assertTrue(controller.latch.await(10, TimeUnit.SECONDS));
+ }
+
+ @Test
+ public void sendMessageToControllerAndReceiveReplyViaTopic() throws Exception {
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/topic/increment");
+ Message message = MessageBuilder.withPayload(ByteBuffer.allocate(0).array())
+ .setHeaders(headers)
+ .build();
+
+ headers = StompHeaderAccessor.create(StompCommand.SEND);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/app/increment");
+ Message message2 = MessageBuilder.withPayload(5).setHeaders(headers).build();
+
+ this.webSocketOutputChannel.send(message);
+ this.webSocketOutputChannel.send(message2);
+
+ Message> receive = webSocketInputChannel.receive(1000);
+ assertNotNull(receive);
+ assertEquals("6", receive.getPayload());
+ }
+
+ @Test
+ public void sendMessageToBrokerAndReceiveReplyViaTopic() throws Exception {
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/topic/foo");
+ Message message = MessageBuilder.withPayload(ByteBuffer.allocate(0).array())
+ .setHeaders(headers)
+ .build();
+
+ headers = StompHeaderAccessor.create(StompCommand.SEND);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/topic/foo");
+ Message message2 = MessageBuilder.withPayload(10).setHeaders(headers).build();
+
+ this.webSocketOutputChannel.send(message);
+ this.webSocketOutputChannel.send(message2);
+
+ Message> receive = webSocketInputChannel.receive(1000);
+ assertNotNull(receive);
+ assertEquals("10", receive.getPayload());
+ }
+
+ @Test
+ public void sendSubscribeToControllerAndReceiveReply() throws Exception {
+
+ String destHeader = "/app/number";
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination(destHeader);
+ Message message = MessageBuilder.withPayload(ByteBuffer.allocate(0).array())
+ .setHeaders(headers)
+ .build();
+
+ this.webSocketOutputChannel.send(message);
+
+ Message> receive = webSocketInputChannel.receive(10000);
+ assertNotNull(receive);
+
+ StompHeaderAccessor stompHeaderAccessor = StompHeaderAccessor.wrap(receive);
+
+ assertEquals("Expected STOMP destination=/app/number, got " + stompHeaderAccessor,
+ destHeader, stompHeaderAccessor.getDestination());
+
+ Object payload = receive.getPayload();
+
+ assertEquals("Expected STOMP Payload=42, got " + payload, "42", payload);
+ }
+
+ @Test
+ public void handleExceptionAndSendToUser() throws Exception {
+ String destHeader = "/user/queue/error";
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination(destHeader);
+ Message message = MessageBuilder.withPayload(ByteBuffer.allocate(0).array())
+ .setHeaders(headers)
+ .build();
+
+ headers = StompHeaderAccessor.create(StompCommand.SEND);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/app/exception");
+ Message message2 = MessageBuilder.withPayload("foo").setHeaders(headers).build();
+
+ this.webSocketOutputChannel.send(message);
+ this.webSocketOutputChannel.send(message2);
+
+
+ Message> receive = webSocketInputChannel.receive(10000);
+ assertNotNull(receive);
+
+ StompHeaderAccessor stompHeaderAccessor = StompHeaderAccessor.wrap(receive);
+
+ assertEquals("Expected STOMP destination=/user/queue/error, got " + stompHeaderAccessor,
+ destHeader, stompHeaderAccessor.getDestination());
+
+ assertEquals("Got error: Bad input", receive.getPayload());
+ }
+
+ @Test
+ public void sendMessageToGateway() throws Exception {
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/user/queue/answer");
+ Message message = MessageBuilder.withPayload(ByteBuffer.allocate(0).array())
+ .setHeaders(headers)
+ .build();
+
+ headers = StompHeaderAccessor.create(StompCommand.SEND);
+ headers.setSubscriptionId("subs1");
+ headers.setDestination("/app/greeting");
+ Message message2 = MessageBuilder.withPayload("Bob").setHeaders(headers).build();
+
+ this.webSocketOutputChannel.send(message);
+ this.webSocketOutputChannel.send(message2);
+
+ Message> receive = webSocketInputChannel.receive(5000);
+ assertNotNull(receive);
+ assertEquals("Hello Bob", receive.getPayload());
+ }
+
+
+ @Configuration
+ @EnableIntegration
+ public static class ContextConfiguration {
+
+ @Bean
+ public JettyWebSocketTestServer server() {
+ return new JettyWebSocketTestServer(ServerConfig.class);
+ }
+
+ @Bean
+ public IntegrationWebSocketContainer clientWebSocketContainer() {
+ return new ClientWebSocketContainer(new JettyWebSocketClient(), server().getWsBaseUrl() + "/ws/websocket");
+ }
+
+ @Bean
+ public SubProtocolHandler stompSubProtocolHandler() {
+ return new StompSubProtocolHandler();
+ }
+
+ @Bean
+ public MessageChannel webSocketInputChannel() {
+ return new QueueChannel();
+ }
+
+ @Bean
+ public MessageChannel webSocketOutputChannel() {
+ return new DirectChannel();
+ }
+
+ @Bean
+ public MessageProducer webSocketInboundChannelAdapter() {
+ WebSocketInboundChannelAdapter webSocketInboundChannelAdapter =
+ new WebSocketInboundChannelAdapter(clientWebSocketContainer(),
+ new SubProtocolHandlerRegistry(stompSubProtocolHandler()));
+ webSocketInboundChannelAdapter.setOutputChannel(webSocketInputChannel());
+ return webSocketInboundChannelAdapter;
+ }
+
+ @Bean
+ @ServiceActivator(inputChannel = "webSocketOutputChannel")
+ public MessageHandler webSocketOutboundMessageHandler() {
+ return new WebSocketOutboundMessageHandler(clientWebSocketContainer(),
+ new SubProtocolHandlerRegistry(stompSubProtocolHandler()));
+ }
+
+ }
+
+ // WebSocket Server part
+
+ @Target({ElementType.TYPE})
+ @Retention(RetentionPolicy.RUNTIME)
+ @Controller
+ private @interface IntegrationTestController {
+ }
+
+ @IntegrationTestController
+ static class SimpleController {
+
+ private CountDownLatch latch = new CountDownLatch(1);
+
+ @MessageMapping(value = "/simple")
+ public void handle() {
+ this.latch.countDown();
+ }
+
+ @MessageMapping(value = "/exception")
+ public void handleWithError() {
+ throw new IllegalArgumentException("Bad input");
+ }
+
+ @MessageExceptionHandler
+ @SendToUser("/queue/error")
+ public String handleException(IllegalArgumentException ex) {
+ return "Got error: " + ex.getMessage();
+ }
+ }
+
+ @IntegrationTestController
+ static class IncrementController {
+
+ @MessageMapping(value = "/increment")
+ public int handle(int i) {
+ return i + 1;
+ }
+
+ @SubscribeMapping("/number")
+ public int number() {
+ return 42;
+ }
+ }
+
+
+ @MessagingGateway
+ @Controller
+ static interface WebSocketGateway {
+
+ @MessageMapping("/greeting")
+ @SendToUser("/queue/answer")
+ @Gateway(requestChannel = "greetingChannel")
+ String greeting(String payload);
+
+ }
+
+ @Configuration
+ @EnableWebSocketMessageBroker
+ @EnableIntegration
+ @ComponentScan(
+ basePackageClasses = StompIntegrationTests.class,
+ useDefaultFilters = false,
+ includeFilters = @ComponentScan.Filter(IntegrationTestController.class))
+ @IntegrationComponentScan
+ static class ServerConfig extends AbstractWebSocketMessageBrokerConfigurer {
+
+ private static final ExpressionParser expressionParser = new SpelExpressionParser();
+
+ @Bean
+ public MessageChannel greetingChannel() {
+ return new DirectChannel();
+ }
+
+ @Bean
+ @Transformer(inputChannel = "greetingChannel")
+ public ExpressionEvaluatingTransformer greetingTransformer() {
+ return new ExpressionEvaluatingTransformer(expressionParser.parseExpression("'Hello ' + payload"));
+ }
+
+ @Bean
+ public DefaultHandshakeHandler handshakeHandler() {
+ return new DefaultHandshakeHandler(new JettyRequestUpgradeStrategy());
+ }
+
+ @Override
+ public void registerStompEndpoints(StompEndpointRegistry registry) {
+ registry.addEndpoint("/ws").setHandshakeHandler(handshakeHandler()).withSockJS();
+ }
+
+ @Override
+ public void configureMessageBroker(MessageBrokerRegistry configurer) {
+ configurer.setApplicationDestinationPrefixes("/app");
+ configurer.enableSimpleBroker("/topic", "/queue");
+ }
+
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/WebSocketClientTests.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/WebSocketClientTests.java
new file mode 100644
index 0000000000..74dd8b08ec
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/client/WebSocketClientTests.java
@@ -0,0 +1,189 @@
+/*
+ * Copyright 2014 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.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 java.util.Collections;
+
+import org.junit.Test;
+import org.junit.runner.RunWith;
+
+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.annotation.Bean;
+import org.springframework.context.annotation.Configuration;
+import org.springframework.integration.annotation.Poller;
+import org.springframework.integration.annotation.ServiceActivator;
+import org.springframework.integration.annotation.Transformer;
+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.transformer.ObjectToStringTransformer;
+import org.springframework.integration.websocket.ClientWebSocketContainer;
+import org.springframework.integration.websocket.IntegrationWebSocketContainer;
+import org.springframework.integration.websocket.JettyWebSocketTestServer;
+import org.springframework.integration.websocket.TestServerConfig;
+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.simp.stomp.StompCommand;
+import org.springframework.messaging.simp.stomp.StompHeaderAccessor;
+import org.springframework.messaging.support.GenericMessage;
+import org.springframework.stereotype.Component;
+import org.springframework.test.annotation.DirtiesContext;
+import org.springframework.test.context.ContextConfiguration;
+import org.springframework.test.context.junit4.SpringJUnit4ClassRunner;
+import org.springframework.web.socket.client.WebSocketClient;
+import org.springframework.web.socket.client.jetty.JettyWebSocketClient;
+import org.springframework.web.socket.messaging.StompSubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+import org.springframework.web.socket.sockjs.client.SockJsClient;
+import org.springframework.web.socket.sockjs.client.Transport;
+import org.springframework.web.socket.sockjs.client.WebSocketTransport;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+@ContextConfiguration
+@RunWith(SpringJUnit4ClassRunner.class)
+@DirtiesContext
+public class WebSocketClientTests {
+
+ @Value("#{server.serverContext}")
+ private ApplicationContext serverContext;
+
+ @Autowired
+ @Qualifier("webSocketOutputChannel")
+ private MessageChannel webSocketOutputChannel;
+
+ @Autowired
+ @Qualifier("webSocketInputChannel")
+ private QueueChannel webSocketInputChannel;
+
+ @Test
+ public void testWebSocketOutboundMessageHandler() throws Exception {
+ this.webSocketOutputChannel.send(new GenericMessage("Spring"));
+
+ Message> received = this.webSocketInputChannel.receive(10000);
+ assertNotNull(received);
+ StompHeaderAccessor stompHeaderAccessor = StompHeaderAccessor.wrap(received);
+ assertEquals(StompCommand.MESSAGE.getMessageType(), stompHeaderAccessor.getMessageType());
+
+ Object receivedPayload = received.getPayload();
+ assertThat(receivedPayload, instanceOf(String.class));
+ assertEquals("Hello Spring", receivedPayload);
+ }
+
+ @Configuration
+ @EnableIntegration
+ public static class ContextConfiguration {
+
+ @Bean
+ public JettyWebSocketTestServer server() {
+ return new JettyWebSocketTestServer(ServerFlowConfig.class);
+ }
+
+ @Bean
+ public WebSocketClient webSocketClient() {
+ return new SockJsClient(Collections.singletonList(new WebSocketTransport(new JettyWebSocketClient())));
+ }
+
+ @Bean
+ public IntegrationWebSocketContainer clientWebSocketContainer() {
+ return new ClientWebSocketContainer(webSocketClient(), server().getWsBaseUrl() + "/ws");
+ }
+
+ @Bean
+ public SubProtocolHandler stompSubProtocolHandler() {
+ return new StompSubProtocolHandler();
+ }
+
+ @Bean
+ public MessageChannel webSocketInputChannel() {
+ return new QueueChannel();
+ }
+
+ @Bean
+ public MessageChannel webSocketOutputChannel() {
+ return new DirectChannel();
+ }
+
+ @Bean
+ public MessageProducer webSocketInboundChannelAdapter() {
+ WebSocketInboundChannelAdapter webSocketInboundChannelAdapter =
+ new WebSocketInboundChannelAdapter(clientWebSocketContainer(),
+ new SubProtocolHandlerRegistry(stompSubProtocolHandler()));
+ webSocketInboundChannelAdapter.setOutputChannel(webSocketInputChannel());
+ return webSocketInboundChannelAdapter;
+ }
+
+ @Bean
+ @ServiceActivator(inputChannel = "webSocketOutputChannel")
+ public MessageHandler webSocketOutboundMessageHandler() {
+ return new WebSocketOutboundMessageHandler(clientWebSocketContainer(),
+ new SubProtocolHandlerRegistry(stompSubProtocolHandler()));
+ }
+
+ }
+
+ // WebSocket Server part
+
+ @Configuration
+ @EnableIntegration
+ static class ServerFlowConfig extends TestServerConfig {
+
+ @Bean
+ @Transformer(inputChannel = "clientInboundChannel", outputChannel = "serviceChannel",
+ poller = @Poller(fixedDelay = "100", maxMessagesPerPoll = "1"))
+ public org.springframework.integration.transformer.Transformer objectToStringTransformer() {
+ return new ObjectToStringTransformer();
+ }
+
+ @Bean
+ public DirectChannel serviceChannel() {
+ return new DirectChannel();
+ }
+
+ @Bean
+ public TestService service() {
+ return new TestService();
+ }
+
+ @Component
+ public static class TestService {
+
+ @ServiceActivator(inputChannel = "serviceChannel", outputChannel = "clientOutboundChannel")
+ public byte[] handle(String payload) {
+ return ("Hello " + payload).getBytes();
+ }
+
+ }
+
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapterTests.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapterTests.java
new file mode 100644
index 0000000000..cd794fef25
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/inbound/WebSocketInboundChannelAdapterTests.java
@@ -0,0 +1,162 @@
+/*
+ * Copyright 2014 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.inbound;
+
+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.nio.ByteBuffer;
+import java.util.Collections;
+import java.util.Map;
+
+import org.junit.Test;
+import org.junit.runner.RunWith;
+
+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.annotation.Bean;
+import org.springframework.context.annotation.Configuration;
+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.test.util.TestUtils;
+import org.springframework.integration.websocket.ClientWebSocketContainer;
+import org.springframework.integration.websocket.IntegrationWebSocketContainer;
+import org.springframework.integration.websocket.JettyWebSocketTestServer;
+import org.springframework.integration.websocket.TestServerConfig;
+import org.springframework.integration.websocket.support.SubProtocolHandlerRegistry;
+import org.springframework.messaging.Message;
+import org.springframework.messaging.MessageChannel;
+import org.springframework.messaging.simp.stomp.StompCommand;
+import org.springframework.messaging.simp.stomp.StompHeaderAccessor;
+import org.springframework.messaging.support.MessageBuilder;
+import org.springframework.test.annotation.DirtiesContext;
+import org.springframework.test.context.ContextConfiguration;
+import org.springframework.test.context.junit4.SpringJUnit4ClassRunner;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.client.WebSocketClient;
+import org.springframework.web.socket.client.jetty.JettyWebSocketClient;
+import org.springframework.web.socket.messaging.StompSubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolWebSocketHandler;
+import org.springframework.web.socket.sockjs.client.SockJsClient;
+import org.springframework.web.socket.sockjs.client.Transport;
+import org.springframework.web.socket.sockjs.client.WebSocketTransport;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+@ContextConfiguration
+@RunWith(SpringJUnit4ClassRunner.class)
+@DirtiesContext
+public class WebSocketInboundChannelAdapterTests {
+
+ @Value("#{server.serverContext.getBean('subProtocolWebSocketHandler')}")
+ private SubProtocolWebSocketHandler subProtocolWebSocketHandler;
+
+ @Value("#{server.serverContext.getBean('clientOutboundChannel')}")
+ private DirectChannel clientOutboundChannel;
+
+ @Autowired
+ IntegrationWebSocketContainer clientWebSocketContainer;
+
+ @Autowired
+ @Qualifier("webSocketChannel")
+ private QueueChannel webSocketChannel;
+
+ @Test
+ @SuppressWarnings("unchecked")
+ public void testWebSocketInboundChannelAdapter() throws Exception {
+ WebSocketSession session = clientWebSocketContainer.getSession(null);
+ assertNotNull(session);
+ assertTrue(session.isOpen());
+ assertEquals("v10.stomp", session.getAcceptedProtocol());
+
+ Map sessions =
+ TestUtils.getPropertyValue(this.subProtocolWebSocketHandler, "sessions", Map.class);
+
+
+ assertEquals(1, sessions.size());
+
+ String sessionId = sessions.keySet().iterator().next();
+
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.MESSAGE);
+ headers.setLeaveMutable(true);
+ headers.setSessionId(sessionId);
+ Message message = MessageBuilder.createMessage(ByteBuffer.allocate(0).array(), headers.getMessageHeaders());
+
+ this.clientOutboundChannel.send(message);
+
+ Message> received = this.webSocketChannel.receive(10000);
+ assertNotNull(received);
+
+ StompHeaderAccessor receivedHeaders = StompHeaderAccessor.wrap(received);
+ assertEquals(StompCommand.MESSAGE, receivedHeaders.getCommand());
+
+ Object receivedPayload = received.getPayload();
+ assertThat(receivedPayload, instanceOf(String.class));
+ assertEquals("", receivedPayload);
+
+ }
+
+ @Configuration
+ @EnableIntegration
+ public static class ContextConfiguration {
+
+ @Bean
+ public JettyWebSocketTestServer server() {
+ return new JettyWebSocketTestServer(TestServerConfig.class);
+ }
+
+ @Bean
+ public WebSocketClient webSocketClient() {
+ return new SockJsClient(Collections.singletonList(new WebSocketTransport(new JettyWebSocketClient())));
+ }
+
+ @Bean
+ public IntegrationWebSocketContainer clientWebSocketContainer() {
+ return new ClientWebSocketContainer(webSocketClient(), server().getWsBaseUrl() + "/ws");
+ }
+
+ @Bean
+ public SubProtocolHandler stompSubProtocolHandler() {
+ return new StompSubProtocolHandler();
+ }
+
+ @Bean
+ public MessageChannel webSocketChannel() {
+ return new QueueChannel();
+ }
+
+ @Bean
+ public MessageProducer webSocketInboundChannelAdapter() {
+ WebSocketInboundChannelAdapter webSocketInboundChannelAdapter =
+ new WebSocketInboundChannelAdapter(clientWebSocketContainer(),
+ new SubProtocolHandlerRegistry(stompSubProtocolHandler()));
+ webSocketInboundChannelAdapter.setOutputChannel(webSocketChannel());
+ return webSocketInboundChannelAdapter;
+ }
+
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/outbound/WebSocketOutboundMessageHandlerTests.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/outbound/WebSocketOutboundMessageHandlerTests.java
new file mode 100644
index 0000000000..cb6507c3bf
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/outbound/WebSocketOutboundMessageHandlerTests.java
@@ -0,0 +1,134 @@
+/*
+ * Copyright 2014 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.outbound;
+
+import static org.hamcrest.Matchers.instanceOf;
+import static org.junit.Assert.assertArrayEquals;
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertNotNull;
+import static org.junit.Assert.assertThat;
+
+import java.util.Collections;
+
+import org.junit.Test;
+import org.junit.runner.RunWith;
+
+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.annotation.Bean;
+import org.springframework.context.annotation.Configuration;
+import org.springframework.integration.channel.QueueChannel;
+import org.springframework.integration.config.EnableIntegration;
+import org.springframework.integration.websocket.ClientWebSocketContainer;
+import org.springframework.integration.websocket.IntegrationWebSocketContainer;
+import org.springframework.integration.websocket.JettyWebSocketTestServer;
+import org.springframework.integration.websocket.TestServerConfig;
+import org.springframework.integration.websocket.support.SubProtocolHandlerRegistry;
+import org.springframework.messaging.Message;
+import org.springframework.messaging.MessageHandler;
+import org.springframework.messaging.simp.stomp.StompCommand;
+import org.springframework.messaging.simp.stomp.StompHeaderAccessor;
+import org.springframework.messaging.support.MessageBuilder;
+import org.springframework.test.annotation.DirtiesContext;
+import org.springframework.test.context.ContextConfiguration;
+import org.springframework.test.context.junit4.SpringJUnit4ClassRunner;
+import org.springframework.web.socket.client.WebSocketClient;
+import org.springframework.web.socket.client.jetty.JettyWebSocketClient;
+import org.springframework.web.socket.messaging.StompSubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+import org.springframework.web.socket.sockjs.client.SockJsClient;
+import org.springframework.web.socket.sockjs.client.Transport;
+import org.springframework.web.socket.sockjs.client.WebSocketTransport;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+@ContextConfiguration
+@RunWith(SpringJUnit4ClassRunner.class)
+@DirtiesContext
+public class WebSocketOutboundMessageHandlerTests {
+
+ @Autowired
+ @Qualifier("webSocketOutboundMessageHandler")
+ private MessageHandler messageHandler;
+
+ @Value("#{server.serverContext.getBean('clientInboundChannel')}")
+ private QueueChannel clientInboundChannel;
+
+ @Test
+ public void testWebSocketOutboundMessageHandler() throws Exception {
+ StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SEND);
+ headers.setMessageId("mess0");
+ headers.setSubscriptionId("sub0");
+ headers.setDestination("/foo");
+ String payload = "Hello World";
+ Message message = MessageBuilder.withPayload(payload).setHeaders(headers).build();
+
+ this.messageHandler.handleMessage(message);
+
+ Message> received = this.clientInboundChannel.receive(10000);
+ assertNotNull(received);
+
+ StompHeaderAccessor receivedHeaders = StompHeaderAccessor.wrap(received);
+ assertEquals("mess0", receivedHeaders.getMessageId());
+ assertEquals("sub0", receivedHeaders.getSubscriptionId());
+ assertEquals("/foo", receivedHeaders.getDestination());
+
+ Object receivedPayload = received.getPayload();
+ assertThat(receivedPayload, instanceOf(byte[].class));
+ assertArrayEquals((byte[]) receivedPayload, payload.getBytes());
+ }
+
+
+ @Configuration
+ @EnableIntegration
+ public static class ContextConfiguration {
+
+ @Bean
+ public JettyWebSocketTestServer server() {
+ return new JettyWebSocketTestServer(TestServerConfig.class);
+ }
+
+ @Bean
+ public WebSocketClient webSocketClient() {
+ return new SockJsClient(Collections.singletonList(new WebSocketTransport(new JettyWebSocketClient())));
+ }
+
+ @Bean
+ public IntegrationWebSocketContainer clientWebSocketContainer() {
+ ClientWebSocketContainer container =
+ new ClientWebSocketContainer(webSocketClient(), server().getWsBaseUrl() + "/ws");
+ container.setAutoStartup(true);
+ return container;
+ }
+
+ @Bean
+ public SubProtocolHandler stompSubProtocolHandler() {
+ return new StompSubProtocolHandler();
+ }
+
+ @Bean
+ public MessageHandler webSocketOutboundMessageHandler() {
+ return new WebSocketOutboundMessageHandler(clientWebSocketContainer(),
+ new SubProtocolHandlerRegistry(stompSubProtocolHandler()));
+ }
+
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/support/SubProtocolHandlerRegistryTests.java b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/support/SubProtocolHandlerRegistryTests.java
new file mode 100644
index 0000000000..ffca3afa13
--- /dev/null
+++ b/spring-integration-websocket/src/test/java/org/springframework/integration/websocket/support/SubProtocolHandlerRegistryTests.java
@@ -0,0 +1,126 @@
+/*
+ * Copyright 2014 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.support;
+
+import static org.hamcrest.Matchers.containsString;
+import static org.hamcrest.Matchers.instanceOf;
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertNotNull;
+import static org.junit.Assert.assertNull;
+import static org.junit.Assert.assertSame;
+import static org.junit.Assert.assertThat;
+import static org.junit.Assert.fail;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.spy;
+import static org.mockito.Mockito.when;
+
+import java.util.Collections;
+
+import org.junit.Test;
+
+import org.springframework.messaging.Message;
+import org.springframework.messaging.simp.SimpMessageHeaderAccessor;
+import org.springframework.messaging.support.MessageBuilder;
+import org.springframework.web.socket.WebSocketSession;
+import org.springframework.web.socket.messaging.StompSubProtocolHandler;
+import org.springframework.web.socket.messaging.SubProtocolHandler;
+
+/**
+ * @author Artem Bilan
+ * @since 4.1
+ */
+public class SubProtocolHandlerRegistryTests {
+
+ @Test
+ public void testProtocolHandlers() {
+ SubProtocolHandler defaultProtocolHandler = mock(SubProtocolHandler.class);
+ SubProtocolHandlerRegistry subProtocolHandlerRegistry =
+ new SubProtocolHandlerRegistry(
+ Collections.singletonList(new StompSubProtocolHandler()),
+ defaultProtocolHandler);
+ WebSocketSession session = mock(WebSocketSession.class);
+ when(session.getAcceptedProtocol()).thenReturn("v10.stomp", (String) null);
+ SubProtocolHandler protocolHandler = subProtocolHandlerRegistry.findProtocolHandler(session);
+ assertNotNull(protocolHandler);
+ assertThat(protocolHandler, instanceOf(StompSubProtocolHandler.class));
+ protocolHandler = subProtocolHandlerRegistry.findProtocolHandler(session);
+ assertNotNull(protocolHandler);
+ assertSame(protocolHandler, defaultProtocolHandler);
+
+ assertEquals(subProtocolHandlerRegistry.getSubProtocols(), new StompSubProtocolHandler().getSupportedProtocols());
+ }
+
+ @Test
+ public void testSingleHandler() {
+ SubProtocolHandler testProtocolHandler = spy(new StompSubProtocolHandler());
+ when(testProtocolHandler.getSupportedProtocols()).thenReturn(Collections.singletonList("foo"));
+ SubProtocolHandlerRegistry subProtocolHandlerRegistry =
+ new SubProtocolHandlerRegistry(Collections.singletonList(testProtocolHandler));
+ WebSocketSession session = mock(WebSocketSession.class);
+ when(session.getAcceptedProtocol()).thenReturn("foo", (String) null);
+ SubProtocolHandler protocolHandler = subProtocolHandlerRegistry.findProtocolHandler(session);
+ assertNotNull(protocolHandler);
+ assertSame(protocolHandler, testProtocolHandler);
+
+ protocolHandler = subProtocolHandlerRegistry.findProtocolHandler(session);
+ assertNotNull(protocolHandler);
+ assertSame(protocolHandler, testProtocolHandler);
+ }
+
+ @Test
+ public void testHandlerSelection() {
+ SubProtocolHandler testProtocolHandler = new StompSubProtocolHandler();
+ SubProtocolHandlerRegistry subProtocolHandlerRegistry =
+ new SubProtocolHandlerRegistry(testProtocolHandler);
+ WebSocketSession session = mock(WebSocketSession.class);
+ when(session.getAcceptedProtocol()).thenReturn("foo", (String) null);
+
+ try {
+ subProtocolHandlerRegistry.findProtocolHandler(session);
+ fail("IllegalStateException expected");
+ }
+ catch (Exception e) {
+ assertThat(e, instanceOf(IllegalStateException.class));
+ assertThat(e.getMessage(), containsString("No handler for sub-protocol 'foo'"));
+ }
+
+ SubProtocolHandler protocolHandler = subProtocolHandlerRegistry.findProtocolHandler(session);
+ assertNotNull(protocolHandler);
+ assertSame(protocolHandler, testProtocolHandler);
+ }
+
+ @Test
+ public void testResolveSessionId() {
+ SubProtocolHandlerRegistry subProtocolHandlerRegistry =
+ new SubProtocolHandlerRegistry(new StompSubProtocolHandler());
+
+ Message message = MessageBuilder.withPayload("foo")
+ .setHeader(SimpMessageHeaderAccessor.SESSION_ID_HEADER, "TEST_SESSION")
+ .build();
+
+ String sessionId = subProtocolHandlerRegistry.resolveSessionId(message);
+ assertEquals(sessionId, "TEST_SESSION");
+
+ message = MessageBuilder.withPayload("foo")
+ .setHeader("MY_SESSION_ID", "TEST_SESSION")
+ .build();
+
+ sessionId = subProtocolHandlerRegistry.resolveSessionId(message);
+ assertNull(sessionId);
+ }
+
+}
diff --git a/spring-integration-websocket/src/test/resources/log4j.properties b/spring-integration-websocket/src/test/resources/log4j.properties
new file mode 100644
index 0000000000..7fcd10a721
--- /dev/null
+++ b/spring-integration-websocket/src/test/resources/log4j.properties
@@ -0,0 +1,8 @@
+log4j.rootCategory=WARN, stdout
+
+log4j.appender.stdout=org.apache.log4j.ConsoleAppender
+log4j.appender.stdout.layout=org.apache.log4j.PatternLayout
+log4j.appender.stdout.layout.ConversionPattern=%d %c{1} [%t] : %m%n
+
+log4j.category.org.springframework.integration=WARN
+log4j.category.org.springframework.integration.websocket=WARN