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 ca50ea0b9c..209030d444 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 @@ -109,6 +109,7 @@ public class AnnotationPubSubMessageHandler extends AbstractPubSubMessageHandler this.argumentResolvers.addResolver(new MessageBodyArgumentResolver(this.messageConverters)); this.returnValueHandlers.addHandler(new MessageReturnValueHandler(this.clientChannel)); + this.returnValueHandlers.addHandler(new PayloadReturnValueHandler(this.clientChannel)); } protected void initHandlerMethods() { 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 85e01e88b9..d148feea36 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 @@ -40,19 +40,9 @@ public class MessageReturnValueHandler implements ReturnValueHandler { @Override public boolean supportsReturnType(MethodParameter returnType) { + // TODO: List return value Class paramType = returnType.getParameterType(); return Message.class.isAssignableFrom(paramType); - -// if (Message.class.isAssignableFrom(paramType)) { -// return true; -// } -// else if (List.class.isAssignableFrom(paramType)) { -// Type type = returnType.getGenericParameterType(); -// if (type instanceof ParameterizedType) { -// Type genericType = ((ParameterizedType) type).getActualTypeArguments()[0]; -// } -// } -// return Message.class.isAssignableFrom(paramType); } @Override diff --git a/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/PayloadReturnValueHandler.java b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/PayloadReturnValueHandler.java new file mode 100644 index 0000000000..1c6c8d6708 --- /dev/null +++ b/spring-websocket/src/main/java/org/springframework/web/messaging/service/method/PayloadReturnValueHandler.java @@ -0,0 +1,69 @@ +/* + * Copyright 2002-2013 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.web.messaging.service.method; + +import org.springframework.core.MethodParameter; +import org.springframework.messaging.Message; +import org.springframework.messaging.MessageChannel; +import org.springframework.messaging.support.MessageBuilder; +import org.springframework.util.Assert; +import org.springframework.web.messaging.support.WebMessageHeaderAccesssor; + + +/** + * @author Rossen Stoyanchev + * @since 4.0 + */ +public class PayloadReturnValueHandler implements ReturnValueHandler { + + private MessageChannel clientChannel; + + + public PayloadReturnValueHandler(MessageChannel clientChannel) { + Assert.notNull(clientChannel, "clientChannel is required"); + this.clientChannel = clientChannel; + } + + @Override + public boolean supportsReturnType(MethodParameter returnType) { + return true; + } + + @Override + public void handleReturnValue(Object returnValue, MethodParameter returnType, Message message) + throws Exception { + + Assert.notNull(this.clientChannel, "No clientChannel to send messages to"); + + if (returnValue == null) { + return; + } + + WebMessageHeaderAccesssor headers = WebMessageHeaderAccesssor.wrap(message); + + WebMessageHeaderAccesssor returnHeaders = WebMessageHeaderAccesssor.create(); + returnHeaders.setDestination(headers.getDestination()); + returnHeaders.setSessionId(headers.getSessionId()); + returnHeaders.setSubscriptionId(headers.getSubscriptionId()); + + Message returnMessage = MessageBuilder.withPayload( + returnValue).copyHeaders(returnHeaders.toMap()).build(); + + this.clientChannel.send(returnMessage); + } + +}