Add ChannelInterceptor to spring-messaging module

Issue: SPR-10866
This commit is contained in:
Rossen Stoyanchev
2013-08-28 14:41:20 -04:00
parent 467a6b9fa7
commit 4b2847d9d1
12 changed files with 844 additions and 61 deletions

View File

@@ -16,8 +16,14 @@
package org.springframework.messaging.simp.stomp;
import java.util.List;
import java.util.Map;
import org.junit.Test;
import org.springframework.http.MediaType;
import org.springframework.messaging.simp.SimpMessageType;
import org.springframework.util.LinkedMultiValueMap;
import org.springframework.util.MultiValueMap;
import static org.junit.Assert.*;
@@ -32,7 +38,8 @@ public class StompHeaderAccessorTests {
@Test
public void testStompCommandSet() {
public void createWithCommand() {
StompHeaderAccessor accessor = StompHeaderAccessor.create(StompCommand.CONNECTED);
assertEquals(StompCommand.CONNECTED, accessor.getCommand());
@@ -40,4 +47,111 @@ public class StompHeaderAccessorTests {
assertEquals(StompCommand.CONNECTED, accessor.getCommand());
}
@Test
public void createWithSubscribeNativeHeaders() {
MultiValueMap<String, String> extHeaders = new LinkedMultiValueMap<>();
extHeaders.add(StompHeaderAccessor.STOMP_ID_HEADER, "s1");
extHeaders.add(StompHeaderAccessor.STOMP_DESTINATION_HEADER, "/d");
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE, extHeaders);
assertEquals(StompCommand.SUBSCRIBE, headers.getCommand());
assertEquals(SimpMessageType.SUBSCRIBE, headers.getMessageType());
assertEquals("/d", headers.getDestination());
assertEquals("s1", headers.getSubscriptionId());
}
@Test
public void createWithUnubscribeNativeHeaders() {
MultiValueMap<String, String> extHeaders = new LinkedMultiValueMap<>();
extHeaders.add(StompHeaderAccessor.STOMP_ID_HEADER, "s1");
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.UNSUBSCRIBE, extHeaders);
assertEquals(StompCommand.UNSUBSCRIBE, headers.getCommand());
assertEquals(SimpMessageType.UNSUBSCRIBE, headers.getMessageType());
assertEquals("s1", headers.getSubscriptionId());
}
@Test
public void createWithMessageFrameNativeHeaders() {
MultiValueMap<String, String> extHeaders = new LinkedMultiValueMap<>();
extHeaders.add(StompHeaderAccessor.DESTINATION_HEADER, "/d");
extHeaders.add(StompHeaderAccessor.STOMP_SUBSCRIPTION_HEADER, "s1");
extHeaders.add(StompHeaderAccessor.STOMP_CONTENT_TYPE_HEADER, "application/json");
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.MESSAGE, extHeaders);
assertEquals(StompCommand.MESSAGE, headers.getCommand());
assertEquals(SimpMessageType.MESSAGE, headers.getMessageType());
assertEquals("s1", headers.getSubscriptionId());
}
@Test
public void toNativeHeadersSubscribe() {
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE);
headers.setSubscriptionId("s1");
headers.setDestination("/d");
Map<String, List<String>> actual = headers.toNativeHeaderMap();
assertEquals(2, actual.size());
assertEquals("s1", actual.get(StompHeaderAccessor.STOMP_ID_HEADER).get(0));
assertEquals("/d", actual.get(StompHeaderAccessor.STOMP_DESTINATION_HEADER).get(0));
}
@Test
public void toNativeHeadersUnsubscribe() {
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.UNSUBSCRIBE);
headers.setSubscriptionId("s1");
Map<String, List<String>> actual = headers.toNativeHeaderMap();
assertEquals(1, actual.size());
assertEquals("s1", actual.get(StompHeaderAccessor.STOMP_ID_HEADER).get(0));
}
@Test
public void toNativeHeadersMessageFrame() {
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.MESSAGE);
headers.setSubscriptionId("s1");
headers.setDestination("/d");
headers.setContentType(MediaType.APPLICATION_JSON);
Map<String, List<String>> actual = headers.toNativeHeaderMap();
assertEquals(4, actual.size());
assertEquals("s1", actual.get(StompHeaderAccessor.STOMP_SUBSCRIPTION_HEADER).get(0));
assertEquals("/d", actual.get(StompHeaderAccessor.STOMP_DESTINATION_HEADER).get(0));
assertEquals("application/json", actual.get(StompHeaderAccessor.STOMP_CONTENT_TYPE_HEADER).get(0));
assertNotNull("message-id was not created", actual.get(StompHeaderAccessor.STOMP_MESSAGE_ID_HEADER).get(0));
}
@Test
public void modifyCustomNativeHeader() {
MultiValueMap<String, String> extHeaders = new LinkedMultiValueMap<>();
extHeaders.add(StompHeaderAccessor.STOMP_ID_HEADER, "s1");
extHeaders.add(StompHeaderAccessor.STOMP_DESTINATION_HEADER, "/d");
extHeaders.add("accountId", "ABC123");
StompHeaderAccessor headers = StompHeaderAccessor.create(StompCommand.SUBSCRIBE, extHeaders);
String accountId = headers.getFirstNativeHeader("accountId");
headers.setNativeHeader("accountId", accountId.toLowerCase());
Map<String, List<String>> actual = headers.toNativeHeaderMap();
assertEquals(3, actual.size());
assertEquals("s1", actual.get(StompHeaderAccessor.STOMP_ID_HEADER).get(0));
assertEquals("/d", actual.get(StompHeaderAccessor.STOMP_DESTINATION_HEADER).get(0));
assertNotNull("abc123", actual.get("accountId").get(0));
}
}

