diff --git a/spring-integration-core/src/main/java/org/springframework/integration/filter/MessageFilter.java b/spring-integration-core/src/main/java/org/springframework/integration/filter/MessageFilter.java index 27450f1afc..f0fe87fc38 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/filter/MessageFilter.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/filter/MessageFilter.java @@ -18,7 +18,6 @@ package org.springframework.integration.filter; import org.springframework.beans.factory.BeanFactoryAware; import org.springframework.integration.Message; -import org.springframework.integration.MessageHeaders; import org.springframework.integration.MessageRejectedException; import org.springframework.integration.core.MessageChannel; import org.springframework.integration.core.MessageSelector; @@ -114,9 +113,8 @@ public class MessageFilter extends AbstractReplyProducingMessageHandler { } @Override - protected void handleResult(Object replyMessage, MessageHeaders requestHeaders) { - Assert.isInstanceOf(Message.class, replyMessage); - this.sendReplyMessage((Message) replyMessage, requestHeaders.getReplyChannel()); + protected boolean shouldCopyRequestHeaders() { + return false; } } diff --git a/spring-integration-core/src/main/java/org/springframework/integration/handler/AbstractReplyProducingMessageHandler.java b/spring-integration-core/src/main/java/org/springframework/integration/handler/AbstractReplyProducingMessageHandler.java index 438b2674cf..f2bfde1c35 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/handler/AbstractReplyProducingMessageHandler.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/handler/AbstractReplyProducingMessageHandler.java @@ -113,18 +113,39 @@ public abstract class AbstractReplyProducingMessageHandler extends AbstractMessa } } - protected void handleResult(Object result, MessageHeaders requestHeaders) { - Message replyMessage = this.createReplyMessage(result, requestHeaders); + private void handleResult(Object result, MessageHeaders requestHeaders) { + if (result instanceof Iterable && this.shouldSplitIterableReply()) { + for (Object o : (Iterable) result) { + this.produceReply(o, requestHeaders); + } + } + else if (result != null) { + this.produceReply(result, requestHeaders); + } + } + + private void produceReply(Object reply, MessageHeaders requestHeaders) { + Message replyMessage = this.createReplyMessage(reply, requestHeaders); this.sendReplyMessage(replyMessage, requestHeaders.getReplyChannel()); } private Message createReplyMessage(Object reply, MessageHeaders requestHeaders) { + MessageBuilder builder = null; if (reply instanceof Message) { - return MessageBuilder.fromMessage((Message) reply).copyHeadersIfAbsent(requestHeaders).build(); + if (!this.shouldCopyRequestHeaders()) { + return (Message) reply; + } + builder = MessageBuilder.fromMessage((Message) reply); + } + else if (reply instanceof MessageBuilder) { + builder = (MessageBuilder) reply; + } + else { + builder = MessageBuilder.withPayload(reply); + } + if (this.shouldCopyRequestHeaders()) { + builder.copyHeadersIfAbsent(requestHeaders); } - MessageBuilder builder = (reply instanceof MessageBuilder) - ? (MessageBuilder) reply : MessageBuilder.withPayload(reply); - builder.copyHeadersIfAbsent(requestHeaders); return builder.build(); } @@ -135,7 +156,7 @@ public abstract class AbstractReplyProducingMessageHandler extends AbstractMessa * @param replyMessage the reply Message to send * @param replyChannelHeaderValue the 'replyChannel' header value from the original request */ - protected final void sendReplyMessage(Message replyMessage, final Object replyChannelHeaderValue) { + private final void sendReplyMessage(Message replyMessage, final Object replyChannelHeaderValue) { if (logger.isDebugEnabled()) { logger.debug("handler '" + this + "' sending reply Message: " + replyMessage); } @@ -167,6 +188,20 @@ public abstract class AbstractReplyProducingMessageHandler extends AbstractMessa } } + /** + * Subclasses may override this. False by default. + */ + protected boolean shouldSplitIterableReply() { + return false; + } + + /** + * Subclasses may override this. True by default. + */ + protected boolean shouldCopyRequestHeaders() { + return true; + } + /** * Subclasses must implement this method to handle the request Message. The return * value may be a Message, a MessageBuilder, or any plain Object. The base class 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 699ad8dfad..1afa43e056 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 @@ -54,8 +54,8 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess else { incomingSequenceDetails = new ArrayList(incomingSequenceDetails); } - incomingSequenceDetails.add(new Object[] { incomingCorrelationId, headers.getSequenceNumber(), - headers.getSequenceSize() }); + incomingSequenceDetails.add(new Object[] { + incomingCorrelationId, headers.getSequenceNumber(), headers.getSequenceSize() }); incomingSequenceDetails = Collections.unmodifiableList(incomingSequenceDetails); } Object correlationId = headers.getId(); @@ -65,8 +65,8 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess int sequenceNumber = 0; int sequenceSize = items.size(); for (Object item : items) { - messageBuilders.add(this.createBuilder(item, incomingSequenceDetails, correlationId, ++sequenceNumber, - sequenceSize)); + messageBuilders.add(this.createBuilder( + item, incomingSequenceDetails, correlationId, ++sequenceNumber, sequenceSize)); } } else if (result.getClass().isArray()) { @@ -74,8 +74,8 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess int sequenceNumber = 0; int sequenceSize = items.length; for (Object item : items) { - messageBuilders.add(this.createBuilder(item, incomingSequenceDetails, correlationId, ++sequenceNumber, - sequenceSize)); + messageBuilders.add(this.createBuilder( + item, incomingSequenceDetails, correlationId, ++sequenceNumber, sequenceSize)); } } else { @@ -84,7 +84,7 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess return messageBuilders; } - @SuppressWarnings("unchecked") + @SuppressWarnings({"unchecked", "rawtypes"}) private MessageBuilder createBuilder(Object item, List incomingSequenceDetails, Object correlationId, int sequenceNumber, int sequenceSize) { MessageBuilder builder = (item instanceof Message) ? MessageBuilder.fromMessage((Message) item) @@ -98,16 +98,8 @@ public abstract class AbstractMessageSplitter extends AbstractReplyProducingMess } @Override - @SuppressWarnings("unchecked") - protected void handleResult(Object result, MessageHeaders requestHeaders) { - if (result instanceof Iterable) { - for (Object o : (Iterable) result) { - super.handleResult(o, requestHeaders); - } - } - else { - super.handleResult(result, requestHeaders); - } + protected boolean shouldSplitIterableReply() { + return true; } @Override diff --git a/spring-integration-core/src/main/java/org/springframework/integration/transformer/MessageTransformingHandler.java b/spring-integration-core/src/main/java/org/springframework/integration/transformer/MessageTransformingHandler.java index cf3bf9be70..16a5b21003 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/transformer/MessageTransformingHandler.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/transformer/MessageTransformingHandler.java @@ -18,7 +18,6 @@ package org.springframework.integration.transformer; import org.springframework.beans.factory.BeanFactoryAware; import org.springframework.integration.Message; -import org.springframework.integration.MessageHeaders; import org.springframework.integration.core.MessageHandler; import org.springframework.integration.handler.AbstractReplyProducingMessageHandler; import org.springframework.util.Assert; @@ -74,8 +73,8 @@ public class MessageTransformingHandler extends AbstractReplyProducingMessageHan } @Override - protected void handleResult(Object replyMessage, MessageHeaders requestHeaders) { - this.sendReplyMessage((Message) replyMessage, requestHeaders.getReplyChannel()); + protected boolean shouldCopyRequestHeaders() { + return false; } }