GH-2731: Fix nested gateway error propagation

Fixes https://github.com/spring-projects/spring-integration/issues/2731

When we have a nested gateway configuration, we are losing the context
of the current request message in case of errors and the downstream
`MessagingException` is just re-thrown as is, without the proper
`failedMessage` and its processed `errorChannel` header.

* Check for exception type and for the `errorChannel` header in the
current `requestMessage` before re-throwing as a new
`MessageHandlingException` in the `MessagingGatewaySupport.sendAndReceive()`
* Extract `.gateway()` tests into the separate `GatewayDslTests` class
This commit is contained in:
Artem Bilan
2019-02-08 09:55:20 -05:00
committed by Gary Russell
parent 37563a140c
commit 4c88ddd429
3 changed files with 303 additions and 216 deletions

View File

@@ -22,6 +22,7 @@ import java.util.concurrent.atomic.AtomicLong;
import org.reactivestreams.Subscriber;
import org.springframework.beans.factory.BeanFactory;
import org.springframework.core.AttributeAccessor;
import org.springframework.integration.MessageTimeoutException;
import org.springframework.integration.channel.ReactiveStreamsSubscribableChannel;
@@ -49,6 +50,7 @@ import org.springframework.lang.Nullable;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.MessageDeliveryException;
import org.springframework.messaging.MessageHandlingException;
import org.springframework.messaging.MessageHeaders;
import org.springframework.messaging.MessagingException;
import org.springframework.messaging.PollableChannel;
@@ -75,9 +77,9 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
private static final long DEFAULT_TIMEOUT = 1000L;
private final SimpleMessageConverter messageConverter = new SimpleMessageConverter();
protected final MessagingTemplate messagingTemplate; // NOSONAR
protected final MessagingTemplate messagingTemplate;
private final SimpleMessageConverter messageConverter = new SimpleMessageConverter();
private final HistoryWritingMessagePostProcessor historyWritingPostProcessor =
new HistoryWritingMessagePostProcessor();
@@ -92,34 +94,33 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
private ErrorMessageStrategy errorMessageStrategy = new DefaultErrorMessageStrategy();
private volatile MessageChannel requestChannel;
private MessageChannel requestChannel;
private volatile String requestChannelName;
private String requestChannelName;
private volatile MessageChannel replyChannel;
private MessageChannel replyChannel;
private volatile String replyChannelName;
private String replyChannelName;
private volatile MessageChannel errorChannel;
private MessageChannel errorChannel;
private volatile String errorChannelName;
private String errorChannelName;
private volatile long replyTimeout = DEFAULT_TIMEOUT;
private long replyTimeout = DEFAULT_TIMEOUT;
@SuppressWarnings("rawtypes")
private volatile InboundMessageMapper requestMapper = new DefaultRequestMapper();
private volatile boolean initialized;
private InboundMessageMapper<Object> requestMapper = new DefaultRequestMapper();
private volatile AbstractEndpoint replyMessageCorrelator;
private volatile String managedType;
private String managedType;
private volatile String managedName;
private String managedName;
private volatile boolean countsEnabled;
private boolean countsEnabled;
private volatile boolean loggingEnabled = true;
private boolean loggingEnabled = true;
private volatile boolean initialized;
/**
@@ -232,8 +233,11 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
* from any object passed in a send or sendAndReceive operation.
* @param requestMapper The request mapper.
*/
@SuppressWarnings("unchecked")
public void setRequestMapper(@Nullable InboundMessageMapper<?> requestMapper) {
this.requestMapper = (requestMapper != null) ? requestMapper : new DefaultRequestMapper();
if (requestMapper != null) {
this.requestMapper = (InboundMessageMapper<Object>) requestMapper;
}
this.messageConverter.setInboundMessageMapper(this.requestMapper);
}
@@ -337,20 +341,22 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
Assert.state(!(this.errorChannelName != null && this.errorChannel != null),
"'errorChannelName' and 'errorChannel' are mutually exclusive.");
this.historyWritingPostProcessor.setTrackableComponent(this);
this.historyWritingPostProcessor.setMessageBuilderFactory(this.getMessageBuilderFactory());
if (this.getBeanFactory() != null) {
this.messagingTemplate.setBeanFactory(this.getBeanFactory());
MessageBuilderFactory messageBuilderFactory = getMessageBuilderFactory();
this.historyWritingPostProcessor.setMessageBuilderFactory(messageBuilderFactory);
BeanFactory beanFactory = getBeanFactory();
if (beanFactory != null) {
this.messagingTemplate.setBeanFactory(beanFactory);
if (this.requestMapper instanceof DefaultRequestMapper) {
((DefaultRequestMapper) this.requestMapper).setMessageBuilderFactory(this.getMessageBuilderFactory());
((DefaultRequestMapper) this.requestMapper).setMessageBuilderFactory(messageBuilderFactory);
}
this.messageConverter.setBeanFactory(this.getBeanFactory());
this.messageConverter.setBeanFactory(beanFactory);
}
this.initialized = true;
}
private void initializeIfNecessary() {
if (!this.initialized) {
this.afterPropertiesSet();
afterPropertiesSet();
}
}
@@ -361,13 +367,8 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
*/
@Nullable
public MessageChannel getRequestChannel() {
if (this.requestChannelName != null) {
synchronized (this) {
if (this.requestChannelName != null) {
this.requestChannel = getChannelResolver().resolveDestination(this.requestChannelName);
this.requestChannelName = null;
}
}
if (this.requestChannel == null && this.requestChannelName != null) {
this.requestChannel = getChannelResolver().resolveDestination(this.requestChannelName);
}
return this.requestChannel;
}
@@ -377,14 +378,10 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
* @return the reply channel instance
* @since 5.1
*/
@Nullable
public MessageChannel getReplyChannel() {
if (this.replyChannelName != null) {
synchronized (this) {
if (this.replyChannelName != null) {
this.replyChannel = getChannelResolver().resolveDestination(this.replyChannelName);
this.replyChannelName = null;
}
}
if (this.replyChannel == null && this.replyChannelName != null) {
this.replyChannel = getChannelResolver().resolveDestination(this.replyChannelName);
}
return this.replyChannel;
}
@@ -397,13 +394,8 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
*/
@Nullable
public MessageChannel getErrorChannel() {
if (this.errorChannelName != null) {
synchronized (this) {
if (this.errorChannelName != null) {
this.errorChannel = getChannelResolver().resolveDestination(this.errorChannelName);
this.errorChannelName = null;
}
}
if (this.errorChannel == null && this.errorChannelName != null) {
this.errorChannel = getChannelResolver().resolveDestination(this.errorChannelName);
}
return this.errorChannel;
}
@@ -426,7 +418,7 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
this.messagingTemplate.send(errorChan, new ErrorMessage(e));
}
else {
this.rethrow(e, "failed to send message");
rethrow(e, "failed to send message");
}
}
}
@@ -435,8 +427,7 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
protected Object receive() {
this.initializeIfNecessary();
MessageChannel channel = getReplyChannel();
Assert.state(channel != null && (channel instanceof PollableChannel),
"receive is not supported, because no pollable reply channel has been configured");
assertPollableChannel(channel);
return this.messagingTemplate.receiveAndConvert(channel, Object.class);
}
@@ -444,8 +435,7 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
protected Message<?> receiveMessage() {
initializeIfNecessary();
MessageChannel channel = getReplyChannel();
Assert.state(channel instanceof PollableChannel,
"receive is not supported, because no pollable reply channel has been configured");
assertPollableChannel(channel);
return this.messagingTemplate.receive(channel);
}
@@ -453,8 +443,7 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
protected Object receive(long timeout) {
this.initializeIfNecessary();
MessageChannel channel = getReplyChannel();
Assert.state(channel != null && (channel instanceof PollableChannel),
"receive is not supported, because no pollable reply channel has been configured");
assertPollableChannel(channel);
return this.messagingTemplate.receiveAndConvert(channel, timeout);
}
@@ -462,25 +451,28 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
protected Message<?> receiveMessage(long timeout) {
initializeIfNecessary();
MessageChannel channel = getReplyChannel();
assertPollableChannel(channel);
return this.messagingTemplate.receive(channel, timeout);
}
private void assertPollableChannel(@Nullable MessageChannel channel) {
Assert.state(channel instanceof PollableChannel,
"receive is not supported, because no pollable reply channel has been configured");
return this.messagingTemplate.receive(channel, timeout);
}
@Nullable
protected Object sendAndReceive(Object object) {
return this.doSendAndReceive(object, true);
return doSendAndReceive(object, true);
}
@Nullable
protected Message<?> sendAndReceiveMessage(Object object) {
return (Message<?>) this.doSendAndReceive(object, false);
return (Message<?>) doSendAndReceive(object, false);
}
@SuppressWarnings("unchecked")
@Nullable
private Object doSendAndReceive(Object object, boolean shouldConvert) {
this.initializeIfNecessary();
initializeIfNecessary();
Assert.notNull(object, "request must not be null");
MessageChannel channel = getRequestChannel();
if (channel == null) {
@@ -522,53 +514,66 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
}
}
}
catch (Exception e) {
catch (Exception ex) {
if (logger.isDebugEnabled()) {
logger.debug("failure occurred in gateway sendAndReceive: " + e.getMessage());
logger.debug("failure occurred in gateway sendAndReceive: " + ex.getMessage());
}
error = e;
error = ex;
}
if (error != null) {
MessageChannel errorChan = getErrorChannel();
if (errorChan != null) {
ErrorMessage errorMessage = buildErrorMessage(requestMessage, error);
Message<?> errorFlowReply = null;
try {
errorFlowReply = this.messagingTemplate.sendAndReceive(errorChan, errorMessage);
}
catch (Exception errorFlowFailure) {
throw new MessagingException(errorMessage, "failure occurred in error-handling flow",
errorFlowFailure);
}
if (shouldConvert) {
Object result = (errorFlowReply != null) ? errorFlowReply.getPayload() : null;
if (result instanceof Throwable) {
this.rethrow((Throwable) result, "error flow returned Exception");
}
return result;
}
if (errorFlowReply != null && errorFlowReply.getPayload() instanceof Throwable) {
this.rethrow((Throwable) errorFlowReply.getPayload(), "error flow returned an Error Message");
}
if (errorFlowReply == null && this.errorOnTimeout) {
if (object instanceof Message) {
throw new MessageTimeoutException((Message<?>) object,
"No reply received from error channel within timeout");
}
else {
throw new MessageTimeoutException("No reply received from error channel within timeout");
}
}
return errorFlowReply;
}
else { // no errorChannel so we'll propagate
this.rethrow(error, "gateway received checked Exception");
}
return handleSendAndReceiveError(object, requestMessage, error, shouldConvert);
}
return reply;
}
@Nullable
private Object handleSendAndReceiveError(Object object, @Nullable Message<?> requestMessage, Throwable error,
boolean shouldConvert) {
MessageChannel errorChan = getErrorChannel();
if (errorChan != null) {
ErrorMessage errorMessage = buildErrorMessage(requestMessage, error);
Message<?> errorFlowReply = null;
try {
errorFlowReply = this.messagingTemplate.sendAndReceive(errorChan, errorMessage);
}
catch (Exception errorFlowFailure) {
throw new MessagingException(errorMessage, "failure occurred in error-handling flow",
errorFlowFailure);
}
if (shouldConvert) {
Object result = (errorFlowReply != null) ? errorFlowReply.getPayload() : null;
if (result instanceof Throwable) {
rethrow((Throwable) result, "error flow returned Exception");
}
return result;
}
if (errorFlowReply != null && errorFlowReply.getPayload() instanceof Throwable) {
rethrow((Throwable) errorFlowReply.getPayload(), "error flow returned an Error Message");
}
if (errorFlowReply == null && this.errorOnTimeout) {
if (object instanceof Message) {
throw new MessageTimeoutException((Message<?>) object,
"No reply received from error channel within timeout");
}
else {
throw new MessageTimeoutException("No reply received from error channel within timeout");
}
}
return errorFlowReply;
}
else {
if (error instanceof MessagingException &&
requestMessage != null && requestMessage.getHeaders().getErrorChannel() != null) {
// We are in nested flow where upstream expects errors in its own errorChannel header.
error = new MessageHandlingException(requestMessage, error);
}
rethrow(error, "gateway received checked Exception");
return null; // unreachable
}
}
protected Mono<Message<?>> sendAndReceiveMessageReactive(Object object) {
initializeIfNecessary();
Assert.notNull(object, "request must not be null");
@@ -582,7 +587,6 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
return doSendAndReceiveMessageReactive(channel, object, false);
}
@SuppressWarnings("unchecked")
private Mono<Message<?>> doSendAndReceiveMessageReactive(MessageChannel requestChannel, Object object,
boolean error) {
@@ -612,52 +616,63 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
.setErrorChannel(replyChan)
.build();
if (requestChannel instanceof ReactiveStreamsSubscribableChannel) {
((ReactiveStreamsSubscribableChannel) requestChannel)
.subscribeTo(Mono.just(requestMessage));
}
else {
long sendTimeout = sendTimeout(requestMessage);
sendMessageForReactiveFlow(requestChannel, requestMessage);
boolean sent =
sendTimeout >= 0
? requestChannel.send(requestMessage, sendTimeout)
: requestChannel.send(requestMessage);
if (!sent) {
throw new MessageDeliveryException(requestMessage,
"Failed to send message to channel '" + requestChannel +
"' within timeout: " + sendTimeout);
}
}
return Mono.fromFuture(replyChan.messageFuture)
.doOnSubscribe(s -> {
if (!error && this.countsEnabled) {
this.messageCount.incrementAndGet();
}
})
.<Message<?>>map(replyMessage -> {
if (!error && replyMessage instanceof ErrorMessage) {
ErrorMessage em = (ErrorMessage) replyMessage;
if (em.getPayload() instanceof MessagingException) {
throw (MessagingException) em.getPayload();
}
else {
throw new MessagingException(requestMessage, em.getPayload());
}
}
else {
return MessageBuilder.fromMessage(replyMessage)
.setHeader(MessageHeaders.REPLY_CHANNEL, originalReplyChannelHeader)
.setHeader(MessageHeaders.ERROR_CHANNEL, originalErrorChannelHeader)
.build();
}
})
.onErrorResume(t -> error ? Mono.error(t) : handleSendError(requestMessage, t));
return buildReplyMono(requestMessage, replyChan, error, originalReplyChannelHeader,
originalErrorChannelHeader);
});
}
private void sendMessageForReactiveFlow(MessageChannel requestChannel, Message<?> requestMessage) {
if (requestChannel instanceof ReactiveStreamsSubscribableChannel) {
((ReactiveStreamsSubscribableChannel) requestChannel)
.subscribeTo(Mono.just(requestMessage));
}
else {
long sendTimeout = sendTimeout(requestMessage);
boolean sent =
sendTimeout >= 0
? requestChannel.send(requestMessage, sendTimeout)
: requestChannel.send(requestMessage);
if (!sent) {
throw new MessageDeliveryException(requestMessage,
"Failed to send message to channel '" + requestChannel +
"' within timeout: " + sendTimeout);
}
}
}
private Mono<Message<?>> buildReplyMono(Message<?> requestMessage, FutureReplyChannel replyChannel, boolean error,
@Nullable Object originalReplyChannelHeader, @Nullable Object originalErrorChannelHeader) {
return Mono.fromFuture(replyChannel.messageFuture)
.doOnSubscribe(s -> {
if (!error && this.countsEnabled) {
this.messageCount.incrementAndGet();
}
})
.<Message<?>>map(replyMessage -> {
if (!error && replyMessage instanceof ErrorMessage) {
ErrorMessage em = (ErrorMessage) replyMessage;
if (em.getPayload() instanceof MessagingException) {
throw (MessagingException) em.getPayload();
}
else {
throw new MessagingException(requestMessage, em.getPayload());
}
}
else {
return MessageBuilder.fromMessage(replyMessage)
.setHeader(MessageHeaders.REPLY_CHANNEL, originalReplyChannelHeader)
.setHeader(MessageHeaders.ERROR_CHANNEL, originalErrorChannelHeader)
.build();
}
})
.onErrorResume(t -> error ? Mono.error(t) : handleSendError(requestMessage, t));
}
private Mono<Message<?>> handleSendError(Message<?> requestMessage, Throwable exception) {
if (logger.isDebugEnabled()) {
logger.debug("failure occurred in gateway sendAndReceiveReactive: " + exception.getMessage());
@@ -669,7 +684,8 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
return doSendAndReceiveMessageReactive(channel, errorMessage, true);
}
catch (Exception errorFlowFailure) {
throw new MessagingException(errorMessage, "failure occurred in error-handling flow", errorFlowFailure);
throw new MessagingException(errorMessage, "failure occurred in error-handling flow",
errorFlowFailure);
}
}
else {
@@ -684,12 +700,6 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
return (sendTimeout != null ? sendTimeout : this.messagingTemplate.getSendTimeout());
}
private long receiveTimeout(Message<?> requestMessage) {
Long receiveTimeout = headerToLong(requestMessage.getHeaders()
.get(this.messagingTemplate.getReceiveTimeoutHeader()));
return (receiveTimeout != null ? receiveTimeout : this.messagingTemplate.getReceiveTimeout());
}
@Nullable
private Long headerToLong(@Nullable Object headerValue) {
if (headerValue instanceof Number) {
@@ -743,15 +753,15 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
protected void registerReplyMessageCorrelatorIfNecessary() {
MessageChannel replyChan = getReplyChannel();
if (replyChan != null && this.replyMessageCorrelator == null) {
boolean shouldStartCorrelator;
synchronized (this.replyMessageCorrelatorMonitor) {
if (this.replyMessageCorrelator != null) {
return;
}
AbstractEndpoint correlator;
BridgeHandler handler = new BridgeHandler();
if (getBeanFactory() != null) {
handler.setBeanFactory(getBeanFactory());
BeanFactory beanFactory = getBeanFactory();
if (beanFactory != null) {
handler.setBeanFactory(beanFactory);
}
handler.afterPropertiesSet();
if (replyChan instanceof SubscribableChannel) {
@@ -759,7 +769,9 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
}
else if (replyChan instanceof PollableChannel) {
PollingConsumer endpoint = new PollingConsumer((PollableChannel) replyChan, handler);
endpoint.setBeanFactory(getBeanFactory());
if (beanFactory != null) {
endpoint.setBeanFactory(beanFactory);
}
endpoint.setReceiveTimeout(this.replyTimeout);
endpoint.afterPropertiesSet();
correlator = endpoint;
@@ -775,12 +787,9 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
+ "SubscribableChannel or PollableChannel type are supported.");
}
this.replyMessageCorrelator = correlator;
shouldStartCorrelator = true;
}
if (shouldStartCorrelator && isRunning()) {
if (isRunning()) {
this.replyMessageCorrelator.start();
}
if (isRunning()) {
this.replyMessageCorrelator.start();
}
}
}
@@ -818,14 +827,11 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
}
@Override
public Message<?> toMessage(Object object, @Nullable Map<String, Object> headers) throws Exception {
public Message<?> toMessage(Object object, @Nullable Map<String, Object> headers) {
if (object instanceof Message<?>) {
return (Message<?>) object;
}
return object != null
? this.messageBuilderFactory.withPayload(object).copyHeadersIfAbsent(headers).build()
: null;
return this.messageBuilderFactory.withPayload(object).copyHeadersIfAbsent(headers).build();
}
}
@@ -834,6 +840,10 @@ public abstract class MessagingGatewaySupport extends AbstractEndpoint
private final CompletableFuture<Message<?>> messageFuture = new CompletableFuture<>();
FutureReplyChannel() {
super();
}
@Override
public boolean send(Message<?> message, long timeout) {
return this.messageFuture.complete(message);

View File

@@ -54,7 +54,6 @@ import org.springframework.context.annotation.ComponentScan;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.annotation.Scope;
import org.springframework.integration.MessageDispatchingException;
import org.springframework.integration.MessageRejectedException;
import org.springframework.integration.annotation.MessageEndpoint;
import org.springframework.integration.annotation.MessagingGateway;
import org.springframework.integration.annotation.ServiceActivator;
@@ -91,7 +90,6 @@ import org.springframework.messaging.MessageHeaders;
import org.springframework.messaging.MessagingException;
import org.springframework.messaging.PollableChannel;
import org.springframework.messaging.SubscribableChannel;
import org.springframework.messaging.support.ErrorMessage;
import org.springframework.messaging.support.GenericMessage;
import org.springframework.scheduling.TaskScheduler;
import org.springframework.scheduling.concurrent.ThreadPoolTaskExecutor;
@@ -180,14 +178,6 @@ public class IntegrationFlowTests {
@Qualifier("lambdasInput")
private MessageChannel lambdasInput;
@Autowired
@Qualifier("gatewayInput")
private MessageChannel gatewayInput;
@Autowired
@Qualifier("gatewayError")
private PollableChannel gatewayError;
@Test
public void testWithSupplierMessageSourceImpliedPoller() {
assertEquals("FOO", this.suppliedChannel.receive(10000).getPayload());
@@ -336,32 +326,6 @@ public class IntegrationFlowTests {
assertSame(message, this.messageStore.getMessage(message.getHeaders().getId()));
}
@Test
public void testGatewayFlow() {
PollableChannel replyChannel = new QueueChannel();
Message<String> message = MessageBuilder.withPayload("foo").setReplyChannel(replyChannel).build();
this.gatewayInput.send(message);
Message<?> receive = replyChannel.receive(10000);
assertNotNull(receive);
assertEquals("From Gateway SubFlow: FOO", receive.getPayload());
assertNull(this.gatewayError.receive(1));
message = MessageBuilder.withPayload("bar").setReplyChannel(replyChannel).build();
this.gatewayInput.send(message);
receive = replyChannel.receive(1);
assertNull(receive);
receive = this.gatewayError.receive(10000);
assertNotNull(receive);
assertThat(receive, instanceOf(ErrorMessage.class));
assertThat(receive.getPayload(), instanceOf(MessageRejectedException.class));
assertThat(((Exception) receive.getPayload()).getMessage(), containsString("' rejected Message"));
}
@Autowired
private SubscribableChannel tappedChannel1;
@@ -845,27 +809,6 @@ public class IntegrationFlowTests {
.get();
}
@Bean
public IntegrationFlow gatewayFlow() {
return IntegrationFlows.from("gatewayInput")
.gateway("gatewayRequest", g -> g.errorChannel("gatewayError").replyTimeout(10L))
.gateway(f -> f.transform("From Gateway SubFlow: "::concat))
.get();
}
@Bean
public IntegrationFlow gatewayRequestFlow() {
return IntegrationFlows.from("gatewayRequest")
.filter("foo"::equals, f -> f.throwExceptionOnRejection(true))
.<String, String>transform(String::toUpperCase)
.get();
}
@Bean
public MessageChannel gatewayError() {
return MessageChannels.queue().get();
}
@Bean
public IntegrationFlow errorRecovererFlow() {
return IntegrationFlows.from(Function.class, "errorRecovererFunction")

View File

@@ -0,0 +1,134 @@
/*
* Copyright 2019 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.dsl.gateway;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Qualifier;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.task.TaskExecutor;
import org.springframework.integration.MessageRejectedException;
import org.springframework.integration.channel.QueueChannel;
import org.springframework.integration.config.EnableIntegration;
import org.springframework.integration.dsl.IntegrationFlow;
import org.springframework.integration.dsl.IntegrationFlows;
import org.springframework.integration.dsl.MessageChannels;
import org.springframework.integration.support.MessageBuilder;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.PollableChannel;
import org.springframework.messaging.support.ErrorMessage;
import org.springframework.messaging.support.GenericMessage;
import org.springframework.test.context.junit.jupiter.SpringJUnitConfig;
/**
* @author Artem Bilan
*
* @since 5.1.3
*/
@SpringJUnitConfig
public class GatewayDslTests {
@Autowired
@Qualifier("gatewayInput")
private MessageChannel gatewayInput;
@Autowired
@Qualifier("gatewayError")
private PollableChannel gatewayError;
@Test
void testGatewayFlow() {
PollableChannel replyChannel = new QueueChannel();
Message<String> message = MessageBuilder.withPayload("foo").setReplyChannel(replyChannel).build();
this.gatewayInput.send(message);
Message<?> receive = replyChannel.receive(10000);
assertThat(receive).isNotNull();
assertThat(receive.getPayload()).isEqualTo("From Gateway SubFlow: FOO");
assertThat(this.gatewayError.receive(1)).isNull();
message = MessageBuilder.withPayload("bar").setReplyChannel(replyChannel).build();
this.gatewayInput.send(message);
assertThat(replyChannel.receive(1)).isNull();
receive = this.gatewayError.receive(10000);
assertThat(receive).isNotNull();
assertThat(receive).isInstanceOf(ErrorMessage.class);
assertThat(receive.getPayload()).isInstanceOf(MessageRejectedException.class);
assertThat(((Exception) receive.getPayload()).getMessage()).contains("' rejected Message");
}
@Autowired
@Qualifier("nestedGatewayErrorPropagationFlow.input")
private MessageChannel nestedGatewayErrorPropagationFlowInput;
@Test
void testNestedGatewayErrorPropagation() {
assertThatThrownBy(() -> this.nestedGatewayErrorPropagationFlowInput.send(new GenericMessage<>("test")))
.hasCauseInstanceOf(RuntimeException.class)
.hasMessageContaining("intentional");
}
@Configuration
@EnableIntegration
public static class ContextConfiguration {
@Bean
public IntegrationFlow gatewayFlow() {
return IntegrationFlows.from("gatewayInput")
.gateway("gatewayRequest", g -> g.errorChannel("gatewayError").replyTimeout(10L))
.gateway((f) -> f.transform("From Gateway SubFlow: "::concat))
.get();
}
@Bean
public IntegrationFlow gatewayRequestFlow() {
return IntegrationFlows.from("gatewayRequest")
.filter("foo"::equals, (f) -> f.throwExceptionOnRejection(true))
.<String, String>transform(String::toUpperCase)
.get();
}
@Bean
public MessageChannel gatewayError() {
return MessageChannels.queue().get();
}
@Bean
public IntegrationFlow nestedGatewayErrorPropagationFlow(TaskExecutor taskExecutor) {
return f -> f
.gateway((gatewayFlow) -> gatewayFlow
.channel((c) -> c.executor(taskExecutor))
.gateway((nestedGatewayFlow) -> nestedGatewayFlow
.transform((m) -> {
throw new RuntimeException("intentional");
})));
}
}
}