Complementary to the fixes made by INT-576. Since the ID of a message will be preserved by components that broadcast messages (e.g. a pub-sub channel), multiple messages in a correlation group may have the same ID. Therefore, organizing the storage support of MessageBarrier as a Map is obsolete. Switched to Collection. Improving the performance of the Aggregator.

This commit is contained in:
Marius Bogoevici
2009-02-13 19:53:08 +00:00
parent b6e697103f
commit 977ab1b34a
6 changed files with 116 additions and 146 deletions

View File

@@ -17,9 +17,7 @@
package org.springframework.integration.aggregator;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import org.springframework.integration.core.Message;
import org.springframework.integration.message.MessageBuilder;
@@ -47,7 +45,7 @@ import org.springframework.util.CollectionUtils;
* @author Marius Bogoevici
*/
public abstract class AbstractMessageAggregator extends
AbstractMessageBarrierHandler<Map<Object, Message<?>>, Object> {
AbstractMessageBarrierHandler<List<Message<?>>> {
private volatile CompletionStrategy completionStrategy = new SequenceSizeCompletionStrategy();
@@ -62,51 +60,31 @@ public abstract class AbstractMessageAggregator extends
}
@Override
protected MessageBarrier<Map<Object, Message<?>>, Object> createMessageBarrier(Object correlationKey) {
return new MessageBarrier<Map<Object, Message<?>>, Object>(new LinkedHashMap<Object, Message<?>>(), correlationKey);
protected MessageBarrier<List<Message<?>>> createMessageBarrier(Object correlationKey) {
return new MessageBarrier<List<Message<?>>>(new ArrayList<Message<?>>(), correlationKey);
}
@Override
protected void processBarrier(MessageBarrier<Map<Object, Message<?>>, Object> barrier) {
ArrayList<Message<?>> messageList = new ArrayList<Message<?>>(barrier.getMessages().values());
if (!barrier.isComplete() && !CollectionUtils.isEmpty(messageList)) {
if (this.completionStrategy.isComplete(messageList)) {
protected void processBarrier(MessageBarrier<List<Message<?>>> barrier) {
if (!barrier.isComplete() && !CollectionUtils.isEmpty(barrier.getMessages())) {
if (this.completionStrategy.isComplete(barrier.getMessages())) {
barrier.setComplete();
}
}
if (barrier.isComplete()) {
this.removeBarrier(barrier.getCorrelationKey());
Message<?> result = this.aggregateMessages(messageList);
Message<?> result = this.aggregateMessages(barrier.getMessages());
if (result != null) {
if (result.getHeaders().getCorrelationId() == null) {
result = MessageBuilder.fromMessage(result)
.setCorrelationId(barrier.getCorrelationKey())
.build();
}
this.sendReply(result, this.resolveReplyChannelFromMessage(messageList.get(0)));
this.sendReply(result, this.resolveReplyChannelFromMessage(barrier.getMessages().get(0)));
}
}
}
@Override
protected boolean canAddMessage(Message<?> message, MessageBarrier<Map<Object, Message<?>>, Object> barrier) {
if (!super.canAddMessage(message, barrier)) {
return false;
}
if (barrier.messages.containsKey(message.getHeaders().getId())) {
logger.debug("The barrier has received message: " + message
+ ", but it already contains a similar message: "
+ barrier.getMessages().get(message.getHeaders().getId()));
return false;
}
return true;
}
@Override
protected void doAddMessage(Message<?> message, MessageBarrier<Map<Object, Message<?>>, Object> barrier) {
barrier.getMessages().put(message.getHeaders().getId(), message);
}
protected abstract Message<?> aggregateMessages(List<Message<?>> messages);
protected abstract Message<?> aggregateMessages(List<Message<?>> messages);
}

View File

@@ -63,14 +63,19 @@ import org.springframework.util.Assert;
* 'discardChannel' if provided unless 'sendPartialResultsOnTimeout' is set to
* true in which case the incomplete group will be sent to the output channel.
* <p>
* Subclasses must decide what kind of a Map they want to use, what is the logic
* for adding messages to the barrier through the '<code>doAddMessage</code>'
* method.
* Subclasses must decide what kind of a Collection they want to use. The semantics
* of adding a Message to the MessageBarrier will be decided by the Collection type.
*
* Note: this class is not part of the Spring Integration API, but
* an internal class, used for implementing components that need to keep
* a list of messages until they are ready to be released or processed
* (e.g. Resequencer or Aggregator). As such it is subject to change in future
* versions.
*
* @author Mark Fisher
* @author Marius Bogoevici
*/
public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>, K>
public abstract class AbstractMessageBarrierHandler<T extends Collection<? extends Message>>
extends AbstractMessageHandler implements BeanFactoryAware, InitializingBean {
public final static long DEFAULT_SEND_TIMEOUT = 1000;
@@ -89,7 +94,7 @@ public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>
private volatile MessageChannel discardChannel;
protected final ConcurrentMap<Object, MessageBarrier<T,K>> barriers = new ConcurrentHashMap<Object, MessageBarrier<T,K>>();
protected final ConcurrentMap<Object, MessageBarrier<T>> barriers = new ConcurrentHashMap<Object, MessageBarrier<T>>();
private volatile long timeout = DEFAULT_TIMEOUT;
@@ -256,13 +261,13 @@ public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>
}
private void processMessage(Message<?> message, Object correlationKey) {
MessageBarrier<T,K> barrier = barriers.putIfAbsent(correlationKey, createMessageBarrier(correlationKey));
MessageBarrier<T> barrier = barriers.putIfAbsent(correlationKey, createMessageBarrier(correlationKey));
if (barrier == null) {
barrier = barriers.get(correlationKey);
}
synchronized (barrier) {
if (canAddMessage(message, barrier)) {
doAddMessage(message, barrier);
((MessageBarrier)barrier).getMessages().add(message);
}
processBarrier(barrier);
}
@@ -327,7 +332,7 @@ public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>
* Verifies that a message can be added to the barrier. To be overridden by subclasses, which may add
* their own verifications. Subclasses overriding this method must call the method from the superclass.
*/
protected boolean canAddMessage(Message<?> message, MessageBarrier<T, K> barrier) {
protected boolean canAddMessage(Message<?> message, MessageBarrier<T> barrier) {
if (barrier.isComplete()) {
if (logger.isDebugEnabled()) {
logger.debug("Message received after aggregation has already completed: " + message);
@@ -341,7 +346,7 @@ public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>
/**
* Factory method for creating a MessageBarrier implementation.
*/
protected abstract MessageBarrier<T, K> createMessageBarrier(Object correlationKey);
protected abstract MessageBarrier<T> createMessageBarrier(Object correlationKey);
/**
* A method for processing the information in the message barrier after a message has been added or on pruning.
@@ -350,24 +355,18 @@ public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>
* flag to true before invoking the method.
* @param barrier the {@link MessageBarrier} to be processed
*/
protected abstract void processBarrier(MessageBarrier<T, K> barrier);
/**
* A method implemented by subclasses to add the incoming message to the message barrier. This is deferred to subclasses,
* as they should have full control over how the messages are indexed in the MessageBarrier.
*/
protected abstract void doAddMessage(Message<?> message, MessageBarrier<T, K> barrier);
protected abstract void processBarrier(MessageBarrier<T> barrier);
/**
/**
* A task that runs periodically, pruning the timed-out message barriers.
*/
private class PrunerTask implements Runnable {
public void run() {
long currentTime = System.currentTimeMillis();
for (Map.Entry<Object, MessageBarrier<T,K>> entry : barriers.entrySet()) {
for (Map.Entry<Object, MessageBarrier<T>> entry : barriers.entrySet()) {
if (currentTime - entry.getValue().getTimestamp() >= timeout) {
MessageBarrier<T,K> barrier = entry.getValue();
MessageBarrier<T> barrier = entry.getValue();
synchronized (barrier) {
removeBarrier(entry.getKey());
if (sendPartialResultOnTimeout) {
@@ -375,12 +374,12 @@ public abstract class AbstractMessageBarrierHandler<T extends Map<K, Message<?>>
processBarrier(barrier);
}
else {
for (Object message : barrier.getMessages().values()) {
for (Message message : barrier.getMessages()) {
if (logger.isDebugEnabled()) {
logger.debug("Handling of Message group with correlationId '" + entry.getKey()
+ "' has timed out.");
}
discardMessage((Message<?>) message);
discardMessage(message);
}
}
}

View File

@@ -17,23 +17,24 @@
package org.springframework.integration.aggregator;
import java.util.Map;
import java.util.Collection;
import org.springframework.integration.core.Message;
/**
* Utility class for AbstractMessageBarrierHandler and its subclasses for
* storing objects while in transit. It is a wrapper around a {@link Map},
* storing objects while in transit. It is a wrapper around a {@link java.util.Collection},
* providing special properties for recording the complete status, the creation
* time (for determining if a group of messages has timed out), and the
* correlation id for a group of messages (available after the first message has
* been added to it). This is a parameterized type, allowing different different
* client classes to use different types of Maps and their respective features.
* client classes to use different types of Collections and their respective features.
*
* This class is not thread-safe and will be synchronized by the calling code.
*
* @author Marius Bogoevici
*/
public class MessageBarrier<T extends Map<K, Message<?>>, K> {
public class MessageBarrier<T extends Collection<? extends Message>> {
protected final T messages;

View File

@@ -17,10 +17,11 @@
package org.springframework.integration.aggregator;
import java.util.ArrayList;
import java.util.Comparator;
import java.util.Iterator;
import java.util.List;
import java.util.SortedMap;
import java.util.TreeMap;
import java.util.SortedSet;
import java.util.TreeSet;
import org.springframework.integration.core.Message;
import org.springframework.integration.message.MessageBuilder;
@@ -38,9 +39,14 @@ import org.springframework.util.CollectionUtils;
* '<code>correlationId</code>' from {@link AbstractMessageBarrierHandler}
* apply here as well.
*
* Note: messages with the same sequence number will be treated as equivalent
* by this class (i.e. after a message with a given sequence number is received,
* further messages from withing the same group, that have the same sequence number,
* will be ignored.
*
* @author Marius Bogoevici
*/
public class Resequencer extends AbstractMessageBarrierHandler<SortedMap<Integer, Message<?>>, Integer> {
public class Resequencer extends AbstractMessageBarrierHandler<SortedSet<Message<?>>> {
private volatile boolean releasePartialSequences = true;
@@ -50,15 +56,19 @@ public class Resequencer extends AbstractMessageBarrierHandler<SortedMap<Integer
}
@Override
protected MessageBarrier<SortedMap<Integer, Message<?>>, Integer> createMessageBarrier(Object correlationKey) {
MessageBarrier<SortedMap<Integer, Message<?>>, Integer> messageBarrier
= new MessageBarrier<SortedMap<Integer, Message<?>>, Integer>(new TreeMap<Integer, Message<?>>(), correlationKey);
messageBarrier.getMessages().put(0, createFlagMessage(0));
protected MessageBarrier<SortedSet<Message<?>>> createMessageBarrier(Object correlationKey) {
MessageBarrier<SortedSet<Message<?>>> messageBarrier
= new MessageBarrier<SortedSet<Message<?>>>(new TreeSet<Message<?>>( new Comparator<Message<?>>() {
public int compare(Message<?> message, Message<?> message1) {
return message.getHeaders().getSequenceNumber().compareTo(message1.getHeaders().getSequenceNumber());
}
}), correlationKey);
messageBarrier.getMessages().add(createFlagMessage(0));
return messageBarrier;
}
@Override
protected void processBarrier(MessageBarrier<SortedMap<Integer, Message<?>>, Integer> barrier) {
protected void processBarrier(MessageBarrier<SortedSet<Message<?>>> barrier) {
if (hasReceivedAllMessages(barrier.getMessages())) {
barrier.setComplete();
}
@@ -72,17 +82,17 @@ public class Resequencer extends AbstractMessageBarrierHandler<SortedMap<Integer
}
}
private boolean hasReceivedAllMessages(SortedMap <Integer, Message<?>> messages) {
Message<?> firstMessage = messages.get(messages.firstKey());
Message<?> lastMessage = messages.get(messages.lastKey());
private boolean hasReceivedAllMessages(SortedSet<Message<?>> messages) {
Message<?> firstMessage = messages.first();
Message<?> lastMessage = messages.last();
return (lastMessage.getHeaders().getSequenceNumber().equals(lastMessage.getHeaders().getSequenceSize())
&& (lastMessage.getHeaders().getSequenceNumber() - firstMessage.getHeaders().getSequenceNumber() == messages.size() - 1));
}
private List<Message<?>> releaseAvailableMessages(MessageBarrier<SortedMap<Integer, Message<?>>, Integer> barrier) {
private List<Message<?>> releaseAvailableMessages(MessageBarrier<SortedSet<Message<?>>> barrier) {
if (this.releasePartialSequences || barrier.isComplete()) {
ArrayList<Message<?>> releasedMessages = new ArrayList<Message<?>>();
Iterator<Message<?>> it = barrier.getMessages().values().iterator();
Iterator<Message<?>> it = barrier.getMessages().iterator();
//remove the initial flag from the list
Message<?> flag = it.next();
it.remove();
@@ -99,7 +109,7 @@ public class Resequencer extends AbstractMessageBarrierHandler<SortedMap<Integer
}
}
//re-insert the flag so that we know where to start releasing next
barrier.getMessages().put(lastReleasedSequenceNumber, createFlagMessage(lastReleasedSequenceNumber));
barrier.getMessages().add(createFlagMessage(lastReleasedSequenceNumber));
return releasedMessages;
}
else {
@@ -109,17 +119,17 @@ public class Resequencer extends AbstractMessageBarrierHandler<SortedMap<Integer
@Override
protected boolean canAddMessage(Message<?> message,
MessageBarrier<SortedMap<Integer, Message<?>>, Integer> barrier) {
MessageBarrier<SortedSet<Message<?>>> barrier) {
if (!super.canAddMessage(message, barrier)) {
return false;
}
Message<?> flagMessage = barrier.getMessages().get(barrier.getMessages().firstKey());
if (barrier.messages.containsKey(message.getHeaders().getSequenceNumber())
Message<?> flagMessage = barrier.getMessages().first();
if (barrier.messages.contains(message)
|| flagMessage.getHeaders().getSequenceNumber() >= message.getHeaders().getSequenceNumber()) {
logger.debug("A message with the same sequence number has been already received: " + message);
return false;
}
Message<?> lastMessage = barrier.getMessages().get(barrier.getMessages().lastKey());
Message<?> lastMessage = barrier.getMessages().last();
if (lastMessage != flagMessage
&& lastMessage.getHeaders().getSequenceSize() < message.getHeaders().getSequenceNumber()) {
logger.debug("The message has a sequence number which is larger than the sequence size: "+ message);
@@ -127,12 +137,6 @@ public class Resequencer extends AbstractMessageBarrierHandler<SortedMap<Integer
}
return true;
}
@Override
protected void doAddMessage(Message<?> message, MessageBarrier<SortedMap<Integer, Message<?>>, Integer> barrier) {
//add the message to the barrier, indexing it by its sequence number
barrier.getMessages().put(message.getHeaders().getSequenceNumber(), message);
}
private static Message<Integer> createFlagMessage(int sequenceNumber) {
return MessageBuilder.withPayload(sequenceNumber).setSequenceNumber(sequenceNumber).build();

View File

@@ -34,6 +34,7 @@ import org.springframework.core.task.TaskExecutor;
import org.springframework.integration.channel.QueueChannel;
import org.springframework.integration.core.Message;
import org.springframework.integration.core.MessageChannel;
import org.springframework.integration.core.MessageHeaders;
import org.springframework.integration.message.MessageBuilder;
import org.springframework.integration.message.MessageHandlingException;
import org.springframework.integration.message.StringMessage;
@@ -66,9 +67,9 @@ public class AggregatorEndpointTests {
@Test
public void testCompleteGroupWithinTimeout() throws InterruptedException {
QueueChannel replyChannel = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel);
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel, null);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel, null);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel, null);
CountDownLatch latch = new CountDownLatch(3);
this.taskExecutor.execute(new AggregatorTestTask(this.aggregator, message1, latch));
this.taskExecutor.execute(new AggregatorTestTask(this.aggregator, message2, latch));
@@ -80,34 +81,15 @@ public class AggregatorEndpointTests {
}
@Test
public void testCompleteGroupWithinTimeoutWithDuplicates() throws InterruptedException {
public void testCompleteGroupWithinTimeoutWithSameId() throws InterruptedException {
QueueChannel replyChannel = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel);
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel, "ID#1");
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel, "ID#1");
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel, "ID#1");
CountDownLatch latch = new CountDownLatch(3);
//for testing the duplication scenario, the messages must be processed synchronously
new AggregatorTestTask(this.aggregator, message1, latch).run();
new AggregatorTestTask(this.aggregator, message2, latch).run();
new AggregatorTestTask(this.aggregator, message2, latch).run();
new AggregatorTestTask(this.aggregator, message3, latch).run();
Message<?> reply = replyChannel.receive(500);
assertNotNull(reply);
assertEquals("123456789", reply.getPayload());
}
@Test
public void testCompleteGroupWithinTimeoutWithInconsistentStructure() throws InterruptedException {
QueueChannel replyChannel = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
Message<?> message4 = createMessage("xyz", "ABC", 4, 3, replyChannel);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel);
CountDownLatch latch = new CountDownLatch(3);
//for testing the duplication scenario, the messages must be processed synchronously
new AggregatorTestTask(this.aggregator, message1, latch).run();
new AggregatorTestTask(this.aggregator, message2, latch).run();
new AggregatorTestTask(this.aggregator, message2, latch).run();
new AggregatorTestTask(this.aggregator, message3, latch).run();
Message<?> reply = replyChannel.receive(500);
assertNotNull(reply);
@@ -121,7 +103,7 @@ public class AggregatorEndpointTests {
this.aggregator.setReaperInterval(10);
this.aggregator.setDiscardChannel(discardChannel);
QueueChannel replyChannel = new QueueChannel();
Message<?> message = createMessage("123", "ABC", 2, 1, replyChannel);
Message<?> message = createMessage("123", "ABC", 2, 1, replyChannel, null);
CountDownLatch latch = new CountDownLatch(1);
AggregatorTestTask task = new AggregatorTestTask(this.aggregator, message, latch);
this.taskExecutor.execute(task);
@@ -140,8 +122,8 @@ public class AggregatorEndpointTests {
this.aggregator.setReaperInterval(10);
this.aggregator.setSendPartialResultOnTimeout(true);
QueueChannel replyChannel = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel, null);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel, null);
CountDownLatch latch = new CountDownLatch(2);
AggregatorTestTask task1 = new AggregatorTestTask(this.aggregator, message1, latch);
AggregatorTestTask task2 = new AggregatorTestTask(this.aggregator, message2, latch);
@@ -160,12 +142,12 @@ public class AggregatorEndpointTests {
public void testMultipleGroupsSimultaneously() throws InterruptedException {
QueueChannel replyChannel1 = new QueueChannel();
QueueChannel replyChannel2 = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel1);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel1);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel1);
Message<?> message4 = createMessage("abc", "XYZ", 3, 1, replyChannel2);
Message<?> message5 = createMessage("def", "XYZ", 3, 2, replyChannel2);
Message<?> message6 = createMessage("ghi", "XYZ", 3, 3, replyChannel2);
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel1, null);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel1, null);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel1, null);
Message<?> message4 = createMessage("abc", "XYZ", 3, 1, replyChannel2, null);
Message<?> message5 = createMessage("def", "XYZ", 3, 2, replyChannel2, null);
Message<?> message6 = createMessage("ghi", "XYZ", 3, 3, replyChannel2, null);
CountDownLatch latch = new CountDownLatch(6);
this.taskExecutor.execute(new AggregatorTestTask(this.aggregator, message1, latch));
this.taskExecutor.execute(new AggregatorTestTask(this.aggregator, message6, latch));
@@ -187,9 +169,9 @@ public class AggregatorEndpointTests {
QueueChannel replyChannel = new QueueChannel();
QueueChannel discardChannel = new QueueChannel();
this.aggregator.setDiscardChannel(discardChannel);
this.aggregator.handleMessage(createMessage("test-1a", 1, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-1a", 1, 1, 1, replyChannel, null));
assertEquals("test-1a", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-1b", 1, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-1b", 1, 1, 1, replyChannel, null));
assertEquals("test-1b", discardChannel.receive(100).getPayload());
}
@@ -199,13 +181,13 @@ public class AggregatorEndpointTests {
QueueChannel discardChannel = new QueueChannel();
this.aggregator.setTrackedCorrelationIdCapacity(3);
this.aggregator.setDiscardChannel(discardChannel);
this.aggregator.handleMessage(createMessage("test-1a", 1, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-1a", 1, 1, 1, replyChannel, null));
assertEquals("test-1a", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-2", 2, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-2", 2, 1, 1, replyChannel, null));
assertEquals("test-2", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-3", 3, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-3", 3, 1, 1, replyChannel, null));
assertEquals("test-3", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-1b", 1, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-1b", 1, 1, 1, replyChannel, null));
assertEquals("test-1b", discardChannel.receive(100).getPayload());
}
@@ -215,32 +197,32 @@ public class AggregatorEndpointTests {
QueueChannel discardChannel = new QueueChannel();
this.aggregator.setTrackedCorrelationIdCapacity(3);
this.aggregator.setDiscardChannel(discardChannel);
this.aggregator.handleMessage(createMessage("test-1a", 1, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-1a", 1, 1, 1, replyChannel, null));
assertEquals("test-1a", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-2", 2, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-2", 2, 1, 1, replyChannel, null));
assertEquals("test-2", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-3", 3, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-3", 3, 1, 1, replyChannel, null));
assertEquals("test-3", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-4", 4, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-4", 4, 1, 1, replyChannel, null));
assertEquals("test-4", replyChannel.receive(100).getPayload());
this.aggregator.handleMessage(createMessage("test-1b", 1, 1, 1, replyChannel));
this.aggregator.handleMessage(createMessage("test-1b", 1, 1, 1, replyChannel, null));
assertEquals("test-1b", replyChannel.receive(100).getPayload());
assertNull(discardChannel.receive(0));
}
@Test(expected = MessageHandlingException.class)
public void testExceptionThrownIfNoCorrelationId() throws InterruptedException {
Message<?> message = createMessage("123", null, 2, 1, new QueueChannel());
Message<?> message = createMessage("123", null, 2, 1, new QueueChannel(), null);
this.aggregator.handleMessage(message);
}
@Test
public void testAdditionalMessageAfterCompletion() throws InterruptedException {
QueueChannel replyChannel = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel);
Message<?> message4 = createMessage("abc", "ABC", 3, 3, replyChannel);
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel, null);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel, null);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel, null);
Message<?> message4 = createMessage("abc", "ABC", 3, 3, replyChannel, null);
CountDownLatch latch = new CountDownLatch(4);
this.taskExecutor.execute(new AggregatorTestTask(this.aggregator, message1, latch));
this.taskExecutor.execute(new AggregatorTestTask(this.aggregator, message2, latch));
@@ -257,9 +239,9 @@ public class AggregatorEndpointTests {
this.aggregator = new NullReturningAggregator();
this.aggregator.setTaskScheduler(this.taskScheduler);
QueueChannel replyChannel = new QueueChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel);
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel, null);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel, null);
Message<?> message3 = createMessage("789", "ABC", 3, 3, replyChannel, null);
CountDownLatch latch = new CountDownLatch(3);
AggregatorTestTask task1 = new AggregatorTestTask(aggregator, message1, latch);
this.taskExecutor.execute(task1);
@@ -278,14 +260,16 @@ public class AggregatorEndpointTests {
private static Message<?> createMessage(String payload, Object correlationId,
int sequenceSize, int sequenceNumber, MessageChannel replyChannel) {
Message<String> message = MessageBuilder.withPayload(payload)
.setCorrelationId(correlationId)
.setSequenceSize(sequenceSize)
.setSequenceNumber(sequenceNumber)
.setReplyChannel(replyChannel)
.build();
return message;
int sequenceSize, int sequenceNumber, MessageChannel replyChannel, String predefinedId) {
MessageBuilder<String> builder = MessageBuilder.withPayload(payload)
.setCorrelationId(correlationId)
.setSequenceSize(sequenceSize)
.setSequenceNumber(sequenceNumber)
.setReplyChannel(replyChannel);
if (predefinedId != null) {
builder.setHeader(MessageHeaders.ID, predefinedId);
}
return builder.build();
}

View File

@@ -17,12 +17,16 @@
package org.springframework.integration.aggregator;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.LinkedHashSet;
import java.util.Set;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertTrue;
import org.junit.Test;
import org.springframework.integration.message.StringMessage;
import org.springframework.integration.core.Message;
/**
* @author Mark Fisher
@@ -32,17 +36,17 @@ public class MessageBarrierTests {
@Test
public void testMessageRetrieval() {
MessageBarrier barrier = new MessageBarrier(new LinkedHashMap(), null);
barrier.getMessages().put("1", new StringMessage("test1"));
MessageBarrier barrier = new MessageBarrier(new LinkedHashSet(), null);
barrier.getMessages().add(new StringMessage("test1"));
assertEquals(1, barrier.getMessages().size());
barrier.getMessages().put("2", new StringMessage("test2"));
barrier.getMessages().add(new StringMessage("test2"));
assertEquals(2, barrier.getMessages().size());
}
@Test
public void testTimestamp() {
long before = System.currentTimeMillis();
MessageBarrier barrier = new MessageBarrier(new LinkedHashMap(), null);
MessageBarrier barrier = new MessageBarrier(new LinkedHashSet(), null);
long timestamp = barrier.getTimestamp();
assertTrue(before <= timestamp);
long after = System.currentTimeMillis();
@@ -51,7 +55,7 @@ public class MessageBarrierTests {
@Test
public void testEmptyMessageList() {
MessageBarrier barrier = new MessageBarrier(new LinkedHashMap(), null);
MessageBarrier barrier = new MessageBarrier(new LinkedHashSet(), null);
assertEquals(0, barrier.getMessages().size());
}