diff --git a/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractAggregatingMessageGroupProcessor.java b/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractAggregatingMessageGroupProcessor.java index 1e1042d218..6fc09dac8c 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractAggregatingMessageGroupProcessor.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/aggregator/AbstractAggregatingMessageGroupProcessor.java @@ -15,10 +15,13 @@ package org.springframework.integration.aggregator; import java.lang.annotation.Annotation; import java.lang.reflect.Method; +import java.util.ArrayList; +import java.util.Arrays; import java.util.Collection; import java.util.HashMap; import java.util.HashSet; import java.util.Iterator; +import java.util.List; import java.util.Map; import java.util.Set; import java.util.concurrent.atomic.AtomicReference; @@ -33,6 +36,7 @@ import org.springframework.integration.annotation.Header; import org.springframework.integration.core.MessageBuilder; import org.springframework.integration.core.MessageChannel; import org.springframework.integration.core.MessagingTemplate; +import org.springframework.integration.splitter.AbstractMessageSplitter; import org.springframework.integration.store.MessageGroup; import org.springframework.util.Assert; import org.springframework.util.ReflectionUtils; @@ -50,8 +54,7 @@ public abstract class AbstractAggregatingMessageGroupProcessor implements Messag private final Log logger = LogFactory.getLog(this.getClass()); @SuppressWarnings("unchecked") - public final void processAndSend(MessageGroup group, MessagingTemplate channelTemplate, - MessageChannel outputChannel) { + public final void processAndSend(MessageGroup group, MessagingTemplate channelTemplate, MessageChannel outputChannel) { Assert.notNull(group, "MessageGroup must not be null"); Assert.notNull(outputChannel, "'outputChannel' must not be null"); Object payload = this.aggregatePayloads(group); @@ -74,13 +77,31 @@ public abstract class AbstractAggregatingMessageGroupProcessor implements Messag MessageHeaders currentHeaders = message.getHeaders(); for (String key : currentHeaders.keySet()) { if (MessageHeaders.ID.equals(key) || MessageHeaders.TIMESTAMP.equals(key) - || MessageHeaders.SEQUENCE_SIZE.equals(key)) { + || MessageHeaders.SEQUENCE_SIZE.equals(key) || MessageHeaders.SEQUENCE_NUMBER.equals(key) + || MessageHeaders.CORRELATION_ID.equals(key)) { + continue; + } + if (AbstractMessageSplitter.SEQUENCE_DETAILS.equals(key) && !aggregatedHeaders.containsKey(MessageHeaders.CORRELATION_ID)) { + @SuppressWarnings("unchecked") + List incomingSequenceDetails = new ArrayList(currentHeaders + .get(key, List.class)); + Object[] sequenceDetails = incomingSequenceDetails.remove(incomingSequenceDetails.size() - 1); + Assert.state(sequenceDetails.length == 3, "Wrong sequence details (not created by splitter?): " + + Arrays.asList(sequenceDetails)); + aggregatedHeaders.put(MessageHeaders.CORRELATION_ID, sequenceDetails[0]); + aggregatedHeaders.put(MessageHeaders.SEQUENCE_NUMBER, sequenceDetails[1]); + aggregatedHeaders.put(MessageHeaders.SEQUENCE_SIZE, sequenceDetails[2]); + if (!incomingSequenceDetails.isEmpty()) { + aggregatedHeaders.put(AbstractMessageSplitter.SEQUENCE_DETAILS, incomingSequenceDetails); + } + System.err.println(aggregatedHeaders); continue; } Object value = currentHeaders.get(key); if (!aggregatedHeaders.containsKey(key)) { aggregatedHeaders.put(key, value); - } else if (!value.equals(aggregatedHeaders.get(key))) { + } + else if (!value.equals(aggregatedHeaders.get(key))) { conflictKeys.add(key); } } diff --git a/spring-integration-core/src/main/java/org/springframework/integration/splitter/AbstractMessageSplitter.java b/spring-integration-core/src/main/java/org/springframework/integration/splitter/AbstractMessageSplitter.java index 93ffb7c28b..699ad8dfad 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/splitter/AbstractMessageSplitter.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/splitter/AbstractMessageSplitter.java @@ -18,6 +18,7 @@ package org.springframework.integration.splitter; import java.util.ArrayList; import java.util.Collection; +import java.util.Collections; import java.util.List; import java.util.UUID; @@ -30,9 +31,12 @@ import org.springframework.integration.handler.AbstractReplyProducingMessageHand * Base class for Message-splitting handlers. * * @author Mark Fisher + * @author Dave Syer */ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMessageHandler { + public static final String SEQUENCE_DETAILS = MessageHeaders.PREFIX + "sequenceDetails"; + @Override @SuppressWarnings("unchecked") protected final Object handleRequestMessage(Message message) { @@ -40,15 +44,29 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess if (result == null) { return null; } - Object correlationId = (message.getHeaders().getCorrelationId() != null) ? - message.getHeaders().getCorrelationId() : message.getHeaders().getId(); + MessageHeaders headers = message.getHeaders(); + Object incomingCorrelationId = headers.getCorrelationId(); + List incomingSequenceDetails = headers.get(SEQUENCE_DETAILS, List.class); + if (incomingCorrelationId != null) { + if (incomingSequenceDetails == null) { + incomingSequenceDetails = new ArrayList(); + } + else { + incomingSequenceDetails = new ArrayList(incomingSequenceDetails); + } + incomingSequenceDetails.add(new Object[] { incomingCorrelationId, headers.getSequenceNumber(), + headers.getSequenceSize() }); + incomingSequenceDetails = Collections.unmodifiableList(incomingSequenceDetails); + } + Object correlationId = headers.getId(); List> messageBuilders = new ArrayList>(); if (result instanceof Collection) { Collection items = (Collection) result; int sequenceNumber = 0; int sequenceSize = items.size(); for (Object item : items) { - messageBuilders.add(this.createBuilder(item, correlationId, ++sequenceNumber, sequenceSize)); + messageBuilders.add(this.createBuilder(item, incomingSequenceDetails, correlationId, ++sequenceNumber, + sequenceSize)); } } else if (result.getClass().isArray()) { @@ -56,23 +74,26 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess int sequenceNumber = 0; int sequenceSize = items.length; for (Object item : items) { - messageBuilders.add(this.createBuilder(item, correlationId, ++sequenceNumber, sequenceSize)); + messageBuilders.add(this.createBuilder(item, incomingSequenceDetails, correlationId, ++sequenceNumber, + sequenceSize)); } } else { - messageBuilders.add(this.createBuilder(result, correlationId, 1, 1)); + messageBuilders.add(this.createBuilder(result, incomingSequenceDetails, correlationId, 1, 1)); } return messageBuilders; } @SuppressWarnings("unchecked") - private MessageBuilder createBuilder(Object item, Object correlationId, int sequenceNumber, int sequenceSize) { - MessageBuilder builder = (item instanceof Message) ? - MessageBuilder.fromMessage((Message) item) : MessageBuilder.withPayload(item); - builder.setCorrelationId(correlationId) - .setSequenceNumber(sequenceNumber) - .setSequenceSize(sequenceSize) + private MessageBuilder createBuilder(Object item, List incomingSequenceDetails, Object correlationId, + int sequenceNumber, int sequenceSize) { + MessageBuilder builder = (item instanceof Message) ? MessageBuilder.fromMessage((Message) item) + : MessageBuilder.withPayload(item); + builder.setCorrelationId(correlationId).setSequenceNumber(sequenceNumber).setSequenceSize(sequenceSize) .setHeader(MessageHeaders.ID, UUID.randomUUID()); + if (incomingSequenceDetails != null) { + builder.setHeader(SEQUENCE_DETAILS, incomingSequenceDetails); + } return builder; } @@ -95,12 +116,10 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess } /** - * Subclasses must override this method to split the received Message. The - * return value may be a Collection or Array. The individual elements may - * be Messages, but it is not necessary. If the elements are not Messages, - * each will be provided as the payload of a Message. It is also acceptable - * to return a single Object or Message. In that case, a single reply - * Message will be produced. + * Subclasses must override this method to split the received Message. The return value may be a Collection or + * Array. The individual elements may be Messages, but it is not necessary. If the elements are not Messages, each + * will be provided as the payload of a Message. It is also acceptable to return a single Object or Message. In that + * case, a single reply Message will be produced. */ protected abstract Object splitMessage(Message message); diff --git a/spring-integration-core/src/test/java/org/springframework/integration/aggregator/scenarios/NestedAggregationTests-context.xml b/spring-integration-core/src/test/java/org/springframework/integration/aggregator/scenarios/NestedAggregationTests-context.xml new file mode 100644 index 0000000000..3c5e7f1a56 --- /dev/null +++ b/spring-integration-core/src/test/java/org/springframework/integration/aggregator/scenarios/NestedAggregationTests-context.xml @@ -0,0 +1,27 @@ + + + + + + + + + + + + + + \ No newline at end of file diff --git a/spring-integration-core/src/test/java/org/springframework/integration/aggregator/scenarios/NestedAggregationTests.java b/spring-integration-core/src/test/java/org/springframework/integration/aggregator/scenarios/NestedAggregationTests.java new file mode 100644 index 0000000000..ae761cf969 --- /dev/null +++ b/spring-integration-core/src/test/java/org/springframework/integration/aggregator/scenarios/NestedAggregationTests.java @@ -0,0 +1,67 @@ +/* + * Copyright 2002-2010 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.aggregator.scenarios; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotNull; + +import java.util.Arrays; +import java.util.List; + +import org.junit.Test; +import org.junit.runner.RunWith; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.integration.Message; +import org.springframework.integration.channel.DirectChannel; +import org.springframework.integration.core.GenericMessage; +import org.springframework.integration.core.MessagingTemplate; +import org.springframework.test.context.ContextConfiguration; +import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; + +/** + * + * @author Dave Syer + */ +@ContextConfiguration +@RunWith(SpringJUnit4ClassRunner.class) +public class NestedAggregationTests { + + @Autowired + DirectChannel input; + + @Test + public void testAggregatorWithNestedSplitter() throws Exception { + List result = sendAndReceiveMessage(input, 2000); + assertNotNull("Expected result and got null", result); + assertEquals("[[foo, bar, spam], [bar, foo]]", result.toString()); + } + + private List sendAndReceiveMessage(DirectChannel channel, int timeout) { + + MessagingTemplate messagingTemplate = new MessagingTemplate(); + messagingTemplate.setReceiveTimeout(timeout); + + @SuppressWarnings("unchecked") + Message> message = (Message>) messagingTemplate.sendAndReceive(channel, + new GenericMessage>>(Arrays.asList(Arrays.asList("foo", "bar", "spam"), Arrays.asList("bar", + "foo")))); + + return message == null ? null : message.getPayload(); + + } + +} diff --git a/spring-integration-core/src/test/java/org/springframework/integration/endpoint/CorrelationIdTests.java b/spring-integration-core/src/test/java/org/springframework/integration/endpoint/CorrelationIdTests.java index ac6f5235de..35d3e959a5 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/endpoint/CorrelationIdTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/endpoint/CorrelationIdTests.java @@ -20,13 +20,13 @@ import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertTrue; import org.junit.Test; - import org.springframework.integration.Message; import org.springframework.integration.channel.DirectChannel; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.core.MessageBuilder; import org.springframework.integration.core.StringMessage; import org.springframework.integration.handler.ServiceActivatingHandler; +import org.springframework.integration.splitter.AbstractMessageSplitter; import org.springframework.integration.splitter.MethodInvokingSplitter; /** @@ -123,8 +123,10 @@ public class CorrelationIdTests { splitter.handleMessage(message); Message reply1 = testChannel.receive(100); Message reply2 = testChannel.receive(100); - assertEquals(correlationIdForTest, reply1.getHeaders().getCorrelationId()); - assertEquals(correlationIdForTest, reply2.getHeaders().getCorrelationId()); + assertEquals(message.getHeaders().getId(), reply1.getHeaders().getCorrelationId()); + assertEquals(message.getHeaders().getId(), reply2.getHeaders().getCorrelationId()); + assertTrue("Sequence details missing", reply1.getHeaders().containsKey(AbstractMessageSplitter.SEQUENCE_DETAILS)); + assertTrue("Sequence details missing", reply2.getHeaders().containsKey(AbstractMessageSplitter.SEQUENCE_DETAILS)); } @SuppressWarnings("unused")