View File

@@ -0,0 +1,78 @@
/*
* 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.messaging.support;
import java.util.Collections;
import java.util.HashMap;
import java.util.Map;
import org.junit.Test;
import org.springframework.messaging.MessageHeaders;
import static org.junit.Assert.*;
/**
* Test fixture for {@link MessageHeaderAccessor}.
*
* @author Rossen Stoyanchev
*/
public class MessageHeaderAccessorTests {
@Test
public void empty() {
MessageHeaderAccessor headers = new MessageHeaderAccessor();
assertEquals(Collections.emptyMap(), headers.toMap());
}
@Test
public void wrapMessage() {
Map<String, Object> original = new HashMap<>();
original.put("foo", "bar");
original.put("bar", "baz");
GenericMessage<String> message = new GenericMessage<>("p", original);
MessageHeaderAccessor headers = new MessageHeaderAccessor(message);
Map<String, Object> actual = headers.toMap();
assertEquals(4, actual.size());
assertNotNull(actual.get(MessageHeaders.ID));
assertNotNull(actual.get(MessageHeaders.TIMESTAMP));
assertEquals("bar", actual.get("foo"));
assertEquals("baz", actual.get("bar"));
}
@Test
public void wrapMessageAndModifyHeaders() {
Map<String, Object> original = new HashMap<>();
original.put("foo", "bar");
original.put("bar", "baz");
GenericMessage<String> message = new GenericMessage<>("p", original);
MessageHeaderAccessor headers = new MessageHeaderAccessor(message);
headers.setHeader("foo", "BAR");
Map<String, Object> actual = headers.toMap();
assertEquals(4, actual.size());
assertNotNull(actual.get(MessageHeaders.ID));
assertNotNull(actual.get(MessageHeaders.TIMESTAMP));
assertEquals("BAR", actual.get("foo"));
assertEquals("baz", actual.get("bar"));
}
}

View File

