From 46ae25d88be8e11035e46ac268dc6f8fc4d92293 Mon Sep 17 00:00:00 2001 From: Artem Bilan Date: Wed, 19 Jun 2019 16:14:12 -0400 Subject: [PATCH] GH-2967: Fix ScatterGatherH for headers copy (#2968) * GH-2967: Fix ScatterGatherH for headers copy Fixes https://github.com/spring-projects/spring-integration/issues/2967 The `ChannelInterceptor` is added into the `this.gatherChannel` on each request message making a subsequent requests for scatter-gather as halting on reply. * Add an interceptor into an injected `this.gatherChannel` only once during `ScatterGatherHandler` initialization * Introduce `ORIGINAL_REPLY_CHANNEL` and `ORIGINAL_ERROR_CHANNEL` headers to carry a request reply and error channels from headers * Populate `REPLY_CHANNEL` and `ERROR_CHANNEL` headers back before sending scattering replies into gatherer * Transfer a `GATHER_RESULT_CHANNEL` header now directly from the scatter message to make it available in the reply from the gatherer * Add note about those headers in the `scatter-gather.adoc` * Modify `ScatterGatherTests` to be sure that `ScatterGatherHandler` works for several requests **Cherry-pick to 5.1.x** * * Fix language in doc --- .../scattergather/ScatterGatherHandler.java | 74 +++++++++---------- .../config/ScatterGatherTests-context.xml | 10 ++- .../config/ScatterGatherTests.java | 53 ++++++++----- src/reference/asciidoc/scatter-gather.adoc | 2 + 4 files changed, 80 insertions(+), 59 deletions(-) diff --git a/spring-integration-core/src/main/java/org/springframework/integration/scattergather/ScatterGatherHandler.java b/spring-integration-core/src/main/java/org/springframework/integration/scattergather/ScatterGatherHandler.java index 1b528b6c8e..b31b0524ab 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/scattergather/ScatterGatherHandler.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/scattergather/ScatterGatherHandler.java @@ -55,6 +55,10 @@ public class ScatterGatherHandler extends AbstractReplyProducingMessageHandler i private static final String GATHER_RESULT_CHANNEL = "gatherResultChannel"; + private static final String ORIGINAL_REPLY_CHANNEL = "originalReplyChannel"; + + private static final String ORIGINAL_ERROR_CHANNEL = "originalErrorChannel"; + private final MessageChannel scatterChannel; private final MessageHandler gatherer; @@ -107,9 +111,24 @@ public class ScatterGatherHandler extends AbstractReplyProducingMessageHandler i protected void doInit() { BeanFactory beanFactory = getBeanFactory(); if (this.gatherChannel == null) { - this.gatherChannel = new FixedSubscriberChannel(this.gatherer); + this.gatherChannel = + new FixedSubscriberChannel((message) -> + this.gatherer.handleMessage(enhanceScatterReplyMessage(message))); } else { + Assert.isInstanceOf(ChannelInterceptorAware.class, this.gatherChannel, + () -> "An injected 'gatherChannel' '" + this.gatherChannel + + "' must be an 'InterceptableChannel' instance."); + ((ChannelInterceptorAware) this.gatherChannel) + .addInterceptor(0, + new ChannelInterceptor() { + + @Override + public Message preSend(Message message, MessageChannel channel) { + return enhanceScatterReplyMessage(message); + } + + }); if (this.gatherChannel instanceof SubscribableChannel) { this.gatherEndpoint = new EventDrivenConsumer((SubscribableChannel) this.gatherChannel, this.gatherer); } @@ -121,7 +140,7 @@ public class ScatterGatherHandler extends AbstractReplyProducingMessageHandler i this.gatherEndpoint = new ReactiveStreamsConsumer(this.gatherChannel, this.gatherer); } else { - throw new BeanInitializationException("Unsupported 'replyChannel' type '" + + throw new BeanInitializationException("Unsupported 'gatherChannel' type '" + this.gatherChannel.getClass() + "'. " + "'SubscribableChannel', 'PollableChannel' or 'ReactiveStreamsSubscribableChannel' " + "types are supported."); @@ -131,7 +150,7 @@ public class ScatterGatherHandler extends AbstractReplyProducingMessageHandler i } ((MessageProducer) this.gatherer) - .setOutputChannel(new FixedSubscriberChannel(message -> { + .setOutputChannel(new FixedSubscriberChannel((message) -> { MessageHeaders headers = message.getHeaders(); MessageChannel gatherResultChannel = headers.get(GATHER_RESULT_CHANNEL, MessageChannel.class); if (gatherResultChannel != null) { @@ -144,35 +163,28 @@ public class ScatterGatherHandler extends AbstractReplyProducingMessageHandler i })); } + private Message enhanceScatterReplyMessage(Message message) { + MessageHeaders headers = message.getHeaders(); + return getMessageBuilderFactory() + .fromMessage(message) + .setHeader(MessageHeaders.REPLY_CHANNEL, headers.get(ORIGINAL_REPLY_CHANNEL)) + .setHeader(MessageHeaders.ERROR_CHANNEL, headers.get(ORIGINAL_ERROR_CHANNEL)) + .removeHeaders(ORIGINAL_REPLY_CHANNEL, ORIGINAL_ERROR_CHANNEL) + .build(); + } + @Override protected Object handleRequestMessage(Message requestMessage) { + MessageHeaders requestMessageHeaders = requestMessage.getHeaders(); PollableChannel gatherResultChannel = new QueueChannel(); - MessageChannel replyChannel = this.gatherChannel; - - if (replyChannel instanceof ChannelInterceptorAware) { - ((ChannelInterceptorAware) replyChannel) - .addInterceptor(0, - new ChannelInterceptor() { - - @Override - public Message preSend(Message message, MessageChannel channel) { - return enhanceScatterReplyMessage(message, gatherResultChannel, requestMessage); - } - - }); - } - else { - replyChannel = - new FixedSubscriberChannel(message -> - this.messagingTemplate.send(this.gatherChannel, - enhanceScatterReplyMessage(message, gatherResultChannel, requestMessage))); - } - Message scatterMessage = getMessageBuilderFactory() .fromMessage(requestMessage) - .setReplyChannel(replyChannel) + .setHeader(GATHER_RESULT_CHANNEL, gatherResultChannel) + .setHeader(ORIGINAL_REPLY_CHANNEL, requestMessageHeaders.getReplyChannel()) + .setHeader(ORIGINAL_ERROR_CHANNEL, requestMessageHeaders.getErrorChannel()) + .setReplyChannel(this.gatherChannel) .setErrorChannelName(this.errorChannelName) .build(); @@ -181,18 +193,6 @@ public class ScatterGatherHandler extends AbstractReplyProducingMessageHandler i return gatherResultChannel.receive(this.gatherTimeout); } - private Message enhanceScatterReplyMessage(Message message, PollableChannel gatherResultChannel, - Message requestMessage) { - - MessageHeaders requestMessageHeaders = requestMessage.getHeaders(); - return getMessageBuilderFactory() - .fromMessage(message) - .setHeader(GATHER_RESULT_CHANNEL, gatherResultChannel) - .setHeader(MessageHeaders.REPLY_CHANNEL, requestMessageHeaders.getReplyChannel()) - .setHeader(MessageHeaders.ERROR_CHANNEL, requestMessageHeaders.getErrorChannel()) - .build(); - } - @Override public void start() { if (this.gatherEndpoint != null) { diff --git a/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests-context.xml b/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests-context.xml index 41ab828291..5255b2eaf9 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests-context.xml +++ b/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests-context.xml @@ -4,7 +4,8 @@ xmlns="http://www.springframework.org/schema/integration" xmlns:task="http://www.springframework.org/schema/task" xsi:schemaLocation="http://www.springframework.org/schema/beans https://www.springframework.org/schema/beans/spring-beans.xsd - http://www.springframework.org/schema/integration https://www.springframework.org/schema/integration/spring-integration.xsd http://www.springframework.org/schema/task https://www.springframework.org/schema/task/spring-task.xsd"> + http://www.springframework.org/schema/integration https://www.springframework.org/schema/integration/spring-integration.xsd + http://www.springframework.org/schema/task https://www.springframework.org/schema/task/spring-task.xsd"> @@ -52,9 +53,12 @@ - + - + + + diff --git a/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests.java b/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests.java index 0bc5c1af75..f549530381 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/scattergather/config/ScatterGatherTests.java @@ -16,10 +16,8 @@ package org.springframework.integration.scattergather.config; -import static org.hamcrest.Matchers.greaterThanOrEqualTo; -import static org.hamcrest.Matchers.instanceOf; +import static org.assertj.core.api.Assertions.assertThat; import static org.junit.Assert.assertNotNull; -import static org.junit.Assert.assertThat; import java.util.List; @@ -32,16 +30,19 @@ import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; import org.springframework.messaging.PollableChannel; import org.springframework.messaging.support.GenericMessage; +import org.springframework.test.annotation.DirtiesContext; import org.springframework.test.context.ContextConfiguration; import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; /** * @author Artem Bilan * @author Gary Russell + * * @since 4.1 */ @ContextConfiguration @RunWith(SpringJUnit4ClassRunner.class) +@DirtiesContext public class ScatterGatherTests { @Autowired @@ -61,36 +62,50 @@ public class ScatterGatherTests { @Test public void testAuction() { - this.inputAuction.send(new GenericMessage("foo")); + this.inputAuction.send(new GenericMessage<>("foo")); Message bestQuoteMessage = this.output.receive(10000); - assertNotNull(bestQuoteMessage); - Object payload = bestQuoteMessage.getPayload(); - assertThat(payload, instanceOf(List.class)); - assertThat(((List) payload).size(), greaterThanOrEqualTo(1)); + assertThat(bestQuoteMessage) + .isNotNull() + .extracting(Message::getPayload) + .isInstanceOf(List.class) + .asList() + .hasSizeGreaterThanOrEqualTo(1); } @Test public void testDistribution() { - this.inputDistribution.send(new GenericMessage("foo")); + this.inputDistribution.send(new GenericMessage<>("foo")); Message bestQuoteMessage = this.output.receive(10000); - assertNotNull(bestQuoteMessage); - Object payload = bestQuoteMessage.getPayload(); - assertThat(payload, instanceOf(List.class)); - assertThat(((List) payload).size(), greaterThanOrEqualTo(1)); + assertThat(bestQuoteMessage) + .isNotNull() + .extracting(Message::getPayload) + .isInstanceOf(List.class) + .asList() + .hasSizeGreaterThanOrEqualTo(1); } @Test public void testGatewayScatterGather() { - Message bestQuoteMessage = this.gateway.exchange(new GenericMessage("foo")); - assertNotNull(bestQuoteMessage); - Object payload = bestQuoteMessage.getPayload(); - assertThat(payload, instanceOf(List.class)); - assertThat(((List) payload).size(), greaterThanOrEqualTo(1)); + Message bestQuoteMessage = this.gateway.exchange(new GenericMessage<>("foo")); + assertThat(bestQuoteMessage) + .isNotNull() + .extracting(Message::getPayload) + .isInstanceOf(List.class) + .asList() + .hasSizeGreaterThanOrEqualTo(1); + + bestQuoteMessage = this.gateway.exchange(new GenericMessage<>("bar")); + assertThat(bestQuoteMessage) + .isNotNull() + .extracting(Message::getPayload) + .isInstanceOf(List.class) + .asList() + .hasSizeGreaterThanOrEqualTo(1); } @Test public void testWithinChain() { - this.scatterGatherWithinChain.send(new GenericMessage("foo")); + this.scatterGatherWithinChain.send(new GenericMessage<>("foo")); for (int i = 0; i < 3; i++) { Message result = this.output.receive(10000); assertNotNull(result); diff --git a/src/reference/asciidoc/scatter-gather.adoc b/src/reference/asciidoc/scatter-gather.adoc index 58ae2cceba..1bc3ef648e 100644 --- a/src/reference/asciidoc/scatter-gather.adoc +++ b/src/reference/asciidoc/scatter-gather.adoc @@ -209,5 +209,7 @@ Such an exception `payload` can be filtered out in the `MessageGroupProcessor` o NOTE: Before sending scattering results to the gatherer, `ScatterGatherHandler` reinstates the request message headers, including reply and error channels if any. This way errors from the `AggregatingMessageHandler` are going to be propagated to the caller, even if an async hand off is applied in scatter recipient subflows. +For successful operation, a `gatherResultChannel`, `originalReplyChannel` and `originalErrorChannel` headers must be transferred back to replies from scatter recipient subflows. In this case a reasonable, finite `gatherTimeout` must be configured for the `ScatterGatherHandler`. Otherwise it is going to be blocked waiting for a reply from the gatherer forever, by default. +