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
This commit is contained in:
Artem Bilan
2019-06-19 16:14:12 -04:00
committed by Gary Russell
parent 13ed8025a5
commit 46ae25d88b
4 changed files with 80 additions and 59 deletions

View File

@@ -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) {

View File

@@ -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">
<channel id="output">
<queue/>
@@ -52,9 +53,12 @@
<!--Sync scenario-->
<gateway id="gateway" default-request-channel="gatewayAuction" default-reply-timeout="10000" />
<gateway id="gateway" default-request-channel="gatewayAuction" default-reply-timeout="10000"/>
<scatter-gather input-channel="gatewayAuction" output-channel="bridgeChannel" scatter-channel="auctionChannel">
<channel id="gatherChannel2"/>
<scatter-gather input-channel="gatewayAuction" output-channel="bridgeChannel" scatter-channel="auctionChannel"
gather-channel="gatherChannel2">
<gatherer release-strategy-expression="messages.^[payload gt 5] != null or size() == 3"/>
</scatter-gather>

View File

@@ -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<String>("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<String>("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<String>("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<String>("foo"));
this.scatterGatherWithinChain.send(new GenericMessage<>("foo"));
for (int i = 0; i < 3; i++) {
Message<?> result = this.output.receive(10000);
assertNotNull(result);

View File

@@ -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.