From 3f9da6f4809431a323a627c719cc28987fb6945e Mon Sep 17 00:00:00 2001 From: Rossen Stoyanchev Date: Tue, 18 Jun 2013 20:35:15 -0400 Subject: [PATCH] Add generic parameters to MessageHandler impls --- .../web/messaging/PubSubChannelRegistry.java | 3 +- .../service/AbstractPubSubMessageHandler.java | 21 +++++---- .../service/ReactorPubSubMessageHandler.java | 27 +++++------ .../AnnotationPubSubMessageHandler.java | 27 ++++++----- .../service/method/ArgumentResolver.java | 5 +- .../method/ArgumentResolverComposite.java | 28 +++++------ .../method/InvocableMessageHandlerMethod.java | 11 +++-- .../method/MessageBodyArgumentResolver.java | 5 +- .../MessageChannelArgumentResolver.java | 11 +++-- .../method/MessageReturnValueHandler.java | 22 +++++---- .../service/method/ReturnValueHandler.java | 5 +- .../method/ReturnValueHandlerComposite.java | 20 ++++---- .../stomp/support/StompMessageConverter.java | 27 +++++++++-- .../StompRelayPubSubMessageHandler.java | 47 ++++++++++--------- .../stomp/support/StompWebSocketHandler.java | 42 +++++++++-------- .../AbstractPubSubChannelRegistry.java | 4 +- .../support/ReactorMessageChannel.java | 13 +++-- .../support/SessionMessageChannel.java | 17 +++---- 18 files changed, 185 insertions(+), 150 deletions(-) diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubChannelRegistry.java b/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubChannelRegistry.java index 8c1e09f522..8cd4143343 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubChannelRegistry.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/PubSubChannelRegistry.java @@ -25,7 +25,8 @@ import org.springframework.messaging.SubscribableChannel; * @author Rossen Stoyanchev * @since 4.0 */ -public interface PubSubChannelRegistry, H extends MessageHandler> { +@SuppressWarnings("rawtypes") +public interface PubSubChannelRegistry> { SubscribableChannel getClientInputChannel(); diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java index a6565b1e1f..ca4c4939b0 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/AbstractPubSubMessageHandler.java @@ -37,7 +37,8 @@ import org.springframework.web.messaging.PubSubHeaders; * @author Rossen Stoyanchev * @since 4.0 */ -public abstract class AbstractPubSubMessageHandler implements MessageHandler> { +@SuppressWarnings("rawtypes") +public abstract class AbstractPubSubMessageHandler implements MessageHandler { protected final Log logger = LogFactory.getLog(getClass()); @@ -67,7 +68,7 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler getSupportedMessageTypes(); - protected boolean canHandle(Message message, MessageType messageType) { + protected boolean canHandle(M message, MessageType messageType) { if (!CollectionUtils.isEmpty(getSupportedMessageTypes())) { if (!getSupportedMessageTypes().contains(messageType)) { @@ -78,7 +79,7 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler message) { + protected boolean isDestinationAllowed(M message) { PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); String destination = headers.getDestination(); @@ -114,7 +115,7 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler message) throws MessagingException { + public final void handleMessage(M message) throws MessagingException { PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); MessageType messageType = headers.getMessageType(); @@ -143,22 +144,22 @@ public abstract class AbstractPubSubMessageHandler implements MessageHandler message) { + protected void handleConnect(M message) { } - protected void handlePublish(Message message) { + protected void handlePublish(M message) { } - protected void handleSubscribe(Message message) { + protected void handleSubscribe(M message) { } - protected void handleUnsubscribe(Message message) { + protected void handleUnsubscribe(M message) { } - protected void handleDisconnect(Message message) { + protected void handleDisconnect(M message) { } - protected void handleOther(Message message) { + protected void handleOther(M message) { } } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/ReactorPubSubMessageHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/ReactorPubSubMessageHandler.java index 57717806fa..e80f12bc33 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/ReactorPubSubMessageHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/ReactorPubSubMessageHandler.java @@ -44,9 +44,10 @@ import reactor.fn.selector.ObjectSelector; * @author Rossen Stoyanchev * @since 4.0 */ -public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { +@SuppressWarnings("rawtypes") +public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { - private MessageChannel> clientChannel; + private MessageChannel clientChannel; private final Reactor reactor; @@ -55,7 +56,7 @@ public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { private Map>> subscriptionsBySession = new ConcurrentHashMap>>(); - public ReactorPubSubMessageHandler(PubSubChannelRegistry registry, Reactor reactor) { + public ReactorPubSubMessageHandler(PubSubChannelRegistry registry, Reactor reactor) { Assert.notNull(reactor, "reactor is required"); this.clientChannel = registry.getClientOutputChannel(); this.reactor = reactor; @@ -72,7 +73,7 @@ public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { } @Override - public void handleSubscribe(Message message) { + public void handleSubscribe(M message) { if (logger.isDebugEnabled()) { logger.debug("Subscribe " + message); @@ -99,7 +100,7 @@ public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { } @Override - public void handlePublish(Message message) { + public void handlePublish(M message) { if (logger.isDebugEnabled()) { logger.debug("Message received: " + message); @@ -109,9 +110,10 @@ public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { // Convert to byte[] payload before the fan-out PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); byte[] payload = payloadConverter.convertToPayload(message.getPayload(), headers.getContentType()); - message = MessageBuilder.fromPayloadAndHeaders(payload, message.getHeaders()).build(); + @SuppressWarnings("unchecked") + M m = (M) MessageBuilder.fromPayloadAndHeaders(payload, message.getHeaders()).build(); - this.reactor.notify(getPublishKey(headers.getDestination()), Event.wrap(message)); + this.reactor.notify(getPublishKey(headers.getDestination()), Event.wrap(m)); } catch (Exception ex) { logger.error("Failed to publish " + message, ex); @@ -119,17 +121,11 @@ public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { } @Override - public void handleDisconnect(Message message) { + public void handleDisconnect(M message) { PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); removeSubscriptions(headers.getSessionId()); } -/* @Override - public void handleClientConnectionClosed(String sessionId) { - removeSubscriptions(sessionId); - } -*/ - private void removeSubscriptions(String sessionId) { List> registrations = this.subscriptionsBySession.remove(sessionId); if (logger.isTraceEnabled()) { @@ -158,7 +154,8 @@ public class ReactorPubSubMessageHandler extends AbstractPubSubMessageHandler { PubSubHeaders clientHeaders = PubSubHeaders.fromMessageHeaders(sentMessage.getHeaders()); clientHeaders.setSubscriptionId(this.subscriptionId); - Message clientMessage = MessageBuilder.fromPayloadAndHeaders(sentMessage.getPayload(), + @SuppressWarnings("unchecked") + M clientMessage = (M) MessageBuilder.fromPayloadAndHeaders(sentMessage.getPayload(), clientHeaders.toMessageHeaders()).build(); clientChannel.send(clientMessage); diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/AnnotationPubSubMessageHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/AnnotationPubSubMessageHandler.java index 9734e792b9..a702dc578b 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/AnnotationPubSubMessageHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/AnnotationPubSubMessageHandler.java @@ -52,10 +52,11 @@ import org.springframework.web.method.HandlerMethodSelector; * @author Rossen Stoyanchev * @since 4.0 */ -public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler +@SuppressWarnings("rawtypes") +public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler implements ApplicationContextAware, InitializingBean { - private PubSubChannelRegistry registry; + private PubSubChannelRegistry registry; private List messageConverters; @@ -67,12 +68,12 @@ public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler private Map unsubscribeMethods = new HashMap(); - private ArgumentResolverComposite argumentResolvers = new ArgumentResolverComposite(); + private ArgumentResolverComposite argumentResolvers = new ArgumentResolverComposite(); - private ReturnValueHandlerComposite returnValueHandlers = new ReturnValueHandlerComposite(); + private ReturnValueHandlerComposite returnValueHandlers = new ReturnValueHandlerComposite(); - public AnnotationPubSubMessageHandler(PubSubChannelRegistry registry) { + public AnnotationPubSubMessageHandler(PubSubChannelRegistry registry) { Assert.notNull(registry, "registry is required"); this.registry = registry; } @@ -96,10 +97,10 @@ public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler initHandlerMethods(); - this.argumentResolvers.addResolver(new MessageChannelArgumentResolver(this.registry.getMessageBrokerChannel())); - this.argumentResolvers.addResolver(new MessageBodyArgumentResolver(this.messageConverters)); + this.argumentResolvers.addResolver(new MessageChannelArgumentResolver(this.registry.getMessageBrokerChannel())); + this.argumentResolvers.addResolver(new MessageBodyArgumentResolver(this.messageConverters)); - this.returnValueHandlers.addHandler(new MessageReturnValueHandler(this.registry.getClientOutputChannel())); + this.returnValueHandlers.addHandler(new MessageReturnValueHandler(this.registry.getClientOutputChannel())); } protected void initHandlerMethods() { @@ -165,21 +166,21 @@ public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler } @Override - public void handlePublish(Message message) { + public void handlePublish(M message) { handleMessageInternal(message, this.messageMethods); } @Override - public void handleSubscribe(Message message) { + public void handleSubscribe(M message) { handleMessageInternal(message, this.subscribeMethods); } @Override - public void handleUnsubscribe(Message message) { + public void handleUnsubscribe(M message) { handleMessageInternal(message, this.unsubscribeMethods); } - private void handleMessageInternal(final Message message, Map handlerMethods) { + private void handleMessageInternal(final M message, Map handlerMethods) { PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); String destination = headers.getDestination(); @@ -192,7 +193,7 @@ public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler HandlerMethod handlerMethod = match.createWithResolvedBean(); // TODO: - InvocableMessageHandlerMethod invocableHandlerMethod = new InvocableMessageHandlerMethod(handlerMethod); + InvocableMessageHandlerMethod invocableHandlerMethod = new InvocableMessageHandlerMethod(handlerMethod); invocableHandlerMethod.setMessageMethodArgumentResolvers(this.argumentResolvers); try { diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolver.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolver.java index b54b3be830..ac8ee3b97a 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolver.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolver.java @@ -27,7 +27,8 @@ import org.springframework.messaging.Message; * @author Rossen Stoyanchev * @since 4.0 */ -public interface ArgumentResolver { +@SuppressWarnings("rawtypes") +public interface ArgumentResolver { /** * Whether the given {@linkplain MethodParameter method parameter} is @@ -53,6 +54,6 @@ public interface ArgumentResolver { * * @throws Exception in case of errors with the preparation of argument values */ - Object resolveArgument(MethodParameter parameter, Message message) throws Exception; + Object resolveArgument(MethodParameter parameter, M message) throws Exception; } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolverComposite.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolverComposite.java index c4100d4433..a3303b3e6d 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolverComposite.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ArgumentResolverComposite.java @@ -36,21 +36,21 @@ import org.springframework.util.Assert; * @author Rossen Stoyanchev * @since 4.0 */ -public class ArgumentResolverComposite implements ArgumentResolver { +@SuppressWarnings("rawtypes") +public class ArgumentResolverComposite implements ArgumentResolver { protected final Log logger = LogFactory.getLog(getClass()); - private final List argumentResolvers = - new LinkedList(); + private final List> argumentResolvers = new LinkedList>(); - private final Map argumentResolverCache = - new ConcurrentHashMap(256); + private final Map> argumentResolverCache = + new ConcurrentHashMap>(256); /** * Return a read-only list with the contained resolvers, or an empty list. */ - public List getResolvers() { + public List> getResolvers() { return Collections.unmodifiableList(this.argumentResolvers); } @@ -68,9 +68,9 @@ public class ArgumentResolverComposite implements ArgumentResolver { * @exception IllegalStateException if no suitable {@link ArgumentResolver} is found. */ @Override - public Object resolveArgument(MethodParameter parameter, Message message) throws Exception { + public Object resolveArgument(MethodParameter parameter, M message) throws Exception { - ArgumentResolver resolver = getArgumentResolver(parameter); + ArgumentResolver resolver = getArgumentResolver(parameter); Assert.notNull(resolver, "Unknown parameter type [" + parameter.getParameterType().getName() + "]"); return resolver.resolveArgument(parameter, message); } @@ -78,10 +78,10 @@ public class ArgumentResolverComposite implements ArgumentResolver { /** * Find a registered {@link ArgumentResolver} that supports the given method parameter. */ - private ArgumentResolver getArgumentResolver(MethodParameter parameter) { - ArgumentResolver result = this.argumentResolverCache.get(parameter); + private ArgumentResolver getArgumentResolver(MethodParameter parameter) { + ArgumentResolver result = this.argumentResolverCache.get(parameter); if (result == null) { - for (ArgumentResolver resolver : this.argumentResolvers) { + for (ArgumentResolver resolver : this.argumentResolvers) { if (resolver.supportsParameter(parameter)) { result = resolver; this.argumentResolverCache.put(parameter, result); @@ -95,7 +95,7 @@ public class ArgumentResolverComposite implements ArgumentResolver { /** * Add the given {@link ArgumentResolver}. */ - public ArgumentResolverComposite addResolver(ArgumentResolver argumentResolver) { + public ArgumentResolverComposite addResolver(ArgumentResolver argumentResolver) { this.argumentResolvers.add(argumentResolver); return this; } @@ -103,9 +103,9 @@ public class ArgumentResolverComposite implements ArgumentResolver { /** * Add the given {@link ArgumentResolver}s. */ - public ArgumentResolverComposite addResolvers(List argumentResolvers) { + public ArgumentResolverComposite addResolvers(List> argumentResolvers) { if (argumentResolvers != null) { - for (ArgumentResolver resolver : argumentResolvers) { + for (ArgumentResolver resolver : argumentResolvers) { this.argumentResolvers.add(resolver); } } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/InvocableMessageHandlerMethod.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/InvocableMessageHandlerMethod.java index 993480d086..37932ba8e3 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/InvocableMessageHandlerMethod.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/InvocableMessageHandlerMethod.java @@ -43,9 +43,10 @@ import org.springframework.web.method.HandlerMethod; * @author Rossen Stoyanchev * @since 4.0 */ -public class InvocableMessageHandlerMethod extends HandlerMethod { +@SuppressWarnings("rawtypes") +public class InvocableMessageHandlerMethod extends HandlerMethod { - private ArgumentResolverComposite argumentResolvers = new ArgumentResolverComposite(); + private ArgumentResolverComposite argumentResolvers = new ArgumentResolverComposite(); private ParameterNameDiscoverer parameterNameDiscoverer = new LocalVariableTableParameterNameDiscoverer(); @@ -75,7 +76,7 @@ public class InvocableMessageHandlerMethod extends HandlerMethod { * Set {@link ArgumentResolver}s to use to use for resolving method * argument values. */ - public void setMessageMethodArgumentResolvers(ArgumentResolverComposite argumentResolvers) { + public void setMessageMethodArgumentResolvers(ArgumentResolverComposite argumentResolvers) { this.argumentResolvers = argumentResolvers; } @@ -97,7 +98,7 @@ public class InvocableMessageHandlerMethod extends HandlerMethod { * @exception Exception raised if no suitable argument resolver can be found, or the * method raised an exception */ - public final Object invoke(Message message) throws Exception { + public final Object invoke(M message) throws Exception { Object[] args = getMethodArgumentValues(message); @@ -120,7 +121,7 @@ public class InvocableMessageHandlerMethod extends HandlerMethod { /** * Get the method argument values for the current request. */ - private Object[] getMethodArgumentValues(Message message) throws Exception { + private Object[] getMethodArgumentValues(M message) throws Exception { MethodParameter[] parameters = getMethodParameters(); Object[] args = new Object[parameters.length]; diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageBodyArgumentResolver.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageBodyArgumentResolver.java index dba7c49f95..4b280e238b 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageBodyArgumentResolver.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageBodyArgumentResolver.java @@ -32,7 +32,8 @@ import org.springframework.web.messaging.converter.MessageConverter; * @author Rossen Stoyanchev * @since 4.0 */ -public class MessageBodyArgumentResolver implements ArgumentResolver { +@SuppressWarnings("rawtypes") +public class MessageBodyArgumentResolver implements ArgumentResolver { private final MessageConverter converter; @@ -47,7 +48,7 @@ public class MessageBodyArgumentResolver implements ArgumentResolver { } @Override - public Object resolveArgument(MethodParameter parameter, Message message) throws Exception { + public Object resolveArgument(MethodParameter parameter, M message) throws Exception { Object arg = null; diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageChannelArgumentResolver.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageChannelArgumentResolver.java index 13d7927495..726236694c 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageChannelArgumentResolver.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageChannelArgumentResolver.java @@ -28,12 +28,13 @@ import org.springframework.web.messaging.support.SessionMessageChannel; * @author Rossen Stoyanchev * @since 4.0 */ -public class MessageChannelArgumentResolver implements ArgumentResolver { +@SuppressWarnings("rawtypes") +public class MessageChannelArgumentResolver implements ArgumentResolver { - private MessageChannel> messageBrokerChannel; + private MessageChannel messageBrokerChannel; - public MessageChannelArgumentResolver(MessageChannel> messageBrokerChannel) { + public MessageChannelArgumentResolver(MessageChannel messageBrokerChannel) { Assert.notNull(messageBrokerChannel, "messageBrokerChannel is required"); this.messageBrokerChannel = messageBrokerChannel; } @@ -44,10 +45,10 @@ public class MessageChannelArgumentResolver implements ArgumentResolver { } @Override - public Object resolveArgument(MethodParameter parameter, Message message) throws Exception { + public Object resolveArgument(MethodParameter parameter, M message) throws Exception { Assert.notNull(this.messageBrokerChannel, "messageBrokerChannel is required"); final String sessionId = PubSubHeaders.fromMessageHeaders(message.getHeaders()).getSessionId(); - return new SessionMessageChannel(this.messageBrokerChannel, sessionId); + return new SessionMessageChannel(this.messageBrokerChannel, sessionId); } } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageReturnValueHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageReturnValueHandler.java index 2541eb33b1..f9aab71517 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageReturnValueHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/MessageReturnValueHandler.java @@ -28,12 +28,13 @@ import org.springframework.web.messaging.PubSubHeaders; * @author Rossen Stoyanchev * @since 4.0 */ -public class MessageReturnValueHandler implements ReturnValueHandler { +@SuppressWarnings("rawtypes") +public class MessageReturnValueHandler implements ReturnValueHandler { - private MessageChannel> clientChannel; + private MessageChannel clientChannel; - public MessageReturnValueHandler(MessageChannel> clientChannel) { + public MessageReturnValueHandler(MessageChannel clientChannel) { Assert.notNull(clientChannel, "clientChannel is required"); this.clientChannel = clientChannel; } @@ -55,14 +56,14 @@ public class MessageReturnValueHandler implements ReturnValueHandler { // return Message.class.isAssignableFrom(paramType); } - @SuppressWarnings("unchecked") @Override - public void handleReturnValue(Object returnValue, MethodParameter returnType, Message message) + public void handleReturnValue(Object returnValue, MethodParameter returnType, M message) throws Exception { Assert.notNull(this.clientChannel, "No clientChannel to send messages to"); - Message returnMessage = (Message) returnValue; + @SuppressWarnings("unchecked") + M returnMessage = (M) returnValue; if (returnMessage == null) { return; } @@ -72,7 +73,7 @@ public class MessageReturnValueHandler implements ReturnValueHandler { this.clientChannel.send(returnMessage); } - protected Message updateReturnMessage(Message returnMessage, Message message) { + protected M updateReturnMessage(M returnMessage, M message) { PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); String sessionId = headers.getSessionId(); @@ -89,7 +90,12 @@ public class MessageReturnValueHandler implements ReturnValueHandler { } Object payload = returnMessage.getPayload(); - return MessageBuilder.fromPayloadAndHeaders(payload, returnHeaders.toMessageHeaders()).build(); + return createMessage(returnHeaders, payload); + } + + @SuppressWarnings("unchecked") + private M createMessage(PubSubHeaders returnHeaders, Object payload) { + return (M) MessageBuilder.fromPayloadAndHeaders(payload, returnHeaders.toMessageHeaders()).build(); } } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandler.java index 72f11b611d..e29761ac00 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandler.java @@ -27,7 +27,8 @@ import org.springframework.messaging.Message; * @author Rossen Stoyanchev * @since 4.0 */ -public interface ReturnValueHandler { +@SuppressWarnings("rawtypes") +public interface ReturnValueHandler { /** * Whether the given {@linkplain MethodParameter method return type} is @@ -50,6 +51,6 @@ public interface ReturnValueHandler { * @param message the message that caused this method to be called * @throws Exception if the return value handling results in an error */ - void handleReturnValue(Object returnValue, MethodParameter returnType, Message message) throws Exception; + void handleReturnValue(Object returnValue, MethodParameter returnType, M message) throws Exception; } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandlerComposite.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandlerComposite.java index 195e755294..109bce1299 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandlerComposite.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/ReturnValueHandlerComposite.java @@ -28,16 +28,16 @@ import org.springframework.util.Assert; * @author Rossen Stoyanchev * @since 4.0 */ -public class ReturnValueHandlerComposite implements ReturnValueHandler { +@SuppressWarnings("rawtypes") +public class ReturnValueHandlerComposite implements ReturnValueHandler { - private final List returnValueHandlers = - new ArrayList(); + private final List> returnValueHandlers = new ArrayList>(); /** * Add the given {@link ReturnValueHandler}. */ - public ReturnValueHandlerComposite addHandler(ReturnValueHandler returnValuehandler) { + public ReturnValueHandlerComposite addHandler(ReturnValueHandler returnValuehandler) { this.returnValueHandlers.add(returnValuehandler); return this; } @@ -45,9 +45,9 @@ public class ReturnValueHandlerComposite implements ReturnValueHandler { /** * Add the given {@link ReturnValueHandler}s. */ - public ReturnValueHandlerComposite addHandlers(List handlers) { + public ReturnValueHandlerComposite addHandlers(List> handlers) { if (handlers != null) { - for (ReturnValueHandler handler : handlers) { + for (ReturnValueHandler handler : handlers) { this.returnValueHandlers.add(handler); } } @@ -59,8 +59,8 @@ public class ReturnValueHandlerComposite implements ReturnValueHandler { return getReturnValueHandler(returnType) != null; } - private ReturnValueHandler getReturnValueHandler(MethodParameter returnType) { - for (ReturnValueHandler handler : this.returnValueHandlers) { + private ReturnValueHandler getReturnValueHandler(MethodParameter returnType) { + for (ReturnValueHandler handler : this.returnValueHandlers) { if (handler.supportsReturnType(returnType)) { return handler; } @@ -69,10 +69,10 @@ public class ReturnValueHandlerComposite implements ReturnValueHandler { } @Override - public void handleReturnValue(Object returnValue, MethodParameter returnType, Message message) + public void handleReturnValue(Object returnValue, MethodParameter returnType, M message) throws Exception { - ReturnValueHandler handler = getReturnValueHandler(returnType); + ReturnValueHandler handler = getReturnValueHandler(returnType); Assert.notNull(handler, "Unknown return value type [" + returnType.getParameterType().getName() + "]"); handler.handleReturnValue(returnValue, returnType, message); } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompMessageConverter.java b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompMessageConverter.java index 124a38eaaf..b3cbbaf984 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompMessageConverter.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompMessageConverter.java @@ -37,7 +37,8 @@ import org.springframework.web.messaging.stomp.StompHeaders; * @author Rossen Stoyanchev * @since 4.0 */ -public class StompMessageConverter { +@SuppressWarnings("rawtypes") +public class StompMessageConverter { private static final Charset STOMP_CHARSET = Charset.forName("UTF-8"); @@ -50,7 +51,7 @@ public class StompMessageConverter { /** * @param stompContent a complete STOMP message (without the trailing 0x00) as byte[] or String. */ - public Message toMessage(Object stompContent, String sessionId) { + public M toMessage(Object stompContent, String sessionId) { byte[] byteContent = null; if (stompContent instanceof String) { @@ -101,7 +102,12 @@ public class StompMessageConverter { byte[] payload = new byte[totalLength - payloadIndex]; System.arraycopy(byteContent, payloadIndex, payload, 0, totalLength - payloadIndex); - return MessageBuilder.fromPayloadAndHeaders(payload, stompHeaders.toMessageHeaders()).build(); + return createMessage(stompHeaders, payload); + } + + @SuppressWarnings("unchecked") + private M createMessage(StompHeaders stompHeaders, byte[] payload) { + return (M) MessageBuilder.fromPayloadAndHeaders(payload, stompHeaders.toMessageHeaders()).build(); } private int findIndexOfPayload(byte[] bytes) { @@ -131,10 +137,21 @@ public class StompMessageConverter { return index; } - public byte[] fromMessage(Message message) { + public byte[] fromMessage(M message) { + + byte[] payload; + if (message.getPayload() instanceof byte[]) { + payload = (byte[]) message.getPayload(); + } + else { + throw new IllegalArgumentException( + "stompContent is not byte[]: " + message.getPayload().getClass()); + } + ByteArrayOutputStream out = new ByteArrayOutputStream(); MessageHeaders messageHeaders = message.getHeaders(); StompHeaders stompHeaders = StompHeaders.fromMessageHeaders(messageHeaders); + try { out.write(stompHeaders.getStompCommand().toString().getBytes("UTF-8")); out.write(LF); @@ -150,7 +167,7 @@ public class StompMessageConverter { } } out.write(LF); - out.write(message.getPayload()); + out.write(payload); out.write(0); return out.toByteArray(); } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompRelayPubSubMessageHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompRelayPubSubMessageHandler.java index 7b64fb5e89..2337814fca 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompRelayPubSubMessageHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompRelayPubSubMessageHandler.java @@ -55,11 +55,12 @@ import reactor.tcp.netty.NettyTcpClient; * @author Rossen Stoyanchev * @since 4.0 */ -public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler { +@SuppressWarnings("rawtypes") +public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler { - private MessageChannel> clientChannel; + private MessageChannel clientChannel; - private final StompMessageConverter stompMessageConverter = new StompMessageConverter(); + private final StompMessageConverter stompMessageConverter = new StompMessageConverter(); private MessageConverter payloadConverter; @@ -72,7 +73,7 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler * @param clientChannel a channel for sending messages from the remote message broker * back to clients */ - public StompRelayPubSubMessageHandler(PubSubChannelRegistry registry) { + public StompRelayPubSubMessageHandler(PubSubChannelRegistry registry) { Assert.notNull(registry, "registry is required"); this.clientChannel = registry.getClientOutputChannel(); @@ -96,7 +97,7 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler } @Override - public void handleConnect(Message message) { + public void handleConnect(M message) { StompHeaders stompHeaders = StompHeaders.fromMessageHeaders(message.getHeaders()); String sessionId = stompHeaders.getSessionId(); if (sessionId == null) { @@ -108,22 +109,22 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler } @Override - public void handlePublish(Message message) { + public void handlePublish(M message) { forwardMessage(message, StompCommand.SEND); } @Override - public void handleSubscribe(Message message) { + public void handleSubscribe(M message) { forwardMessage(message, StompCommand.SUBSCRIBE); } @Override - public void handleUnsubscribe(Message message) { + public void handleUnsubscribe(M message) { forwardMessage(message, StompCommand.UNSUBSCRIBE); } @Override - public void handleDisconnect(Message message) { + public void handleDisconnect(M message) { StompHeaders stompHeaders = StompHeaders.fromMessageHeaders(message.getHeaders()); if (stompHeaders.getStompCommand() != null) { forwardMessage(message, StompCommand.DISCONNECT); @@ -136,13 +137,13 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler } @Override - public void handleOther(Message message) { + public void handleOther(M message) { StompCommand command = (StompCommand) message.getHeaders().get(PubSubHeaders.PROTOCOL_MESSAGE_TYPE); Assert.notNull(command, "Expected STOMP command: " + message.getHeaders()); forwardMessage(message, command); } - private void forwardMessage(Message message, StompCommand command) { + private void forwardMessage(M message, StompCommand command) { StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders()); headers.setStompCommandIfNotSet(command); @@ -172,10 +173,10 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler private final AtomicBoolean isConnected = new AtomicBoolean(false); - private final BlockingQueue> messageQueue = new LinkedBlockingQueue>(50); + private final BlockingQueue messageQueue = new LinkedBlockingQueue(50); - public RelaySession(final Message message, final StompHeaders stompHeaders) { + public RelaySession(final M message, final StompHeaders stompHeaders) { Assert.notNull(message, "message is required"); Assert.notNull(stompHeaders, "stompHeaders is required"); @@ -216,7 +217,7 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler return; } - Message message = stompMessageConverter.toMessage(stompFrame, this.sessionId); + M message = stompMessageConverter.toMessage(stompFrame, this.sessionId); if (logger.isTraceEnabled()) { logger.trace("Reading message " + message); } @@ -240,19 +241,20 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler StompHeaders stompHeaders = StompHeaders.create(StompCommand.ERROR); stompHeaders.setSessionId(sessionId); stompHeaders.setMessage(errorText); - Message errorMessage = MessageBuilder.fromPayloadAndHeaders( - new byte[0], stompHeaders.toMessageHeaders()).build(); + @SuppressWarnings("unchecked") + M errorMessage = (M) MessageBuilder.fromPayloadAndHeaders(new byte[0], stompHeaders.toMessageHeaders()).build(); clientChannel.send(errorMessage); } - public void forward(Message message, StompHeaders headers) { + public void forward(M message, StompHeaders headers) { if (!this.isConnected.get()) { - message = MessageBuilder.fromPayloadAndHeaders(message.getPayload(), headers.toMessageHeaders()).build(); + @SuppressWarnings("unchecked") + M m = (M) MessageBuilder.fromPayloadAndHeaders(message.getPayload(), headers.toMessageHeaders()).build(); if (logger.isTraceEnabled()) { - logger.trace("Adding to queue message " + message + ", queue size=" + this.messageQueue.size()); + logger.trace("Adding to queue message " + m + ", queue size=" + this.messageQueue.size()); } - this.messageQueue.add(message); + this.messageQueue.add(m); return; } @@ -268,7 +270,7 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler } private void flushMessages(TcpConnection connection) { - List> messages = new ArrayList>(); + List messages = new ArrayList(); this.messageQueue.drainTo(messages); for (Message message : messages) { StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders()); @@ -284,7 +286,8 @@ public class StompRelayPubSubMessageHandler extends AbstractPubSubMessageHandler MediaType contentType = headers.getContentType(); byte[] payload = payloadConverter.convertToPayload(message.getPayload(), contentType); - Message byteMessage = MessageBuilder.fromPayloadAndHeaders(payload, headers.toMessageHeaders()).build(); + @SuppressWarnings("unchecked") + M byteMessage = (M) MessageBuilder.fromPayloadAndHeaders(payload, headers.toMessageHeaders()).build(); if (logger.isTraceEnabled()) { logger.trace("Forwarding message " + byteMessage); diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java index 002384a74c..ecd3089676 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/stomp/support/StompWebSocketHandler.java @@ -49,22 +49,24 @@ import reactor.util.Assert; * @author Rossen Stoyanchev * @since 4.0 */ -public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implements MessageHandler> { +@SuppressWarnings("rawtypes") +public class StompWebSocketHandler extends TextWebSocketHandlerAdapter + implements MessageHandler { private static final byte[] EMPTY_PAYLOAD = new byte[0]; private static Log logger = LogFactory.getLog(StompWebSocketHandler.class); - private MessageChannel outputChannel; + private MessageChannel outputChannel; - private final StompMessageConverter stompMessageConverter = new StompMessageConverter(); + private final StompMessageConverter stompMessageConverter = new StompMessageConverter(); private final Map sessions = new ConcurrentHashMap(); private MessageConverter payloadConverter = new CompositeMessageConverter(null); - public StompWebSocketHandler(PubSubChannelRegistry registry) { + public StompWebSocketHandler(PubSubChannelRegistry registry) { Assert.notNull(registry, "registry is required"); this.outputChannel = registry.getClientInputChannel(); } @@ -73,7 +75,7 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement this.payloadConverter = new CompositeMessageConverter(converters); } - public StompMessageConverter getStompMessageConverter() { + public StompMessageConverter getStompMessageConverter() { return this.stompMessageConverter; } @@ -91,12 +93,11 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement /** * Handle incoming WebSocket messages from clients. */ - @SuppressWarnings("unchecked") @Override protected void handleTextMessage(WebSocketSession session, TextMessage textMessage) { try { String payload = textMessage.getPayload(); - Message message = this.stompMessageConverter.toMessage(payload, session.getId()); + M message = this.stompMessageConverter.toMessage(payload, session.getId()); // TODO: validate size limits // http://stomp.github.io/stomp-specification-1.2.html#Size_Limits @@ -139,7 +140,7 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement } } - protected void handleConnect(final WebSocketSession session, Message message) throws IOException { + protected void handleConnect(final WebSocketSession session, M message) throws IOException { StompHeaders connectHeaders = StompHeaders.fromMessageHeaders(message.getHeaders()); StompHeaders connectedHeaders = StompHeaders.create(StompCommand.CONNECTED); @@ -161,25 +162,26 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement // TODO: security - Message connectedMessage = MessageBuilder.fromPayloadAndHeaders(EMPTY_PAYLOAD, + @SuppressWarnings("unchecked") + M connectedMessage = (M) MessageBuilder.fromPayloadAndHeaders(EMPTY_PAYLOAD, connectedHeaders.toMessageHeaders()).build(); byte[] bytes = getStompMessageConverter().fromMessage(connectedMessage); session.sendMessage(new TextMessage(new String(bytes, Charset.forName("UTF-8")))); } - protected void handlePublish(Message stompMessage) { + protected void handlePublish(M stompMessage) { } - protected void handleSubscribe(Message message) { + protected void handleSubscribe(M message) { // TODO: need a way to communicate back if subscription was successfully created or // not in which case an ERROR should be sent back and close the connection // http://stomp.github.io/stomp-specification-1.2.html#SUBSCRIBE } - protected void handleUnsubscribe(Message message) { + protected void handleUnsubscribe(M message) { } - protected void handleDisconnect(Message stompMessage) { + protected void handleDisconnect(M stompMessage) { } protected void sendErrorMessage(WebSocketSession session, Throwable error) { @@ -187,8 +189,8 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement StompHeaders headers = StompHeaders.create(StompCommand.ERROR); headers.setMessage(error.getMessage()); - Message message = MessageBuilder.fromPayloadAndHeaders(EMPTY_PAYLOAD, - headers.toMessageHeaders()).build(); + @SuppressWarnings("unchecked") + M message = (M) MessageBuilder.fromPayloadAndHeaders(EMPTY_PAYLOAD, headers.toMessageHeaders()).build(); byte[] bytes = this.stompMessageConverter.fromMessage(message); try { @@ -199,13 +201,13 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement } } - @SuppressWarnings("unchecked") @Override public void afterConnectionClosed(WebSocketSession session, CloseStatus status) throws Exception { this.sessions.remove(session.getId()); PubSubHeaders headers = PubSubHeaders.create(MessageType.DISCONNECT); headers.setSessionId(session.getId()); - Message message = MessageBuilder.fromPayloadAndHeaders(new byte[0], headers.toMessageHeaders()).build(); + @SuppressWarnings("unchecked") + M message = (M) MessageBuilder.fromPayloadAndHeaders(new byte[0], headers.toMessageHeaders()).build(); this.outputChannel.send(message); } @@ -213,7 +215,7 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement * Handle STOMP messages going back out to WebSocket clients. */ @Override - public void handleMessage(Message message) { + public void handleMessage(M message) { StompHeaders headers = StompHeaders.fromMessageHeaders(message.getHeaders()); headers.setStompCommandIfNotSet(StompCommand.MESSAGE); @@ -243,8 +245,8 @@ public class StompWebSocketHandler extends TextWebSocketHandlerAdapter implement } try { - Message byteMessage = MessageBuilder.fromPayloadAndHeaders(payload, - headers.toMessageHeaders()).build(); + @SuppressWarnings("unchecked") + M byteMessage = (M) MessageBuilder.fromPayloadAndHeaders(payload, headers.toMessageHeaders()).build(); byte[] bytes = getStompMessageConverter().fromMessage(byteMessage); session.sendMessage(new TextMessage(new String(bytes, Charset.forName("UTF-8")))); } diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/support/AbstractPubSubChannelRegistry.java b/spring-websocket/src/main/java/org/springframework/web/messaging/support/AbstractPubSubChannelRegistry.java index 429048dc55..7f07da8b46 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/support/AbstractPubSubChannelRegistry.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/support/AbstractPubSubChannelRegistry.java @@ -28,7 +28,9 @@ import org.springframework.web.messaging.PubSubChannelRegistry; * @author Rossen Stoyanchev * @since 4.0 */ -public class AbstractPubSubChannelRegistry, H extends MessageHandler> implements PubSubChannelRegistry, InitializingBean { +@SuppressWarnings("rawtypes") +public class AbstractPubSubChannelRegistry> + implements PubSubChannelRegistry, InitializingBean { private SubscribableChannel clientInputChannel; diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorMessageChannel.java b/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorMessageChannel.java index c5a99c124c..304cc04163 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorMessageChannel.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/support/ReactorMessageChannel.java @@ -47,8 +47,8 @@ public class ReactorMessageChannel implements SubscribableChannel, Me private String name = toString(); // TODO - private final Map> registrations = - new HashMap>(); + private final Map>, Registration> registrations = + new HashMap>, Registration>(); public ReactorMessageChannel(Reactor reactor) { @@ -78,7 +78,7 @@ public class ReactorMessageChannel implements SubscribableChannel, Me } @Override - public boolean subscribe(final MessageHandler handler) { + public boolean subscribe(final MessageHandler> handler) { if (this.registrations.containsKey(handler)) { logger.warn("Channel " + getName() + ", handler already subscribed " + handler); @@ -98,7 +98,7 @@ public class ReactorMessageChannel implements SubscribableChannel, Me } @Override - public boolean unsubscribe(MessageHandler handler) { + public boolean unsubscribe(MessageHandler> handler) { if (logger.isTraceEnabled()) { logger.trace("Channel " + getName() + ", removing subscription for handler " + handler); @@ -119,13 +119,12 @@ public class ReactorMessageChannel implements SubscribableChannel, Me private static final class MessageHandlerConsumer implements Consumer>> { - private final MessageHandler handler; + private final MessageHandler> handler; - private MessageHandlerConsumer(MessageHandler handler) { + private MessageHandlerConsumer(MessageHandler> handler) { this.handler = handler; } - @SuppressWarnings("unchecked") @Override public void accept(Event> event) { Message message = event.getData(); diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/support/SessionMessageChannel.java b/spring-websocket/src/main/java/org/springframework/web/messaging/support/SessionMessageChannel.java index 998a42e3cc..c466c031b8 100644 --- a/spring-websocket/src/main/java/org/springframework/web/messaging/support/SessionMessageChannel.java +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/support/SessionMessageChannel.java @@ -28,14 +28,15 @@ import reactor.util.Assert; * @author Rossen Stoyanchev * @since 4.0 */ -public class SessionMessageChannel implements MessageChannel> { +@SuppressWarnings("rawtypes") +public class SessionMessageChannel implements MessageChannel { - private MessageChannel> delegate; + private MessageChannel delegate; private final String sessionId; - public SessionMessageChannel(MessageChannel> delegate, String sessionId) { + public SessionMessageChannel(MessageChannel delegate, String sessionId) { Assert.notNull(delegate, "delegate is required"); Assert.notNull(sessionId, "sessionId is required"); this.sessionId = sessionId; @@ -43,17 +44,17 @@ public class SessionMessageChannel implements MessageChannel> { } @Override - public boolean send(Message message) { + public boolean send(M message) { return send(message, -1); } @Override - public boolean send(Message message, long timeout) { + public boolean send(M message, long timeout) { PubSubHeaders headers = PubSubHeaders.fromMessageHeaders(message.getHeaders()); headers.setSessionId(this.sessionId); - MessageBuilder messageToSend = MessageBuilder.fromPayloadAndHeaders( - message.getPayload(), headers.toMessageHeaders()); - this.delegate.send(messageToSend.build()); + @SuppressWarnings("unchecked") + M messageToSend = (M) MessageBuilder.fromPayloadAndHeaders(message.getPayload(), headers.toMessageHeaders()).build(); + this.delegate.send(messageToSend); return true; } }