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