diff --git a/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractCorrelatingMessageHandler.java b/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractCorrelatingMessageHandler.java index d331560ad5..5cbf688acb 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractCorrelatingMessageHandler.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractCorrelatingMessageHandler.java @@ -752,6 +752,8 @@ public abstract class AbstractCorrelatingMessageHandler extends AbstractMessageP protected static class SequenceAwareMessageGroup extends SimpleMessageGroup { + private final SimpleMessageGroup sourceGroup; + public SequenceAwareMessageGroup(MessageGroup messageGroup) { /* * Since this group is temporary, and never added to, we simply use the @@ -760,6 +762,12 @@ public abstract class AbstractCorrelatingMessageHandler extends AbstractMessageP */ super(messageGroup.getMessages(), null, messageGroup.getGroupId(), messageGroup.getTimestamp(), messageGroup.isComplete(), true); + if (messageGroup instanceof SimpleMessageGroup) { + this.sourceGroup = (SimpleMessageGroup) messageGroup; + } + else { + this.sourceGroup = null; + } } /** @@ -773,20 +781,25 @@ public abstract class AbstractCorrelatingMessageHandler extends AbstractMessageP if (this.size() == 0) { return true; } - IntegrationMessageHeaderAccessor messageHeaderAccessor = new IntegrationMessageHeaderAccessor(message); - Integer messageSequenceNumber = messageHeaderAccessor.getSequenceNumber(); + Integer messageSequenceNumber = message.getHeaders().get(IntegrationMessageHeaderAccessor.SEQUENCE_NUMBER, + Integer.class); if (messageSequenceNumber != null && messageSequenceNumber > 0) { - Integer messageSequenceSize = messageHeaderAccessor.getSequenceSize(); - return messageSequenceSize.equals(this.getSequenceSize()) - && !this.containsSequenceNumber(this.getMessages(), messageSequenceNumber); + Integer messageSequenceSize = message.getHeaders().get(IntegrationMessageHeaderAccessor.SEQUENCE_SIZE, + Integer.class); + if (messageSequenceSize == null) { + messageSequenceSize = Integer.valueOf(0); + } + return messageSequenceSize.equals(getSequenceSize()) + && !(this.sourceGroup != null ? this.sourceGroup.containsSequence(messageSequenceNumber) + : containsSequenceNumber(this.getMessages(), messageSequenceNumber)); } return true; } private boolean containsSequenceNumber(Collection> messages, Integer messageSequenceNumber) { for (Message member : messages) { - Integer memberSequenceNumber = new IntegrationMessageHeaderAccessor(member).getSequenceNumber(); - if (messageSequenceNumber.equals(memberSequenceNumber)) { + if (messageSequenceNumber.equals(member.getHeaders().get( + IntegrationMessageHeaderAccessor.SEQUENCE_NUMBER, Integer.class))) { return true; } } diff --git a/spring-integration-core/src/main/java/org/springframework/integration/store/SimpleMessageGroup.java b/spring-integration-core/src/main/java/org/springframework/integration/store/SimpleMessageGroup.java index f5a5b4b4ad..b3315d8039 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/store/SimpleMessageGroup.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/store/SimpleMessageGroup.java @@ -18,8 +18,10 @@ package org.springframework.integration.store; import java.util.Collection; import java.util.Collections; +import java.util.HashSet; import java.util.Iterator; import java.util.LinkedHashSet; +import java.util.Set; import org.springframework.integration.IntegrationMessageHeaderAccessor; import org.springframework.messaging.Message; @@ -44,6 +46,8 @@ public class SimpleMessageGroup implements MessageGroup { private final Collection> messages; + private final Set sequences = new HashSet<>(); + private final long timestamp; private volatile int lastReleasedMessageSequence; @@ -114,6 +118,7 @@ public class SimpleMessageGroup implements MessageGroup { @Override public boolean remove(Message message) { + this.sequences.remove(message.getHeaders().get(IntegrationMessageHeaderAccessor.SEQUENCE_NUMBER)); return this.messages.remove(message); } @@ -123,6 +128,8 @@ public class SimpleMessageGroup implements MessageGroup { } private boolean addMessage(Message message) { + Integer sequence = message.getHeaders().get(IntegrationMessageHeaderAccessor.SEQUENCE_NUMBER, Integer.class); + this.sequences.add(sequence != null ? sequence : 0); return this.messages.add(message); } @@ -175,6 +182,18 @@ public class SimpleMessageGroup implements MessageGroup { @Override public void clear() { this.messages.clear(); + this.sequences.clear(); + } + + /** + * Return true if a message with this sequence number header exists in + * the group. + * @param sequence the sequence number. + * @return true if it exists. + * @since 4.3.7 + */ + public boolean containsSequence(Integer sequence) { + return this.sequences.contains(sequence); } @Override diff --git a/spring-integration-core/src/test/java/org/springframework/integration/aggregator/AggregatorTests.java b/spring-integration-core/src/test/java/org/springframework/integration/aggregator/AggregatorTests.java index b51f29241c..052c6d06b8 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/aggregator/AggregatorTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/aggregator/AggregatorTests.java @@ -17,6 +17,7 @@ package org.springframework.integration.aggregator; import static org.hamcrest.CoreMatchers.is; +import static org.hamcrest.Matchers.lessThan; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertNull; @@ -129,6 +130,55 @@ public class AggregatorTests { assertEquals(60000, result.size()); } + @Test + public void testAggPerfDefaultPartial() throws InterruptedException, ExecutionException, TimeoutException { + AggregatingMessageHandler handler = new AggregatingMessageHandler(new DefaultAggregatingMessageGroupProcessor()); + handler.setCorrelationStrategy(message -> "foo"); + handler.setReleasePartialSequences(true); + DirectChannel outputChannel = new DirectChannel(); + handler.setOutputChannel(outputChannel); + + final CompletableFuture> resultFuture = new CompletableFuture<>(); + outputChannel.subscribe(message -> { + Collection payload = (Collection) message.getPayload(); + logger.warn("Received " + payload.size()); + resultFuture.complete(payload); + }); + + SimpleMessageStore store = new SimpleMessageStore(); + + SimpleMessageGroupFactory messageGroupFactory = + new SimpleMessageGroupFactory(SimpleMessageGroupFactory.GroupType.BLOCKING_QUEUE); + + store.setMessageGroupFactory(messageGroupFactory); + + handler.setMessageStore(store); + + + StopWatch stopwatch = new StopWatch(); + stopwatch.start(); + for (int i = 0; i < 120000; i++) { + if (i % 10000 == 0) { + stopwatch.stop(); + logger.warn("Sent " + i + " in " + stopwatch.getTotalTimeSeconds() + + " (10k in " + stopwatch.getLastTaskTimeMillis() + "ms)"); + stopwatch.start(); + } + handler.handleMessage(MessageBuilder.withPayload("foo") + .setSequenceSize(120000) + .setSequenceNumber(i + 1) + .build()); + } + stopwatch.stop(); + logger.warn("Sent " + 120000 + " in " + stopwatch.getTotalTimeSeconds() + + " (10k in " + stopwatch.getLastTaskTimeMillis() + "ms)"); + + Collection result = resultFuture.get(10, TimeUnit.SECONDS); + assertNotNull(result); + assertEquals(120000, result.size()); + assertThat(stopwatch.getTotalTimeSeconds(), lessThan(60.0)); // actually < 2.0, was many minutes + } + @Test public void testCustomAggPerf() throws InterruptedException, ExecutionException, TimeoutException { class CustomHandler extends AbstractMessageHandler { diff --git a/spring-integration-core/src/test/java/org/springframework/integration/store/SimpleMessageGroupTests.java b/spring-integration-core/src/test/java/org/springframework/integration/store/SimpleMessageGroupTests.java index 8af222874a..55afe2c024 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/store/SimpleMessageGroupTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/store/SimpleMessageGroupTests.java @@ -20,20 +20,20 @@ import static org.hamcrest.CoreMatchers.is; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertThat; import static org.junit.Assert.assertTrue; +import static org.mockito.BDDMockito.willReturn; import static org.mockito.Mockito.mock; import java.lang.reflect.Constructor; import java.util.ArrayList; import java.util.Collection; -import java.util.Collections; -import java.util.HashSet; import java.util.List; +import java.util.Map; import org.junit.Test; -import org.springframework.beans.DirectFieldAccessor; import org.springframework.integration.support.MessageBuilder; import org.springframework.messaging.Message; +import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.support.GenericMessage; import org.springframework.util.StopWatch; @@ -48,7 +48,9 @@ public class SimpleMessageGroupTests { private final Object key = new Object(); - private SimpleMessageGroup group = new SimpleMessageGroup(Collections.>emptyList(), key); + private final SimpleMessageGroup group = new SimpleMessageGroup(new ArrayList>(), key); + + private MessageGroup sequenceAwareGroup; @SuppressWarnings("unchecked") public void prepareForSequenceAwareMessageGroup() throws Exception { @@ -56,8 +58,7 @@ public class SimpleMessageGroupTests { (Class) Class.forName("org.springframework.integration.aggregator.AbstractCorrelatingMessageHandler$SequenceAwareMessageGroup"); Constructor ctr = clazz.getDeclaredConstructor(MessageGroup.class); ctr.setAccessible(true); - group = ctr.newInstance(group); - new DirectFieldAccessor(group).setPropertyValue("messages", new HashSet>()); + this.sequenceAwareGroup = ctr.newInstance(this.group); } @Test @@ -65,10 +66,11 @@ public class SimpleMessageGroupTests { prepareForSequenceAwareMessageGroup(); final Message message1 = MessageBuilder.withPayload("test").setSequenceNumber(1).build(); final Message message2 = MessageBuilder.fromMessage(message1).setSequenceNumber(1).build(); - assertThat(group.canAdd(message1), is(true)); - group.add(message1); - group.add(message2); - assertThat(group.canAdd(message1), is(false)); + assertThat(this.sequenceAwareGroup.canAdd(message1), is(true)); + this.group.add(message1); + this.group.add(message2); + prepareForSequenceAwareMessageGroup(); + assertThat(this.sequenceAwareGroup.canAdd(message1), is(false)); } @Test @@ -76,16 +78,20 @@ public class SimpleMessageGroupTests { prepareForSequenceAwareMessageGroup(); final Message message1 = MessageBuilder.withPayload("test").build(); final Message message2 = MessageBuilder.fromMessage(message1).build(); - assertThat(group.canAdd(message1), is(true)); - group.add(message1); - group.add(message2); - assertThat(group.canAdd(message1), is(true)); + assertThat(this.sequenceAwareGroup.canAdd(message1), is(true)); + this.group.add(message1); + this.group.add(message2); + prepareForSequenceAwareMessageGroup(); + assertThat(this.sequenceAwareGroup.canAdd(message1), is(true)); } + @SuppressWarnings("unchecked") @Test // should not fail with NPE (see INT-2666) public void shouldIgnoreNullValuesWhenInitializedWithCollectionContainingNulls() throws Exception { Message m1 = mock(Message.class); + willReturn(new MessageHeaders(mock(Map.class))).given(m1).getHeaders(); Message m2 = mock(Message.class); + willReturn(new MessageHeaders(mock(Map.class))).given(m2).getHeaders(); final List> messages = new ArrayList>(); messages.add(m1); messages.add(null);