@@ -0,0 +1,127 @@
/*
* 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.messaging.support;
import java.util.Arrays;
import java.util.Collections;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import org.junit.Test;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageHeaders;
import org.springframework.util.LinkedMultiValueMap;
import org.springframework.util.MultiValueMap;
import static org.junit.Assert.*;
/**
* Test fixture for {@link NativeMessageHeaderAccessor}.
*
* @author Rossen Stoyanchev
*/
public class NativeMessageHeaderAccessorTests {
@Test
public void originalNativeHeaders() {
MultiValueMap<String, String> original = new LinkedMultiValueMap<>();
original.add("foo", "bar");
original.add("bar", "baz");
NativeMessageHeaderAccessor headers = new NativeMessageHeaderAccessor(original);
Map<String, Object> actual = headers.toMap();
assertEquals(1, actual.size());
assertNotNull(actual.get(NativeMessageHeaderAccessor.NATIVE_HEADERS));
assertEquals(original, actual.get(NativeMessageHeaderAccessor.NATIVE_HEADERS));
}
@Test
public void wrapMessage() {
MultiValueMap<String, String> originalNativeHeaders = new LinkedMultiValueMap<>();
originalNativeHeaders.add("foo", "bar");
originalNativeHeaders.add("bar", "baz");
Map<String, Object> original = new HashMap<String, Object>();
original.put("a", "b");
original.put(NativeMessageHeaderAccessor.NATIVE_HEADERS, originalNativeHeaders);
GenericMessage<String> message = new GenericMessage<>("p", original);
NativeMessageHeaderAccessor headers = new NativeMessageHeaderAccessor(message);
Map<String, Object> actual = headers.toMap();
assertEquals(4, actual.size());
assertNotNull(actual.get(MessageHeaders.ID));
assertNotNull(actual.get(MessageHeaders.TIMESTAMP));
assertEquals("b", actual.get("a"));
assertNotNull(actual.get(NativeMessageHeaderAccessor.NATIVE_HEADERS));
assertEquals(originalNativeHeaders, actual.get(NativeMessageHeaderAccessor.NATIVE_HEADERS));
}
@Test
public void wrapNullMessage() {
NativeMessageHeaderAccessor headers = new NativeMessageHeaderAccessor((Message<?>) null);
Map<String, Object> actual = headers.toMap();
assertEquals(1, actual.size());
@SuppressWarnings("unchecked")
Map<String, List<String>> actualNativeHeaders =
(Map<String, List<String>>) actual.get(NativeMessageHeaderAccessor.NATIVE_HEADERS);
assertEquals(Collections.emptyMap(), actualNativeHeaders);
}
@Test
public void wrapMessageAndModifyHeaders() {
MultiValueMap<String, String> originalNativeHeaders = new LinkedMultiValueMap<>();
originalNativeHeaders.add("foo", "bar");
originalNativeHeaders.add("bar", "baz");
Map<String, Object> original = new HashMap<String, Object>();
original.put("a", "b");
original.put(NativeMessageHeaderAccessor.NATIVE_HEADERS, originalNativeHeaders);
GenericMessage<String> message = new GenericMessage<>("p", original);
NativeMessageHeaderAccessor headers = new NativeMessageHeaderAccessor(message);
headers.setHeader("a", "B");
headers.setNativeHeader("foo", "BAR");
Map<String, Object> actual = headers.toMap();
assertEquals(4, actual.size());
assertNotNull(actual.get(MessageHeaders.ID));
assertNotNull(actual.get(MessageHeaders.TIMESTAMP));
assertEquals("B", actual.get("a"));
@SuppressWarnings("unchecked")
Map<String, List<String>> actualNativeHeaders =
(Map<String, List<String>>) actual.get(NativeMessageHeaderAccessor.NATIVE_HEADERS);
assertNotNull(actualNativeHeaders);
assertEquals(Arrays.asList("BAR"), actualNativeHeaders.get("foo"));
assertEquals(Arrays.asList("baz"), actualNativeHeaders.get("bar"));
}
}

View File

