diff --git a/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitExpressionEvaluatingInterceptor.java b/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitExpressionEvaluatingInterceptor.java new file mode 100644 index 000000000..ef9376ba8 --- /dev/null +++ b/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitExpressionEvaluatingInterceptor.java @@ -0,0 +1,90 @@ +/* + * Copyright 2017 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.cloud.stream.binder.rabbit; + +import org.springframework.expression.EvaluationContext; +import org.springframework.expression.Expression; +import org.springframework.expression.ExpressionParser; +import org.springframework.expression.spel.standard.SpelExpressionParser; +import org.springframework.integration.support.MessageBuilder; +import org.springframework.messaging.Message; +import org.springframework.messaging.MessageChannel; +import org.springframework.messaging.support.ChannelInterceptorAdapter; +import org.springframework.util.Assert; + +/** + * Interceptor to evaluate expressions for outbound messages before serialization. + * + * @author Gary Russell + * @since 2.0 + * + */ +public class RabbitExpressionEvaluatingInterceptor extends ChannelInterceptorAdapter { + + public static final ExpressionParser PARSER = new SpelExpressionParser(); + + public static final String ROUTING_KEY_HEADER = "scst_routingKey"; + + public static final String DELAY_HEADER = "scst_delay"; + + private final Expression routingKeyExpression; + + private final Expression delayExpression; + + private final EvaluationContext evaluationContext; + + /** + * Construct an instance with the provided expressions and evaluation context. + * At least one expression muse be non-null. + * @param routingKeyExpression the routing key expresssion. + * @param delayExpression the delay expression. + * @param evaluationContext the evaluation context. + */ + public RabbitExpressionEvaluatingInterceptor(String routingKeyExpression, String delayExpression, + EvaluationContext evaluationContext) { + Assert.isTrue(routingKeyExpression != null || delayExpression != null, + "At least one expression is required"); + Assert.notNull(evaluationContext, "the 'evaluationContext' cannot be null"); + if (routingKeyExpression != null) { + this.routingKeyExpression = PARSER.parseExpression(routingKeyExpression); + } + else { + this.routingKeyExpression = null; + } + if (delayExpression != null) { + this.delayExpression = PARSER.parseExpression(delayExpression); + } + else { + this.delayExpression = null; + } + this.evaluationContext = evaluationContext; + } + + @Override + public Message preSend(Message message, MessageChannel channel) { + MessageBuilder builder = MessageBuilder.fromMessage(message); + if (this.routingKeyExpression != null) { + builder.setHeader(ROUTING_KEY_HEADER, + this.routingKeyExpression.getValue(this.evaluationContext, message)); + } + if (this.delayExpression != null) { + builder.setHeader(DELAY_HEADER, this.delayExpression.getValue(this.evaluationContext, message)); + } + return builder.build(); + } + +} diff --git a/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitMessageChannelBinder.java b/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitMessageChannelBinder.java index e499d2164..01fb21b38 100644 --- a/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitMessageChannelBinder.java +++ b/spring-cloud-stream-binder-rabbit/src/main/java/org/springframework/cloud/stream/binder/rabbit/RabbitMessageChannelBinder.java @@ -64,6 +64,7 @@ import org.springframework.integration.amqp.inbound.AmqpInboundChannelAdapter; import org.springframework.integration.amqp.outbound.AmqpOutboundEndpoint; import org.springframework.integration.amqp.support.AmqpMessageHeaderErrorMessageStrategy; import org.springframework.integration.amqp.support.DefaultAmqpHeaderMapper; +import org.springframework.integration.channel.AbstractMessageChannel; import org.springframework.integration.context.IntegrationContextUtils; import org.springframework.integration.core.MessageProducer; import org.springframework.integration.support.DefaultErrorMessageStrategy; @@ -303,13 +304,20 @@ public class RabbitMessageChannelBinder buildRabbitTemplate(producerProperties.getExtension(), errorChannel != null)); endpoint.setExchangeName(producerDestination.getName()); RabbitProducerProperties extendedProperties = producerProperties.getExtension(); + boolean expresssionInterceptorNeeded = expresssionInterceptorNeeded(extendedProperties); String routingKeyExpression = extendedProperties.getRoutingKeyExpression(); if (!producerProperties.isPartitioned()) { if (routingKeyExpression == null) { endpoint.setRoutingKey(destination); } else { - endpoint.setRoutingKeyExpressionString(routingKeyExpression); + if (expresssionInterceptorNeeded) { + endpoint.setRoutingKeyExpressionString("headers['" + + RabbitExpressionEvaluatingInterceptor.ROUTING_KEY_HEADER + "']"); + } + else { + endpoint.setRoutingKeyExpressionString(routingKeyExpression); + } } } else { @@ -317,11 +325,24 @@ public class RabbitMessageChannelBinder endpoint.setRoutingKeyExpressionString(buildPartitionRoutingExpression(destination, false)); } else { - endpoint.setRoutingKeyExpressionString(buildPartitionRoutingExpression(routingKeyExpression, true)); + if (expresssionInterceptorNeeded) { + endpoint.setRoutingKeyExpressionString(buildPartitionRoutingExpression("headers['" + + RabbitExpressionEvaluatingInterceptor.ROUTING_KEY_HEADER + "']", true)); + } + else { + endpoint.setRoutingKeyExpressionString(buildPartitionRoutingExpression(routingKeyExpression, + true)); + } } } if (extendedProperties.getDelayExpression() != null) { - endpoint.setDelayExpressionString(extendedProperties.getDelayExpression()); + if (expresssionInterceptorNeeded) { + endpoint.setDelayExpressionString("headers['" + + RabbitExpressionEvaluatingInterceptor.DELAY_HEADER + "']"); + } + else { + endpoint.setDelayExpressionString(extendedProperties.getDelayExpression()); + } } DefaultAmqpHeaderMapper mapper = DefaultAmqpHeaderMapper.outboundMapper(); List headerPatterns = new ArrayList<>(extendedProperties.getHeaderPatterns().length + 1); @@ -342,6 +363,25 @@ public class RabbitMessageChannelBinder return endpoint; } + + @Override + protected void postProcessOutputChannel(MessageChannel outputChannel, + ExtendedProducerProperties producerProperties) { + RabbitProducerProperties extendedProperties = producerProperties.getExtension(); + if (expresssionInterceptorNeeded(extendedProperties)) { + ((AbstractMessageChannel) outputChannel).addInterceptor(0, + new RabbitExpressionEvaluatingInterceptor(extendedProperties.getRoutingKeyExpression(), + extendedProperties.getDelayExpression(), getEvaluationContext())); + } + } + + public boolean expresssionInterceptorNeeded(RabbitProducerProperties extendedProperties) { + return extendedProperties.getRoutingKeyExpression() != null + && extendedProperties.getRoutingKeyExpression().contains("payload") + || (extendedProperties.getDelayExpression() != null + && extendedProperties.getDelayExpression().contains("payload")); + } + private void checkConnectionFactoryIsErrorCapable() { if (!(this.connectionFactory instanceof CachingConnectionFactory)) { logger.warn("Unknown connection factory type, cannot determine error capabilities: " diff --git a/spring-cloud-stream-binder-rabbit/src/test/java/org/springframework/cloud/stream/binder/rabbit/RabbitBinderTests.java b/spring-cloud-stream-binder-rabbit/src/test/java/org/springframework/cloud/stream/binder/rabbit/RabbitBinderTests.java index e12a15ee1..9c021c275 100644 --- a/spring-cloud-stream-binder-rabbit/src/test/java/org/springframework/cloud/stream/binder/rabbit/RabbitBinderTests.java +++ b/spring-cloud-stream-binder-rabbit/src/test/java/org/springframework/cloud/stream/binder/rabbit/RabbitBinderTests.java @@ -16,14 +16,8 @@ package org.springframework.cloud.stream.binder.rabbit; -import static org.assertj.core.api.Assertions.assertThat; -import static org.assertj.core.api.Assertions.fail; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.spy; -import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.when; - import java.lang.reflect.Constructor; +import java.nio.charset.StandardCharsets; import java.util.Arrays; import java.util.HashMap; import java.util.List; @@ -34,22 +28,23 @@ import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; import java.util.zip.Deflater; -import com.rabbitmq.http.client.domain.QueueInfo; - import org.apache.commons.logging.Log; import org.junit.Rule; import org.junit.Test; import org.junit.rules.TestName; import org.mockito.ArgumentCaptor; + import org.springframework.amqp.AmqpIOException; import org.springframework.amqp.core.AcknowledgeMode; import org.springframework.amqp.core.AmqpTemplate; +import org.springframework.amqp.core.AnonymousQueue; import org.springframework.amqp.core.BindingBuilder; import org.springframework.amqp.core.DirectExchange; import org.springframework.amqp.core.Exchange; import org.springframework.amqp.core.ExchangeTypes; import org.springframework.amqp.core.MessageDeliveryMode; import org.springframework.amqp.core.Queue; +import org.springframework.amqp.core.TopicExchange; import org.springframework.amqp.rabbit.connection.CachingConnectionFactory; import org.springframework.amqp.rabbit.connection.ConnectionFactory; import org.springframework.amqp.rabbit.core.RabbitAdmin; @@ -87,6 +82,7 @@ import org.springframework.integration.amqp.support.ReturnedAmqpMessageException import org.springframework.integration.channel.DirectChannel; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.context.IntegrationContextUtils; +import org.springframework.integration.expression.ValueExpression; import org.springframework.integration.support.MessageBuilder; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; @@ -94,12 +90,22 @@ import org.springframework.messaging.MessageHandler; import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.MessagingException; import org.springframework.messaging.SubscribableChannel; +import org.springframework.messaging.support.ChannelInterceptorAdapter; import org.springframework.messaging.support.ErrorMessage; import org.springframework.messaging.support.GenericMessage; import org.springframework.retry.support.RetryTemplate; import org.springframework.util.MimeTypeUtils; import org.springframework.util.ReflectionUtils; +import com.rabbitmq.http.client.domain.QueueInfo; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.fail; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.spy; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + /** * @author Mark Fisher * @author Gary Russell @@ -1037,7 +1043,7 @@ public class RabbitBinderTests extends producerProperties.getExtension().setCompress(true); producerProperties.setRequiredGroups("default"); - DirectChannel output = createBindableChannel("input", createProducerBindingProperties(producerProperties)); + DirectChannel output = createBindableChannel("output", createProducerBindingProperties(producerProperties)); output.setBeanName("batchingProducer"); Binding producerBinding = binder.bindProducer("batching.0", output, producerProperties); @@ -1237,6 +1243,87 @@ public class RabbitBinderTests extends } } + @Test + public void testRoutingKeyExpression() throws Exception { + RabbitTestBinder binder = getBinder(); + ExtendedProducerProperties producerProperties = createProducerProperties(); + producerProperties.getExtension().setRoutingKeyExpression("payload.field"); + + DirectChannel output = createBindableChannel("output", createProducerBindingProperties(producerProperties)); + output.setBeanName("rkeProducer"); + Binding producerBinding = binder.bindProducer("rke", output, producerProperties); + + RabbitAdmin admin = new RabbitAdmin(this.rabbitAvailableRule.getResource()); + Queue queue = new AnonymousQueue(); + TopicExchange exchange = new TopicExchange("rke"); + org.springframework.amqp.core.Binding binding = BindingBuilder.bind(queue).to(exchange).with("rkeTest"); + admin.declareQueue(queue); + admin.declareBinding(binding); + + output.addInterceptor(new ChannelInterceptorAdapter() { + + @Override + public Message preSend(Message message, MessageChannel channel) { + assertThat(message.getHeaders().get(RabbitExpressionEvaluatingInterceptor.ROUTING_KEY_HEADER)) + .isEqualTo("rkeTest"); + return message; + } + + }); + + output.send(new GenericMessage<>(new Pojo("rkeTest"))); + + Object out = spyOn(queue.getName()).receive(false); + assertThat(out).isInstanceOf(byte[].class); + assertThat(new String((byte[]) out, StandardCharsets.UTF_8)).isEqualTo("{\"field\":\"rkeTest\"}"); + + producerBinding.unbind(); + } + + @Test + public void testRoutingKeyExpressionPartitionedAndDelay() throws Exception { + RabbitTestBinder binder = getBinder(); + ExtendedProducerProperties producerProperties = createProducerProperties(); + producerProperties.getExtension().setRoutingKeyExpression("payload.field"); + // requires delayed message exchange plugin; tested locally +// producerProperties.getExtension().setDelayedExchange(true); + producerProperties.getExtension().setDelayExpression("1000"); + producerProperties.setPartitionKeyExpression(new ValueExpression<>(0)); + + DirectChannel output = createBindableChannel("output", createProducerBindingProperties(producerProperties)); + output.setBeanName("rkeProducer"); + Binding producerBinding = binder.bindProducer("rkep", output, producerProperties); + + RabbitAdmin admin = new RabbitAdmin(this.rabbitAvailableRule.getResource()); + Queue queue = new AnonymousQueue(); + TopicExchange exchange = new TopicExchange("rkep"); + org.springframework.amqp.core.Binding binding = + BindingBuilder.bind(queue).to(exchange).with("rkepTest-0"); + admin.declareQueue(queue); + admin.declareBinding(binding); + + output.addInterceptor(new ChannelInterceptorAdapter() { + + @Override + public Message preSend(Message message, MessageChannel channel) { + assertThat(message.getHeaders().get(RabbitExpressionEvaluatingInterceptor.ROUTING_KEY_HEADER)) + .isEqualTo("rkepTest"); + assertThat(message.getHeaders().get(RabbitExpressionEvaluatingInterceptor.DELAY_HEADER)) + .isEqualTo(1000); + return message; + } + + }); + + output.send(new GenericMessage<>(new Pojo("rkepTest"))); + + Object out = spyOn(queue.getName()).receive(false); + assertThat(out).isInstanceOf(byte[].class); + assertThat(new String((byte[]) out, StandardCharsets.UTF_8)).isEqualTo("{\"field\":\"rkepTest\"}"); + + producerBinding.unbind(); + } + private SimpleMessageListenerContainer verifyContainer(Lifecycle endpoint) { SimpleMessageListenerContainer container; RetryTemplate retry; @@ -1336,4 +1423,29 @@ public class RabbitBinderTests extends } + public static class Pojo { + + private String field; + + public Pojo() { + super(); + } + + + public Pojo(String field) { + this.field = field; + } + + + public String getField() { + return this.field; + } + + + public void setField(String field) { + this.field = field; + } + + } + }