diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandler.java index db524ca50da..f799ebbe23c 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandler.java @@ -19,6 +19,7 @@ package org.springframework.messaging.simp.annotation.support; import java.lang.annotation.Annotation; import java.security.Principal; import java.util.Collections; +import java.util.List; import java.util.Map; import java.util.function.Predicate; @@ -41,6 +42,7 @@ import org.springframework.messaging.simp.SimpMessageType; import org.springframework.messaging.simp.annotation.SendToUser; import org.springframework.messaging.simp.user.DestinationUserNameProvider; import org.springframework.messaging.support.MessageHeaderInitializer; +import org.springframework.messaging.support.NativeMessageHeaderAccessor; import org.springframework.util.Assert; import org.springframework.util.ObjectUtils; import org.springframework.util.PropertyPlaceholderHelper; @@ -138,7 +140,8 @@ public class SendToMethodReturnValueHandler implements HandlerMethodReturnValueH /** * Add a filter to determine which headers from the input message should be - * propagated to the output message. Multiple filters are combined with + * propagated to the output message. The filter is applied to the "native + * headers" submap. Multiple filters are combined with * {@link Predicate#or(Predicate)}. *

By default, no headers are propagated if this is not set. * @since 7.0.4 @@ -257,14 +260,15 @@ public class SendToMethodReturnValueHandler implements HandlerMethodReturnValueH new String[] {defaultPrefix + destination} : new String[] {defaultPrefix + '/' + destination}); } - private MessageHeaders createHeaders(@Nullable String sessionId, MethodParameter returnType, @Nullable Message inputMessage) { + private MessageHeaders createHeaders( + @Nullable String sessionId, MethodParameter returnType, @Nullable Message inputMessage) { + SimpMessageHeaderAccessor headerAccessor = SimpMessageHeaderAccessor.create(SimpMessageType.MESSAGE); if (getHeaderInitializer() != null) { getHeaderInitializer().initHeaders(headerAccessor); } if (inputMessage != null && this.headerFilter != null) { - SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.wrap(inputMessage); - inputAccessor.toNativeHeaderMap().forEach((name, values) -> { + getNativeHeaders(inputMessage).forEach((name, values) -> { if (this.headerFilter.test(name)) { headerAccessor.setNativeHeaderValues(name, values); } @@ -278,6 +282,12 @@ public class SendToMethodReturnValueHandler implements HandlerMethodReturnValueH return headerAccessor.getMessageHeaders(); } + @SuppressWarnings("unchecked") + private static Map> getNativeHeaders(Message message) { + Object value = message.getHeaders().get(NativeMessageHeaderAccessor.NATIVE_HEADERS); + return (value != null ? (Map>) value : Collections.emptyMap()); + } + @Override public String toString() { diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SimpAnnotationMethodMessageHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SimpAnnotationMethodMessageHandler.java index dc234f60301..cb65109f8db 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SimpAnnotationMethodMessageHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SimpAnnotationMethodMessageHandler.java @@ -275,8 +275,8 @@ public class SimpAnnotationMethodMessageHandler extends AbstractMethodMessageHan * Add a filter to determine which headers from the input message should be * propagated to the output message. Applies to return value handling of * {@code @SendTo}, {@code @SendToUser}, and {@code @SubscribeMapping} - * controller methods. Multiple filters are combined with - * {@link Predicate#or(Predicate)}. + * controller methods. The filter is applied to the "native headers" submap. + * Multiple filters are combined with {@link Predicate#or(Predicate)}. *

By default, no headers are propagated if this is not set. * @since 7.0.4 * @see SendToMethodReturnValueHandler#addHeaderFilter(Predicate) diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandler.java index f133daaa2b1..591f27db8c3 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandler.java @@ -16,6 +16,9 @@ package org.springframework.messaging.simp.annotation.support; +import java.util.Collections; +import java.util.List; +import java.util.Map; import java.util.function.Predicate; import org.apache.commons.logging.Log; @@ -34,6 +37,7 @@ import org.springframework.messaging.simp.SimpMessageType; import org.springframework.messaging.simp.annotation.SendToUser; import org.springframework.messaging.simp.annotation.SubscribeMapping; import org.springframework.messaging.support.MessageHeaderInitializer; +import org.springframework.messaging.support.NativeMessageHeaderAccessor; import org.springframework.util.Assert; /** @@ -99,7 +103,8 @@ public class SubscriptionMethodReturnValueHandler implements HandlerMethodReturn /** * Add a filter to determine which headers from the input message should be - * propagated to the output message. Multiple filters are combined with + * propagated to the output message. The filter is applied to the "native + * headers" submap. Multiple filters are combined with * {@link Predicate#or(Predicate)}. *

By default, no headers are propagated if this is not set. * @since 7.0.4 @@ -154,14 +159,16 @@ public class SubscriptionMethodReturnValueHandler implements HandlerMethodReturn this.messagingTemplate.convertAndSend(destination, returnValue, headersToSend); } - private MessageHeaders createHeaders(@Nullable String sessionId, String subscriptionId, MethodParameter returnType, @Nullable Message inputMessage) { + private MessageHeaders createHeaders( + @Nullable String sessionId, String subscriptionId, MethodParameter returnType, + @Nullable Message inputMessage) { + SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create(SimpMessageType.MESSAGE); if (getHeaderInitializer() != null) { getHeaderInitializer().initHeaders(accessor); } if (inputMessage != null && this.headerFilter != null) { - SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.wrap(inputMessage); - inputAccessor.toNativeHeaderMap().forEach((name, values) -> { + getNativeHeaders(inputMessage).forEach((name, values) -> { if (this.headerFilter.test(name)) { accessor.setNativeHeaderValues(name, values); } @@ -176,4 +183,10 @@ public class SubscriptionMethodReturnValueHandler implements HandlerMethodReturn return accessor.getMessageHeaders(); } + @SuppressWarnings("unchecked") + private static Map> getNativeHeaders(Message message) { + Object value = message.getHeaders().get(NativeMessageHeaderAccessor.NATIVE_HEADERS); + return (value != null ? (Map>) value : Collections.emptyMap()); + } + } diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandlerTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandlerTests.java index b3d67cbc26b..9a921e3fed2 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandlerTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SendToMethodReturnValueHandlerTests.java @@ -298,26 +298,26 @@ public class SendToMethodReturnValueHandlerTests { given(this.messageChannel.send(any(Message.class))).willReturn(true); String sessionId = "sess1"; - String nativeHeaderName = "x-custom-header"; - String nativeHeaderValue = "custom-value"; + String headerName = "x-custom-header"; + String headerValue = "custom-value"; - SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); - inputAccessor.setSessionId(sessionId); - inputAccessor.setSubscriptionId("sub1"); - inputAccessor.setNativeHeader(nativeHeaderName, nativeHeaderValue); - Message inputMessage = MessageBuilder.createMessage(new byte[0], inputAccessor.getMessageHeaders()); + SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create(); + accessor.setSessionId(sessionId); + accessor.setSubscriptionId("sub1"); + accessor.setNativeHeader(headerName, headerValue); + Message inputMessage = MessageBuilder.createMessage(new byte[0], accessor.getMessageHeaders()); SimpMessagingTemplate template = new SimpMessagingTemplate(this.messageChannel); SendToMethodReturnValueHandler handler = new SendToMethodReturnValueHandler(template, true); - handler.addHeaderFilter(name -> name.equals(nativeHeaderName)); + handler.addHeaderFilter(name -> name.equals(headerName)); handler.handleReturnValue(PAYLOAD, this.sendToReturnType, inputMessage); verify(this.messageChannel, times(2)).send(this.messageCaptor.capture()); for (Message sent : this.messageCaptor.getAllValues()) { - SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class); - assertThat(sentAccessor).isNotNull(); - assertThat(sentAccessor.getFirstNativeHeader(nativeHeaderName)).isEqualTo(nativeHeaderValue); + accessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class); + assertThat(accessor).isNotNull(); + assertThat(accessor.getFirstNativeHeader(headerName)).isEqualTo(headerValue); } } @@ -329,12 +329,12 @@ public class SendToMethodReturnValueHandlerTests { String headerA = "x-header-a"; String headerB = "x-header-b"; - SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); - inputAccessor.setSessionId(sessionId); - inputAccessor.setSubscriptionId("sub1"); - inputAccessor.setNativeHeader(headerA, "A-value"); - inputAccessor.setNativeHeader(headerB, "B-value"); - Message inputMessage = MessageBuilder.createMessage(new byte[0], inputAccessor.getMessageHeaders()); + SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create(); + accessor.setSessionId(sessionId); + accessor.setSubscriptionId("sub1"); + accessor.setNativeHeader(headerA, "A-value"); + accessor.setNativeHeader(headerB, "B-value"); + Message inputMessage = MessageBuilder.createMessage(new byte[0], accessor.getMessageHeaders()); SimpMessagingTemplate template = new SimpMessagingTemplate(this.messageChannel); SendToMethodReturnValueHandler handler = new SendToMethodReturnValueHandler(template, true); @@ -345,10 +345,10 @@ public class SendToMethodReturnValueHandlerTests { verify(this.messageChannel, times(2)).send(this.messageCaptor.capture()); for (Message sent : this.messageCaptor.getAllValues()) { - SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class); - assertThat(sentAccessor).isNotNull(); - assertThat(sentAccessor.getFirstNativeHeader(headerA)).isEqualTo("A-value"); - assertThat(sentAccessor.getFirstNativeHeader(headerB)).isEqualTo("B-value"); + accessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class); + assertThat(accessor).isNotNull(); + assertThat(accessor.getFirstNativeHeader(headerA)).isEqualTo("A-value"); + assertThat(accessor.getFirstNativeHeader(headerB)).isEqualTo("B-value"); } } diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandlerTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandlerTests.java index 3ed01bbe38f..83b8ca68a11 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandlerTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/annotation/support/SubscriptionMethodReturnValueHandlerTests.java @@ -191,28 +191,28 @@ class SubscriptionMethodReturnValueHandlerTests { String sessionId = "sess1"; String subscriptionId = "subs1"; String destination = "/dest"; - String nativeHeaderName = "x-custom-header"; - String nativeHeaderValue = "custom-value"; + String headerName = "x-custom-header"; + String headerValue = "custom-value"; - SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); - inputAccessor.setSessionId(sessionId); - inputAccessor.setSubscriptionId(subscriptionId); - inputAccessor.setDestination(destination); - inputAccessor.setNativeHeader(nativeHeaderName, nativeHeaderValue); - Message inputMessage = MessageBuilder.createMessage(PAYLOAD, inputAccessor.getMessageHeaders()); + SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create(); + accessor.setSessionId(sessionId); + accessor.setSubscriptionId(subscriptionId); + accessor.setDestination(destination); + accessor.setNativeHeader(headerName, headerValue); + Message inputMessage = MessageBuilder.createMessage(PAYLOAD, accessor.getMessageHeaders()); MessageSendingOperations template = mock(); SubscriptionMethodReturnValueHandler handler = new SubscriptionMethodReturnValueHandler(template); - handler.addHeaderFilter(name -> name.equals(nativeHeaderName)); + handler.addHeaderFilter(name -> name.equals(headerName)); handler.handleReturnValue(PAYLOAD, this.subscribeEventReturnType, inputMessage); ArgumentCaptor captor = ArgumentCaptor.forClass(MessageHeaders.class); verify(template).convertAndSend(eq(destination), eq(PAYLOAD), captor.capture()); - SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class); - assertThat(sentAccessor).isNotNull(); - assertThat(sentAccessor.getFirstNativeHeader(nativeHeaderName)).isEqualTo(nativeHeaderValue); + accessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class); + assertThat(accessor).isNotNull(); + assertThat(accessor.getFirstNativeHeader(headerName)).isEqualTo(headerValue); } @Test @@ -223,13 +223,13 @@ class SubscriptionMethodReturnValueHandlerTests { String headerA = "x-header-a"; String headerB = "x-header-b"; - SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); - inputAccessor.setSessionId(sessionId); - inputAccessor.setSubscriptionId(subscriptionId); - inputAccessor.setDestination(destination); - inputAccessor.setNativeHeader(headerA, "A-value"); - inputAccessor.setNativeHeader(headerB, "B-value"); - Message inputMessage = MessageBuilder.createMessage(PAYLOAD, inputAccessor.getMessageHeaders()); + SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create(); + accessor.setSessionId(sessionId); + accessor.setSubscriptionId(subscriptionId); + accessor.setDestination(destination); + accessor.setNativeHeader(headerA, "A-value"); + accessor.setNativeHeader(headerB, "B-value"); + Message inputMessage = MessageBuilder.createMessage(PAYLOAD, accessor.getMessageHeaders()); MessageSendingOperations template = mock(); SubscriptionMethodReturnValueHandler handler = new SubscriptionMethodReturnValueHandler(template); @@ -241,10 +241,10 @@ class SubscriptionMethodReturnValueHandlerTests { ArgumentCaptor captor = ArgumentCaptor.forClass(MessageHeaders.class); verify(template).convertAndSend(eq(destination), eq(PAYLOAD), captor.capture()); - SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class); - assertThat(sentAccessor).isNotNull(); - assertThat(sentAccessor.getFirstNativeHeader(headerA)).isEqualTo("A-value"); - assertThat(sentAccessor.getFirstNativeHeader(headerB)).isEqualTo("B-value"); + accessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class); + assertThat(accessor).isNotNull(); + assertThat(accessor.getFirstNativeHeader(headerA)).isEqualTo("A-value"); + assertThat(accessor.getFirstNativeHeader(headerB)).isEqualTo("B-value"); } private Message createInputMessage(String sessId, String subsId, String dest, Principal principal) {