@@ -0,0 +1,156 @@
/*
* 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.messaging.support.channel;
import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.concurrent.atomic.AtomicInteger;
import org.junit.Before;
import org.junit.Test;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.MessageHandler;
import org.springframework.messaging.MessagingException;
import org.springframework.messaging.support.MessageBuilder;
import static org.junit.Assert.*;
/**
* Test fixture for the use of {@link ChannelInterceptor}s.
* @author Rossen Stoyanchev
*/
public class ChannelInterceptorTests {
private ExecutorSubscribableChannel channel;
private TestMessageHandler messageHandler;
@Before
public void setup() {
this.channel = new ExecutorSubscribableChannel();
this.messageHandler = new TestMessageHandler();
this.channel.subscribe(this.messageHandler);
}
@Test
public void preSendInterceptorReturningModifiedMessage() {
this.channel.addInterceptor(new PreSendReturnsMessageInterceptor());
this.channel.send(MessageBuilder.withPayload("test").build());
assertEquals(1, this.messageHandler.messages.size());
Message<?> result = this.messageHandler.messages.get(0);
assertNotNull(result);
assertEquals("test", result.getPayload());
assertEquals(1, result.getHeaders().get(PreSendReturnsMessageInterceptor.class.getSimpleName()));
}
@Test
public void preSendInterceptorReturningNull() {
PreSendReturnsNullInterceptor interceptor = new PreSendReturnsNullInterceptor();
this.channel.addInterceptor(interceptor);
Message<?> message = MessageBuilder.withPayload("test").build();
this.channel.send(message);
assertEquals(1, interceptor.counter.get());
assertEquals(0, this.messageHandler.messages.size());
}
@Test
public void postSendInterceptorMessageWasSent() {
final AtomicBoolean invoked = new AtomicBoolean(false);
this.channel.addInterceptor(new ChannelInterceptorAdapter() {
@Override
public void postSend(Message<?> message, MessageChannel channel, boolean sent) {
assertNotNull(message);
assertNotNull(channel);
assertSame(ChannelInterceptorTests.this.channel, channel);
assertTrue(sent);
invoked.set(true);
}
});
this.channel.send(MessageBuilder.withPayload("test").build());
assertTrue(invoked.get());
}
@Test
public void postSendInterceptorMessageWasNotSent() {
final AbstractMessageChannel testChannel = new AbstractMessageChannel() {
@Override
protected boolean sendInternal(Message<?> message, long timeout) {
return false;
}
};
final AtomicBoolean invoked = new AtomicBoolean(false);
testChannel.addInterceptor(new ChannelInterceptorAdapter() {
@Override
public void postSend(Message<?> message, MessageChannel channel, boolean sent) {
assertNotNull(message);
assertNotNull(channel);
assertSame(testChannel, channel);
assertFalse(sent);
invoked.set(true);
}
});
testChannel.send(MessageBuilder.withPayload("test").build());
assertTrue(invoked.get());
}
private static class TestMessageHandler implements MessageHandler {
private List<Message<?>> messages = new ArrayList<Message<?>>();
@Override
public void handleMessage(Message<?> message) throws MessagingException {
this.messages.add(message);
}
}
private static class PreSendReturnsMessageInterceptor extends ChannelInterceptorAdapter {
private AtomicInteger counter = new AtomicInteger();
private String foo;
@Override
public Message<?> preSend(Message<?> message, MessageChannel channel) {
assertNotNull(message);
return MessageBuilder.fromMessage(message).setHeader(
this.getClass().getSimpleName(), counter.incrementAndGet()).build();
}
}
private static class PreSendReturnsNullInterceptor extends ChannelInterceptorAdapter {
private AtomicInteger counter = new AtomicInteger();
@Override
public Message<?> preSend(Message<?> message, MessageChannel channel) {
assertNotNull(message);
counter.incrementAndGet();
return null;
}
}
}

View File

@@ -26,6 +26,7 @@ import org.mockito.Mock;
import org.mockito.MockitoAnnotations;
import org.springframework.core.task.TaskExecutor;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageDeliveryException;
import org.springframework.messaging.MessageHandler;
import org.springframework.messaging.MessagingException;
import org.springframework.messaging.support.MessageBuilder;
@@ -117,8 +118,8 @@ public class PublishSubscibeChannelTests {
try {
this.channel.send(message);
}
catch(RuntimeException actualException) {
assertThat(actualException, equalTo(ex));
catch(MessageDeliveryException actualException) {
assertThat((RuntimeException) actualException.getCause(), equalTo(ex));
}
verifyZeroInteractions(secondHandler);
}