Polishing contribution

Closes gh-36179
This commit is contained in:
rstoyanchev
2026-02-10 15:28:54 +00:00
parent fee7e6b7e2
commit 6507367d10
5 changed files with 77 additions and 54 deletions
@@ -19,6 +19,7 @@ package org.springframework.messaging.simp.annotation.support;
import java.lang.annotation.Annotation; import java.lang.annotation.Annotation;
import java.security.Principal; import java.security.Principal;
import java.util.Collections; import java.util.Collections;
import java.util.List;
import java.util.Map; import java.util.Map;
import java.util.function.Predicate; 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.annotation.SendToUser;
import org.springframework.messaging.simp.user.DestinationUserNameProvider; import org.springframework.messaging.simp.user.DestinationUserNameProvider;
import org.springframework.messaging.support.MessageHeaderInitializer; import org.springframework.messaging.support.MessageHeaderInitializer;
import org.springframework.messaging.support.NativeMessageHeaderAccessor;
import org.springframework.util.Assert; import org.springframework.util.Assert;
import org.springframework.util.ObjectUtils; import org.springframework.util.ObjectUtils;
import org.springframework.util.PropertyPlaceholderHelper; 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 * 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)}. * {@link Predicate#or(Predicate)}.
* <p>By default, no headers are propagated if this is not set. * <p>By default, no headers are propagated if this is not set.
* @since 7.0.4 * @since 7.0.4
@@ -257,14 +260,15 @@ public class SendToMethodReturnValueHandler implements HandlerMethodReturnValueH
new String[] {defaultPrefix + destination} : new String[] {defaultPrefix + '/' + destination}); 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); SimpMessageHeaderAccessor headerAccessor = SimpMessageHeaderAccessor.create(SimpMessageType.MESSAGE);
if (getHeaderInitializer() != null) { if (getHeaderInitializer() != null) {
getHeaderInitializer().initHeaders(headerAccessor); getHeaderInitializer().initHeaders(headerAccessor);
} }
if (inputMessage != null && this.headerFilter != null) { if (inputMessage != null && this.headerFilter != null) {
SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.wrap(inputMessage); getNativeHeaders(inputMessage).forEach((name, values) -> {
inputAccessor.toNativeHeaderMap().forEach((name, values) -> {
if (this.headerFilter.test(name)) { if (this.headerFilter.test(name)) {
headerAccessor.setNativeHeaderValues(name, values); headerAccessor.setNativeHeaderValues(name, values);
} }
@@ -278,6 +282,12 @@ public class SendToMethodReturnValueHandler implements HandlerMethodReturnValueH
return headerAccessor.getMessageHeaders(); return headerAccessor.getMessageHeaders();
} }
@SuppressWarnings("unchecked")
private static Map<String, List<String>> getNativeHeaders(Message<?> message) {
Object value = message.getHeaders().get(NativeMessageHeaderAccessor.NATIVE_HEADERS);
return (value != null ? (Map<String, List<String>>) value : Collections.emptyMap());
}
@Override @Override
public String toString() { public String toString() {
@@ -275,8 +275,8 @@ public class SimpAnnotationMethodMessageHandler extends AbstractMethodMessageHan
* Add a filter to determine which headers from the input message should be * Add a filter to determine which headers from the input message should be
* propagated to the output message. Applies to return value handling of * propagated to the output message. Applies to return value handling of
* {@code @SendTo}, {@code @SendToUser}, and {@code @SubscribeMapping} * {@code @SendTo}, {@code @SendToUser}, and {@code @SubscribeMapping}
* controller methods. Multiple filters are combined with * controller methods. The filter is applied to the "native headers" submap.
* {@link Predicate#or(Predicate)}. * Multiple filters are combined with {@link Predicate#or(Predicate)}.
* <p>By default, no headers are propagated if this is not set. * <p>By default, no headers are propagated if this is not set.
* @since 7.0.4 * @since 7.0.4
* @see SendToMethodReturnValueHandler#addHeaderFilter(Predicate) * @see SendToMethodReturnValueHandler#addHeaderFilter(Predicate)
@@ -16,6 +16,9 @@
package org.springframework.messaging.simp.annotation.support; 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 java.util.function.Predicate;
import org.apache.commons.logging.Log; 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.SendToUser;
import org.springframework.messaging.simp.annotation.SubscribeMapping; import org.springframework.messaging.simp.annotation.SubscribeMapping;
import org.springframework.messaging.support.MessageHeaderInitializer; import org.springframework.messaging.support.MessageHeaderInitializer;
import org.springframework.messaging.support.NativeMessageHeaderAccessor;
import org.springframework.util.Assert; 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 * 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)}. * {@link Predicate#or(Predicate)}.
* <p>By default, no headers are propagated if this is not set. * <p>By default, no headers are propagated if this is not set.
* @since 7.0.4 * @since 7.0.4
@@ -154,14 +159,16 @@ public class SubscriptionMethodReturnValueHandler implements HandlerMethodReturn
this.messagingTemplate.convertAndSend(destination, returnValue, headersToSend); 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); SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create(SimpMessageType.MESSAGE);
if (getHeaderInitializer() != null) { if (getHeaderInitializer() != null) {
getHeaderInitializer().initHeaders(accessor); getHeaderInitializer().initHeaders(accessor);
} }
if (inputMessage != null && this.headerFilter != null) { if (inputMessage != null && this.headerFilter != null) {
SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.wrap(inputMessage); getNativeHeaders(inputMessage).forEach((name, values) -> {
inputAccessor.toNativeHeaderMap().forEach((name, values) -> {
if (this.headerFilter.test(name)) { if (this.headerFilter.test(name)) {
accessor.setNativeHeaderValues(name, values); accessor.setNativeHeaderValues(name, values);
} }
@@ -176,4 +183,10 @@ public class SubscriptionMethodReturnValueHandler implements HandlerMethodReturn
return accessor.getMessageHeaders(); return accessor.getMessageHeaders();
} }
@SuppressWarnings("unchecked")
private static Map<String, List<String>> getNativeHeaders(Message<?> message) {
Object value = message.getHeaders().get(NativeMessageHeaderAccessor.NATIVE_HEADERS);
return (value != null ? (Map<String, List<String>>) value : Collections.emptyMap());
}
} }
@@ -298,26 +298,26 @@ public class SendToMethodReturnValueHandlerTests {
given(this.messageChannel.send(any(Message.class))).willReturn(true); given(this.messageChannel.send(any(Message.class))).willReturn(true);
String sessionId = "sess1"; String sessionId = "sess1";
String nativeHeaderName = "x-custom-header"; String headerName = "x-custom-header";
String nativeHeaderValue = "custom-value"; String headerValue = "custom-value";
SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create();
inputAccessor.setSessionId(sessionId); accessor.setSessionId(sessionId);
inputAccessor.setSubscriptionId("sub1"); accessor.setSubscriptionId("sub1");
inputAccessor.setNativeHeader(nativeHeaderName, nativeHeaderValue); accessor.setNativeHeader(headerName, headerValue);
Message<?> inputMessage = MessageBuilder.createMessage(new byte[0], inputAccessor.getMessageHeaders()); Message<?> inputMessage = MessageBuilder.createMessage(new byte[0], accessor.getMessageHeaders());
SimpMessagingTemplate template = new SimpMessagingTemplate(this.messageChannel); SimpMessagingTemplate template = new SimpMessagingTemplate(this.messageChannel);
SendToMethodReturnValueHandler handler = new SendToMethodReturnValueHandler(template, true); SendToMethodReturnValueHandler handler = new SendToMethodReturnValueHandler(template, true);
handler.addHeaderFilter(name -> name.equals(nativeHeaderName)); handler.addHeaderFilter(name -> name.equals(headerName));
handler.handleReturnValue(PAYLOAD, this.sendToReturnType, inputMessage); handler.handleReturnValue(PAYLOAD, this.sendToReturnType, inputMessage);
verify(this.messageChannel, times(2)).send(this.messageCaptor.capture()); verify(this.messageChannel, times(2)).send(this.messageCaptor.capture());
for (Message<?> sent : this.messageCaptor.getAllValues()) { for (Message<?> sent : this.messageCaptor.getAllValues()) {
SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class); accessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class);
assertThat(sentAccessor).isNotNull(); assertThat(accessor).isNotNull();
assertThat(sentAccessor.getFirstNativeHeader(nativeHeaderName)).isEqualTo(nativeHeaderValue); assertThat(accessor.getFirstNativeHeader(headerName)).isEqualTo(headerValue);
} }
} }
@@ -329,12 +329,12 @@ public class SendToMethodReturnValueHandlerTests {
String headerA = "x-header-a"; String headerA = "x-header-a";
String headerB = "x-header-b"; String headerB = "x-header-b";
SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create();
inputAccessor.setSessionId(sessionId); accessor.setSessionId(sessionId);
inputAccessor.setSubscriptionId("sub1"); accessor.setSubscriptionId("sub1");
inputAccessor.setNativeHeader(headerA, "A-value"); accessor.setNativeHeader(headerA, "A-value");
inputAccessor.setNativeHeader(headerB, "B-value"); accessor.setNativeHeader(headerB, "B-value");
Message<?> inputMessage = MessageBuilder.createMessage(new byte[0], inputAccessor.getMessageHeaders()); Message<?> inputMessage = MessageBuilder.createMessage(new byte[0], accessor.getMessageHeaders());
SimpMessagingTemplate template = new SimpMessagingTemplate(this.messageChannel); SimpMessagingTemplate template = new SimpMessagingTemplate(this.messageChannel);
SendToMethodReturnValueHandler handler = new SendToMethodReturnValueHandler(template, true); SendToMethodReturnValueHandler handler = new SendToMethodReturnValueHandler(template, true);
@@ -345,10 +345,10 @@ public class SendToMethodReturnValueHandlerTests {
verify(this.messageChannel, times(2)).send(this.messageCaptor.capture()); verify(this.messageChannel, times(2)).send(this.messageCaptor.capture());
for (Message<?> sent : this.messageCaptor.getAllValues()) { for (Message<?> sent : this.messageCaptor.getAllValues()) {
SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class); accessor = MessageHeaderAccessor.getAccessor(sent, SimpMessageHeaderAccessor.class);
assertThat(sentAccessor).isNotNull(); assertThat(accessor).isNotNull();
assertThat(sentAccessor.getFirstNativeHeader(headerA)).isEqualTo("A-value"); assertThat(accessor.getFirstNativeHeader(headerA)).isEqualTo("A-value");
assertThat(sentAccessor.getFirstNativeHeader(headerB)).isEqualTo("B-value"); assertThat(accessor.getFirstNativeHeader(headerB)).isEqualTo("B-value");
} }
} }
@@ -191,28 +191,28 @@ class SubscriptionMethodReturnValueHandlerTests {
String sessionId = "sess1"; String sessionId = "sess1";
String subscriptionId = "subs1"; String subscriptionId = "subs1";
String destination = "/dest"; String destination = "/dest";
String nativeHeaderName = "x-custom-header"; String headerName = "x-custom-header";
String nativeHeaderValue = "custom-value"; String headerValue = "custom-value";
SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create();
inputAccessor.setSessionId(sessionId); accessor.setSessionId(sessionId);
inputAccessor.setSubscriptionId(subscriptionId); accessor.setSubscriptionId(subscriptionId);
inputAccessor.setDestination(destination); accessor.setDestination(destination);
inputAccessor.setNativeHeader(nativeHeaderName, nativeHeaderValue); accessor.setNativeHeader(headerName, headerValue);
Message<?> inputMessage = MessageBuilder.createMessage(PAYLOAD, inputAccessor.getMessageHeaders()); Message<?> inputMessage = MessageBuilder.createMessage(PAYLOAD, accessor.getMessageHeaders());
MessageSendingOperations template = mock(); MessageSendingOperations template = mock();
SubscriptionMethodReturnValueHandler handler = new SubscriptionMethodReturnValueHandler(template); SubscriptionMethodReturnValueHandler handler = new SubscriptionMethodReturnValueHandler(template);
handler.addHeaderFilter(name -> name.equals(nativeHeaderName)); handler.addHeaderFilter(name -> name.equals(headerName));
handler.handleReturnValue(PAYLOAD, this.subscribeEventReturnType, inputMessage); handler.handleReturnValue(PAYLOAD, this.subscribeEventReturnType, inputMessage);
ArgumentCaptor<MessageHeaders> captor = ArgumentCaptor.forClass(MessageHeaders.class); ArgumentCaptor<MessageHeaders> captor = ArgumentCaptor.forClass(MessageHeaders.class);
verify(template).convertAndSend(eq(destination), eq(PAYLOAD), captor.capture()); verify(template).convertAndSend(eq(destination), eq(PAYLOAD), captor.capture());
SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class); accessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class);
assertThat(sentAccessor).isNotNull(); assertThat(accessor).isNotNull();
assertThat(sentAccessor.getFirstNativeHeader(nativeHeaderName)).isEqualTo(nativeHeaderValue); assertThat(accessor.getFirstNativeHeader(headerName)).isEqualTo(headerValue);
} }
@Test @Test
@@ -223,13 +223,13 @@ class SubscriptionMethodReturnValueHandlerTests {
String headerA = "x-header-a"; String headerA = "x-header-a";
String headerB = "x-header-b"; String headerB = "x-header-b";
SimpMessageHeaderAccessor inputAccessor = SimpMessageHeaderAccessor.create(); SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.create();
inputAccessor.setSessionId(sessionId); accessor.setSessionId(sessionId);
inputAccessor.setSubscriptionId(subscriptionId); accessor.setSubscriptionId(subscriptionId);
inputAccessor.setDestination(destination); accessor.setDestination(destination);
inputAccessor.setNativeHeader(headerA, "A-value"); accessor.setNativeHeader(headerA, "A-value");
inputAccessor.setNativeHeader(headerB, "B-value"); accessor.setNativeHeader(headerB, "B-value");
Message<?> inputMessage = MessageBuilder.createMessage(PAYLOAD, inputAccessor.getMessageHeaders()); Message<?> inputMessage = MessageBuilder.createMessage(PAYLOAD, accessor.getMessageHeaders());
MessageSendingOperations template = mock(); MessageSendingOperations template = mock();
SubscriptionMethodReturnValueHandler handler = new SubscriptionMethodReturnValueHandler(template); SubscriptionMethodReturnValueHandler handler = new SubscriptionMethodReturnValueHandler(template);
@@ -241,10 +241,10 @@ class SubscriptionMethodReturnValueHandlerTests {
ArgumentCaptor<MessageHeaders> captor = ArgumentCaptor.forClass(MessageHeaders.class); ArgumentCaptor<MessageHeaders> captor = ArgumentCaptor.forClass(MessageHeaders.class);
verify(template).convertAndSend(eq(destination), eq(PAYLOAD), captor.capture()); verify(template).convertAndSend(eq(destination), eq(PAYLOAD), captor.capture());
SimpMessageHeaderAccessor sentAccessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class); accessor = MessageHeaderAccessor.getAccessor(captor.getValue(), SimpMessageHeaderAccessor.class);
assertThat(sentAccessor).isNotNull(); assertThat(accessor).isNotNull();
assertThat(sentAccessor.getFirstNativeHeader(headerA)).isEqualTo("A-value"); assertThat(accessor.getFirstNativeHeader(headerA)).isEqualTo("A-value");
assertThat(sentAccessor.getFirstNativeHeader(headerB)).isEqualTo("B-value"); assertThat(accessor.getFirstNativeHeader(headerB)).isEqualTo("B-value");
} }
private Message<?> createInputMessage(String sessId, String subsId, String dest, Principal principal) { private Message<?> createInputMessage(String sessId, String subsId, String dest, Principal principal) {