Major refactoring of AggregatingMessageHandler and related classes.

This commit is contained in:
Mark Fisher
2008-02-27 02:33:44 +00:00
parent 4c0b4f41f9
commit d2580a3b0b
9 changed files with 491 additions and 245 deletions

View File

@@ -16,12 +16,26 @@
package org.springframework.integration.router;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ArrayBlockingQueue;
import java.util.concurrent.BlockingQueue;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ConcurrentMap;
import java.util.concurrent.Executors;
import java.util.concurrent.ScheduledExecutorService;
import java.util.concurrent.TimeUnit;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.integration.channel.MessageChannel;
import org.springframework.integration.handler.MessageHandler;
import org.springframework.integration.message.Message;
import org.springframework.integration.message.MessageHandlingException;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
/**
* A {@link MessageHandler} implementation that waits for a <em>complete</em>
@@ -35,93 +49,249 @@ import org.springframework.util.Assert;
* <p>
* The default strategy for determining whether a group is complete is based on
* the '<code>sequenceSize</code>' property of the header. Alternatively, a
* custom implementation of the {@link RoutingBarrierCompletionStrategy} may be
* provided.
* custom implementation of the {@link CompletionStrategy} may be provided.
* <p>
* The '<code>timeout</code>' value determines how long to wait for the
* complete group after the arrival of the first {@link Message} of the group.
* The default value is 1 minute. If the timeout elapses prior to completion,
* the handler will throw a {@link MessageHandlingException} by default. To
* prevent the exception and aggregate the group even when incomplete, set the
* '<code>shouldFailOnTimeout</code>' property to '<code>false</code>'.
* then Messages with that timed-out 'correlationId' will be sent to the
* 'discardChannel' if provided.
*
* @author Mark Fisher
* @author Marius Bogoevici
*/
public class AggregatingMessageHandler implements MessageHandler {
public class AggregatingMessageHandler implements MessageHandler, InitializingBean {
private long timeout = 60000;
private final Log logger = LogFactory.getLog(this.getClass());
private boolean shouldFailOnTimeout = true;
private final Aggregator aggregator;
private Aggregator aggregator;
private volatile MessageChannel defaultReplyChannel;
private RoutingBarrierCompletionStrategy completionStrategy = new SequenceSizeCompletionStrategy();
private volatile MessageChannel discardChannel;
private ConcurrentHashMap<Object, RoutingBarrier> barriers = new ConcurrentHashMap<Object, RoutingBarrier>();
private volatile long sendTimeout = 1000;
private volatile CompletionStrategy completionStrategy = new SequenceSizeCompletionStrategy();
private final ConcurrentMap<Object, AggregationBarrier> barriers = new ConcurrentHashMap<Object, AggregationBarrier>();
private volatile long timeout = 60000;
private volatile boolean sendPartialResultOnTimeout = false;
private volatile long reaperInterval = 1000;
private volatile int trackedCorrelationIdCapacity = 1000;
private volatile BlockingQueue<Object> trackedCorrelationIds;
private final ScheduledExecutorService executor;
private volatile boolean initialized;
/**
* Create a handler that delegates to the provided aggregator to combine a
* group of messages into a single message.
* group of messages into a single message. The executor will be used for
* scheduling a background maintenance thread. If <code>null</code>, a new
* single-threaded executor will be created.
*/
public AggregatingMessageHandler(Aggregator aggregator) {
public AggregatingMessageHandler(Aggregator aggregator, ScheduledExecutorService executor) {
Assert.notNull(aggregator, "'aggregator' must not be null");
this.aggregator = aggregator;
this.executor = (executor != null) ? executor : Executors.newSingleThreadScheduledExecutor();
}
public AggregatingMessageHandler(Aggregator aggregator) {
this(aggregator, null);
}
/**
* Set the default channel for sending aggregated Messages. Note that
* precedence will be given to the 'returnAddress' of the aggregated
* message itself, then to the 'returnAddress' of the original message.
*/
public void setDefaultReplyChannel(MessageChannel defaultReplyChannel) {
this.defaultReplyChannel = defaultReplyChannel;
}
/**
* Specify a channel for sending Messages that arrive after their aggregation
* group has either completed or timed-out.
*/
public void setDiscardChannel(MessageChannel discardChannel) {
this.discardChannel = discardChannel;
}
/**
* Set the timeout for sending aggregation results and discarded Messages.
*/
public void setSendTimeout(long sendTimeout) {
this.sendTimeout = sendTimeout;
}
/**
* Specify whether to aggregate and send the resulting Message when the
* timeout elapses prior to the CompletionStrategy.
*/
public void setSendPartialResultOnTimeout(boolean sendPartialResultOnTimeout) {
this.sendPartialResultOnTimeout = sendPartialResultOnTimeout;
}
/**
* Set the interval in milliseconds for the reaper thread. Default is 1000.
*/
public void setReaperInterval(long reaperInterval) {
Assert.isTrue(reaperInterval > 0, "'reaperInterval' must be a positive value");
this.reaperInterval = reaperInterval;
}
/**
* Set the number of completed correlationIds to track. Default is 1000.
*/
public void setTrackedCorrelationIdCapacity(int trackedCorrelationIdCapacity) {
Assert.isTrue(trackedCorrelationIdCapacity > 0, "'trackedCorrelationIdCapacity' must be a positive value");
this.trackedCorrelationIdCapacity = trackedCorrelationIdCapacity;
}
/**
* Initialize this handler.
*/
public void afterPropertiesSet() {
this.trackedCorrelationIds = new ArrayBlockingQueue<Object>(this.trackedCorrelationIdCapacity);
this.executor.scheduleWithFixedDelay(new ReaperTask(),
this.reaperInterval, this.reaperInterval, TimeUnit.MILLISECONDS);
this.initialized = true;
}
/**
* Strategy to determine whether the group of messages is complete.
*/
public void setCompletionStrategy(RoutingBarrierCompletionStrategy completionStrategy) {
public void setCompletionStrategy(CompletionStrategy completionStrategy) {
Assert.notNull(completionStrategy, "'completionStrategy' must not be null");
this.completionStrategy = completionStrategy;
}
/**
* Maximum time to wait (in milliseconds) for the completion strategy to
* become true.
* become true. The default is 60000 (1 minute).
*/
public void setTimeout(long timeout) {
Assert.isTrue(timeout >= 0, "'timeout' must not be negative");
this.timeout = timeout;
}
/**
* Specify whether this handler should throw a {@link MessageHandlingException}
* when a message group does not reach completion within the allotted time. The
* default is '<code>true</code>'. Setting this to '<code>false</code>' will cause
* the {@link Aggregator} to be invoked even when the group is incomplete.
*/
public void setShouldFailOnTimeout(boolean setShouldFailOnTimeout) {
this.shouldFailOnTimeout = setShouldFailOnTimeout;
}
public Message<?> handle(Message<?> message) {
if (!this.initialized) {
this.afterPropertiesSet();
}
Object correlationId = message.getHeader().getCorrelationId();
if (correlationId == null) {
throw new MessageHandlingException(this.getClass().getSimpleName() +
" requires the 'correlationId' property");
}
RoutingBarrier barrier = barriers.putIfAbsent(correlationId, new RoutingBarrier(this.completionStrategy));
if (barrier == null) {
try {
barrier = barriers.get(correlationId);
barrier.addMessage(message);
if (!barrier.waitForCompletion(this.timeout) && this.shouldFailOnTimeout) {
throw new MessageHandlingException("aggregation did not complete "
+ "within the allotted time limit of " + this.timeout + " milliseconds");
}
Message<?> result = aggregator.aggregate(barrier.getMessages());
return result;
if (this.trackedCorrelationIds.contains(correlationId)) {
if (logger.isDebugEnabled()) {
logger.debug("Aggregation for correlationId '" + correlationId +
"' has already completed or timed out.");
}
finally {
this.barriers.remove(correlationId);
this.sendToDiscardChannelIfAvailable(message);
return null;
}
AggregationBarrier barrier = barriers.putIfAbsent(correlationId,
new AggregationBarrier(this.completionStrategy));
if (barrier == null) {
barrier = barriers.get(correlationId);
}
List<Message<?>> releasedMessages = barrier.addAndRelease(message);
if (CollectionUtils.isEmpty(releasedMessages)) {
return null;
}
this.removeBarrier(correlationId);
this.aggregationCompleted(correlationId, releasedMessages);
return null;
}
private void sendToDiscardChannelIfAvailable(Message<?> message) {
if (this.discardChannel != null) {
if (!this.discardChannel.send(message, this.sendTimeout)) {
if (logger.isWarnEnabled()) {
logger.warn("unable to send to 'discardChannel', message: " + message);
}
}
}
else {
barriers.get(correlationId).addMessage(message);
return null;
}
private void aggregationCompleted(Object correlationId, List<Message<?>> messages) {
if (CollectionUtils.isEmpty(messages)) {
if (logger.isDebugEnabled()) {
logger.debug("no messages to aggregate");
}
return;
}
Message<?> result = aggregator.aggregate(messages);
MessageChannel replyChannel = this.resolveReplyChannelFromMessage(result);
if (replyChannel == null) {
replyChannel = this.resolveReplyChannelFromMessage(messages.get(0));
if (replyChannel == null) {
replyChannel = this.defaultReplyChannel;
}
}
if (replyChannel != null) {
replyChannel.send(result, this.sendTimeout);
}
else if (logger.isWarnEnabled()) {
logger.warn("unable to determine reply channel for aggregation result: " + result);
}
}
private void removeBarrier(Object correlationId) {
if (this.barriers.remove(correlationId) != null) {
synchronized (this.trackedCorrelationIds) {
boolean added = this.trackedCorrelationIds.offer(correlationId);
if (!added) {
this.trackedCorrelationIds.poll();
this.trackedCorrelationIds.offer(correlationId);
}
}
}
}
private MessageChannel resolveReplyChannelFromMessage(Message<?> message) {
Object returnAddress = message.getHeader().getReturnAddress();
if (returnAddress != null) {
if (returnAddress instanceof MessageChannel) {
return (MessageChannel) returnAddress;
}
if (logger.isWarnEnabled()) {
logger.warn("Aggregator can only reply to a 'returnAddress' of type MessageChannel.");
}
}
return null;
}
private class ReaperTask implements Runnable {
public void run() {
long currentTime = System.currentTimeMillis();
for (Map.Entry<Object, AggregationBarrier> entry : barriers.entrySet()) {
if (currentTime - entry.getValue().getTimestamp() >= timeout) {
Object correlationId = entry.getKey();
List<Message<?>> messages = entry.getValue().getMessages();
removeBarrier(correlationId);
if (sendPartialResultOnTimeout) {
aggregationCompleted(correlationId, messages);
}
else {
for (Message<?> message : messages) {
sendToDiscardChannelIfAvailable(message);
}
}
}
}
}
}

View File

@@ -0,0 +1,96 @@
/*
* Copyright 2002-2007 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.integration.router;
import java.util.List;
import java.util.concurrent.CopyOnWriteArrayList;
import java.util.concurrent.locks.ReentrantLock;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.integration.message.Message;
import org.springframework.util.Assert;
/**
* MessageBarrier implementation for message aggregation. Delegates to a
* {@link CompletionStrategy} to determine when the group of messages is ready
* for aggregation.
*
* @author Marius Bogoevici
* @author Mark Fisher
*/
public class AggregationBarrier implements MessageBarrier {
private final Log logger = LogFactory.getLog(this.getClass());
private final List<Message<?>> messages = new CopyOnWriteArrayList<Message<?>>();
private final CompletionStrategy completionStrategy;
private volatile boolean complete = false;
private final ReentrantLock lock = new ReentrantLock();
private final long timestamp = System.currentTimeMillis();
/**
* Create an AggregationBarrier with the given {@link CompletionStrategy}.
*/
public AggregationBarrier(CompletionStrategy completionStrategy) {
Assert.notNull(completionStrategy, "'completionStrategy' must not be null");
this.completionStrategy = completionStrategy;
}
/**
* Returns the creation time of this barrier as the number of milliseconds
* since January 1, 1970.
* @see java.lang.System#currentTimeMillis()
*/
public long getTimestamp() {
return this.timestamp;
}
/**
* Adds a message to the aggregation group and releases <em>if complete</em>.
* Otherwise, the return value will be <code>null</code>.
*/
public List<Message<?>> addAndRelease(Message<?> message) {
try {
this.lock.lock();
if (this.complete) {
if (logger.isDebugEnabled()) {
logger.debug("Message received after aggregation has already completed: " + message);
}
return null;
}
this.messages.add(message);
boolean complete = completionStrategy.isComplete(this.messages);
return (complete) ? this.messages : null;
}
finally {
this.lock.unlock();
}
}
public List<Message<?>> getMessages() {
return this.messages;
}
}

View File

@@ -25,9 +25,9 @@ import org.springframework.integration.message.Message;
* completion (i.e. can trip a barrier).
*
* @author Mark Fisher
* @see RoutingBarrier
* @see AggregationBarrier
*/
public interface RoutingBarrierCompletionStrategy {
public interface CompletionStrategy {
boolean isComplete(List<Message<?>> messages);

View File

@@ -0,0 +1,34 @@
/*
* Copyright 2002-2007 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.integration.router;
import java.util.List;
import org.springframework.integration.message.Message;
/**
* Common interface for routing components that release a list of
* {@link Message Messages} based upon a condition that is met when a
* {@link Message} arrives.
*
* @author Mark Fisher
*/
public interface MessageBarrier {
List<Message<?>> addAndRelease(Message<?> message);
}

View File

@@ -1,104 +0,0 @@
/*
* Copyright 2002-2007 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.integration.router;
import java.util.List;
import java.util.concurrent.CopyOnWriteArrayList;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.locks.Condition;
import java.util.concurrent.locks.ReentrantLock;
import org.springframework.integration.message.Message;
import org.springframework.util.Assert;
/**
* A rendezvous point for {@link Message Messages} that delegates to a
* {@link RoutingBarrierCompletionStrategy} to determine when a
* <em>complete</em> message group is available.
*
* @author Mark Fisher
*/
public class RoutingBarrier {
private final List<Message<?>> messages = new CopyOnWriteArrayList<Message<?>>();
private final RoutingBarrierCompletionStrategy completionStrategy;
private volatile boolean complete = false;
private final ReentrantLock lock = new ReentrantLock();
private final Condition condition = lock.newCondition();
public RoutingBarrier(RoutingBarrierCompletionStrategy completionStrategy) {
Assert.notNull(completionStrategy, "'completionStrategy' must not be null");
this.completionStrategy = completionStrategy;
}
public void addMessage(Message<?> message) {
this.messages.add(message);
if (this.completionStrategy.isComplete(this.messages)) {
try {
this.lock.lock();
if (!this.complete) {
this.complete = true;
this.condition.signalAll();
}
}
finally {
this.lock.unlock();
}
}
}
public boolean waitForCompletion(long timeout) {
if (this.complete) {
return true;
}
lock.lock();
try {
if (this.complete) {
return true;
}
if (timeout >= 0) {
return this.condition.await(timeout, TimeUnit.MILLISECONDS);
}
else {
this.condition.await();
return true;
}
}
catch (InterruptedException e) {
Thread.currentThread().interrupt();
return false;
}
finally {
lock.unlock();
}
}
public boolean isComplete() {
return this.complete;
}
public List<Message<?>> getMessages() {
return this.messages;
}
}

View File

@@ -22,13 +22,13 @@ import org.springframework.integration.message.Message;
import org.springframework.util.CollectionUtils;
/**
* An implementation of {@link RoutingBarrierCompletionStrategy} that simply
* An implementation of {@link CompletionStrategy} that simply
* compares the current size of the message list to the expected 'sequenceSize'
* according to the first {@link Message} in the list.
*
* @author Mark Fisher
*/
public class SequenceSizeCompletionStrategy implements RoutingBarrierCompletionStrategy {
public class SequenceSizeCompletionStrategy implements CompletionStrategy {
public boolean isComplete(List<Message<?>> messages) {
if (CollectionUtils.isEmpty(messages)) {

View File

@@ -68,26 +68,31 @@ public class AggregatingMessageHandlerTests {
}
@Test
public void testShouldFailOnTimeoutByDefault() throws InterruptedException {
public void testShouldNotSendPartialResultOnTimeoutByDefault() throws InterruptedException {
SimpleChannel discardChannel = new SimpleChannel();
AggregatingMessageHandler aggregator = new AggregatingMessageHandler(new TestAggregator());
aggregator.setTimeout(10);
aggregator.setTimeout(50);
aggregator.setReaperInterval(10);
aggregator.setDiscardChannel(discardChannel);
SimpleChannel replyChannel = new SimpleChannel();
Message<?> message1 = createMessage("123", "ABC", 2, 1, replyChannel);
Message<?> message = createMessage("123", "ABC", 2, 1, replyChannel);
CountDownLatch latch = new CountDownLatch(1);
AggregatorTestTask task = new AggregatorTestTask(aggregator, message1, latch);
AggregatorTestTask task = new AggregatorTestTask(aggregator, message, latch);
executor.execute(task);
latch.await(1000, TimeUnit.MILLISECONDS);
Message<?> reply = replyChannel.receive(0);
assertNull(reply);
assertNotNull(task.getException());
assertEquals(MessageHandlingException.class, task.getException().getClass());
Message<?> discardedMessage = discardChannel.receive(500);
assertNotNull(discardedMessage);
assertEquals(message, discardedMessage);
}
@Test
public void testShouldFailOnTimeoutFalse() throws InterruptedException {
public void testShouldSendPartialResultOnTimeoutTrue() throws InterruptedException {
AggregatingMessageHandler aggregator = new AggregatingMessageHandler(new TestAggregator());
aggregator.setTimeout(10);
aggregator.setShouldFailOnTimeout(false);
aggregator.setTimeout(50);
aggregator.setReaperInterval(10);
aggregator.setSendPartialResultOnTimeout(true);
SimpleChannel replyChannel = new SimpleChannel();
Message<?> message1 = createMessage("123", "ABC", 3, 1, replyChannel);
Message<?> message2 = createMessage("456", "ABC", 3, 2, replyChannel);
@@ -131,6 +136,62 @@ public class AggregatingMessageHandlerTests {
assertEquals("abcdefghi", reply2.getPayload());
}
@Test
public void testDiscardChannelForTrackedCorrelationId() {
SimpleChannel replyChannel = new SimpleChannel();
SimpleChannel discardChannel = new SimpleChannel();
AggregatingMessageHandler aggregator = new AggregatingMessageHandler(new TestAggregator());
aggregator.setDiscardChannel(discardChannel);
aggregator.handle(createMessage("test-1a", 1, 1, 1, replyChannel));
assertEquals("test-1a", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-1b", 1, 1, 1, replyChannel));
assertEquals("test-1b", discardChannel.receive(100).getPayload());
}
@Test
public void testTrackedCorrelationIdsCapacityAtLimit() {
SimpleChannel replyChannel = new SimpleChannel();
SimpleChannel discardChannel = new SimpleChannel();
AggregatingMessageHandler aggregator = new AggregatingMessageHandler(new TestAggregator());
aggregator.setTrackedCorrelationIdCapacity(3);
aggregator.setDiscardChannel(discardChannel);
aggregator.handle(createMessage("test-1a", 1, 1, 1, replyChannel));
assertEquals("test-1a", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-2", 2, 1, 1, replyChannel));
assertEquals("test-2", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-3", 3, 1, 1, replyChannel));
assertEquals("test-3", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-1b", 1, 1, 1, replyChannel));
assertEquals("test-1b", discardChannel.receive(100).getPayload());
}
@Test
public void testTrackedCorrelationIdsCapacityPassesLimit() {
SimpleChannel replyChannel = new SimpleChannel();
SimpleChannel discardChannel = new SimpleChannel();
AggregatingMessageHandler aggregator = new AggregatingMessageHandler(new TestAggregator());
aggregator.setTrackedCorrelationIdCapacity(3);
aggregator.setDiscardChannel(discardChannel);
aggregator.handle(createMessage("test-1a", 1, 1, 1, replyChannel));
assertEquals("test-1a", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-2", 2, 1, 1, replyChannel));
assertEquals("test-2", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-3", 3, 1, 1, replyChannel));
assertEquals("test-3", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-4", 4, 1, 1, replyChannel));
assertEquals("test-4", replyChannel.receive(100).getPayload());
aggregator.handle(createMessage("test-1b", 1, 1, 1, replyChannel));
assertEquals("test-1b", replyChannel.receive(100).getPayload());
assertNull(discardChannel.receive(0));
}
@Test(expected=MessageHandlingException.class)
public void testExceptionThrownIfNoCorrelationId() throws InterruptedException {
AggregatingMessageHandler aggregator = new AggregatingMessageHandler(new TestAggregator());
Message<?> message = createMessage("123", null, 2, 1, new SimpleChannel());
aggregator.handle(message);
}
private static Message<?> createMessage(String payload, Object correlationId,
int sequenceSize, int sequenceNumber, MessageChannel replyChannel) {

View File

@@ -0,0 +1,76 @@
/*
* Copyright 2002-2007 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.integration.router;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertNotNull;
import static org.junit.Assert.assertNull;
import static org.junit.Assert.assertTrue;
import java.util.List;
import org.junit.Test;
import org.springframework.integration.message.Message;
import org.springframework.integration.message.StringMessage;
/**
* @author Mark Fisher
*/
public class AggregationBarrierTests {
@Test
public void testBasicCompletionCheck() {
AggregationBarrier barrier = new AggregationBarrier(new TwoMessageCompletionStrategy());
assertNull(barrier.addAndRelease(new StringMessage("test1")));
assertNotNull(barrier.addAndRelease(new StringMessage("test2")));
}
@Test
public void testMessageRetrieval() {
AggregationBarrier barrier = new AggregationBarrier(new TwoMessageCompletionStrategy());
barrier.addAndRelease(new StringMessage("test1"));
assertEquals(1, barrier.getMessages().size());
barrier.addAndRelease(new StringMessage("test2"));
assertEquals(2, barrier.getMessages().size());
}
@Test
public void testTimestamp() {
long before = System.currentTimeMillis();
AggregationBarrier barrier = new AggregationBarrier(new TwoMessageCompletionStrategy());
long timestamp = barrier.getTimestamp();
assertTrue(before <= timestamp);
long after = System.currentTimeMillis();
assertTrue(after >= timestamp);
}
@Test
public void testEmptyMessageList() {
AggregationBarrier barrier = new AggregationBarrier(new TwoMessageCompletionStrategy());
assertEquals(0, barrier.getMessages().size());
}
private static class TwoMessageCompletionStrategy implements CompletionStrategy {
public boolean isComplete(List<Message<?>> messages) {
return (messages.size() == 2);
}
}
}

View File

@@ -1,87 +0,0 @@
/*
* Copyright 2002-2007 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.integration.router;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertFalse;
import static org.junit.Assert.assertTrue;
import java.util.List;
import org.junit.Test;
import org.springframework.integration.message.Message;
import org.springframework.integration.message.StringMessage;
/**
* @author Mark Fisher
*/
public class RoutingBarrierTests {
@Test
public void testBasicCompletionCheck() {
RoutingBarrier barrier = new RoutingBarrier(new TwoMessageCompletionStrategy());
barrier.addMessage(new StringMessage("test1"));
assertFalse(barrier.isComplete());
barrier.addMessage(new StringMessage("test2"));
assertTrue(barrier.isComplete());
}
@Test
public void testMessageRetrieval() {
RoutingBarrier barrier = new RoutingBarrier(new TwoMessageCompletionStrategy());
barrier.addMessage(new StringMessage("test1"));
assertEquals(1, barrier.getMessages().size());
barrier.addMessage(new StringMessage("test2"));
assertEquals(2, barrier.getMessages().size());
}
@Test
public void testWaitForCompletionTimesOut() {
RoutingBarrier barrier = new RoutingBarrier(new TwoMessageCompletionStrategy());
barrier.addMessage(new StringMessage("test1"));
assertFalse(barrier.isComplete());
assertFalse(barrier.waitForCompletion(10));
}
@Test
public void testWaitForCompletionReturnsTrueImmediately() {
RoutingBarrier barrier = new RoutingBarrier(new TwoMessageCompletionStrategy());
barrier.addMessage(new StringMessage("test1"));
assertFalse(barrier.isComplete());
barrier.addMessage(new StringMessage("test2"));
assertTrue(barrier.waitForCompletion(0));
assertTrue(barrier.isComplete());
}
@Test
public void testEmptyMessageList() {
RoutingBarrier barrier = new RoutingBarrier(new TwoMessageCompletionStrategy());
assertFalse(barrier.isComplete());
assertFalse(barrier.waitForCompletion(0));
assertEquals(0, barrier.getMessages().size());
}
private static class TwoMessageCompletionStrategy implements RoutingBarrierCompletionStrategy {
public boolean isComplete(List<Message<?>> messages) {
return (messages.size() == 2);
}
}
}