From 0a9359c90e18136117cb9b67a19c1fd47a171992 Mon Sep 17 00:00:00 2001 From: anshlykov Date: Tue, 14 Apr 2020 18:20:39 +0300 Subject: [PATCH] RabbitProducerProperties#batchingStrategyBeanName fix checkstyle Resolves #287 --- .../properties/RabbitProducerProperties.java | 14 +++ .../rabbit/RabbitMessageChannelBinder.java | 27 ++++-- .../binder/rabbit/RabbitBinderTests.java | 95 +++++++++++++++++++ 3 files changed, 129 insertions(+), 7 deletions(-) diff --git a/spring-cloud-stream-binder-rabbit-core/src/main/java/org/springframework/cloud/stream/binder/rabbit/properties/RabbitProducerProperties.java b/spring-cloud-stream-binder-rabbit-core/src/main/java/org/springframework/cloud/stream/binder/rabbit/properties/RabbitProducerProperties.java index 71428fe1f..4023ba612 100644 --- a/spring-cloud-stream-binder-rabbit-core/src/main/java/org/springframework/cloud/stream/binder/rabbit/properties/RabbitProducerProperties.java +++ b/spring-cloud-stream-binder-rabbit-core/src/main/java/org/springframework/cloud/stream/binder/rabbit/properties/RabbitProducerProperties.java @@ -52,6 +52,12 @@ public class RabbitProducerProperties extends RabbitCommonProperties { */ private int batchTimeout = 5000; + /** + * the bean name of a custom batching strategy to use instead of the + * {@link org.springframework.amqp.rabbit.batch.SimpleBatchingStrategy}. + */ + private String batchingStrategyBeanName; + /** * true to use transacted channels. */ @@ -195,4 +201,12 @@ public class RabbitProducerProperties extends RabbitCommonProperties { this.confirmAckChannel = confirmAckChannel; } + public String getBatchingStrategyBeanName() { + return batchingStrategyBeanName; + } + + public void setBatchingStrategyBeanName(String batchingStrategyBeanName) { + this.batchingStrategyBeanName = batchingStrategyBeanName; + } + } 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 7593c9eee..902e32fe0 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 @@ -823,13 +823,10 @@ public class RabbitMessageChannelBinder extends boolean mandatory) { RabbitTemplate rabbitTemplate; if (properties.isBatchingEnabled()) { - BatchingStrategy batchingStrategy = new SimpleBatchingStrategy( - properties.getBatchSize(), properties.getBatchBufferLimit(), - properties.getBatchTimeout()); - rabbitTemplate = new BatchingRabbitTemplate(batchingStrategy, - getApplicationContext().getBean( - IntegrationContextUtils.TASK_SCHEDULER_BEAN_NAME, - TaskScheduler.class)); + BatchingStrategy batchingStrategy = getBatchingStrategy(properties); + TaskScheduler taskScheduler = getApplicationContext() + .getBean(IntegrationContextUtils.TASK_SCHEDULER_BEAN_NAME, TaskScheduler.class); + rabbitTemplate = new BatchingRabbitTemplate(batchingStrategy, taskScheduler); } else { rabbitTemplate = new RabbitTemplate(); @@ -859,6 +856,22 @@ public class RabbitMessageChannelBinder extends return rabbitTemplate; } + private BatchingStrategy getBatchingStrategy(RabbitProducerProperties properties) { + BatchingStrategy batchingStrategy; + if (properties.getBatchingStrategyBeanName() != null) { + batchingStrategy = getApplicationContext() + .getBean(properties.getBatchingStrategyBeanName(), BatchingStrategy.class); + } + else { + batchingStrategy = new SimpleBatchingStrategy( + properties.getBatchSize(), + properties.getBatchBufferLimit(), + properties.getBatchTimeout() + ); + } + return batchingStrategy; + } + private String getStackTraceAsString(Throwable cause) { StringWriter stringWriter = new StringWriter(); PrintWriter printWriter = new PrintWriter(stringWriter, true); 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 978d358cf..26cdab37d 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 @@ -19,8 +19,13 @@ package org.springframework.cloud.stream.binder.rabbit; import java.io.PrintWriter; import java.io.StringWriter; import java.lang.reflect.Constructor; +import java.nio.ByteBuffer; import java.nio.charset.StandardCharsets; +import java.util.ArrayList; import java.util.Arrays; +import java.util.Collection; +import java.util.Collections; +import java.util.Date; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -51,8 +56,11 @@ import org.springframework.amqp.core.BindingBuilder; import org.springframework.amqp.core.DirectExchange; import org.springframework.amqp.core.ExchangeTypes; import org.springframework.amqp.core.MessageDeliveryMode; +import org.springframework.amqp.core.MessageProperties; import org.springframework.amqp.core.Queue; import org.springframework.amqp.core.TopicExchange; +import org.springframework.amqp.rabbit.batch.BatchingStrategy; +import org.springframework.amqp.rabbit.batch.MessageBatch; import org.springframework.amqp.rabbit.connection.CachingConnectionFactory; import org.springframework.amqp.rabbit.connection.CachingConnectionFactory.ConfirmType; import org.springframework.amqp.rabbit.connection.ConnectionFactory; @@ -2049,6 +2057,43 @@ public class RabbitBinderTests extends binding.unbind(); } + @Test + public void testCustomBatchingStrategy() throws Exception { + RabbitTestBinder binder = getBinder(); + ExtendedProducerProperties producerProperties = createProducerProperties(); + producerProperties.getExtension().setDeliveryMode(MessageDeliveryMode.NON_PERSISTENT); + producerProperties.getExtension().setBatchingEnabled(true); + producerProperties.getExtension().setBatchingStrategyBeanName("testCustomBatchingStrategy"); + producerProperties.setRequiredGroups("default"); + + ConfigurableListableBeanFactory beanFactory = binder.getApplicationContext().getBeanFactory(); + beanFactory.registerSingleton("testCustomBatchingStrategy", new TestBatchingStrategy()); + + DirectChannel output = createBindableChannel("output", createProducerBindingProperties(producerProperties)); + output.setBeanName("batchingProducer"); + Binding producerBinding = binder.bindProducer("batching.0", output, producerProperties); + + Log logger = spy(TestUtils.getPropertyValue(binder, "binder.compressingPostProcessor.logger", Log.class)); + new DirectFieldAccessor(TestUtils.getPropertyValue(binder, "binder.compressingPostProcessor")) + .setPropertyValue("logger", logger); + when(logger.isTraceEnabled()).thenReturn(true); + + assertThat(TestUtils.getPropertyValue(binder, "binder.compressingPostProcessor.level")) + .isEqualTo(Deflater.BEST_SPEED); + + output.send(new GenericMessage<>("0".getBytes())); + output.send(new GenericMessage<>("1".getBytes())); + output.send(new GenericMessage<>("2".getBytes())); + output.send(new GenericMessage<>("3".getBytes())); + output.send(new GenericMessage<>("4".getBytes())); + + Object out = spyOn("batching.0.default").receive(false); + assertThat(out).isInstanceOf(byte[].class); + assertThat(new String((byte[]) out)).isEqualTo("0\u0000\n1\u0000\n2\u0000\n3\u0000\n4\u0000\n"); + + producerBinding.unbind(); + } + private SimpleMessageListenerContainer verifyContainer(Lifecycle endpoint) { SimpleMessageListenerContainer container; RetryTemplate retry; @@ -2204,4 +2249,54 @@ public class RabbitBinderTests extends } // @checkstyle:on + public static class TestBatchingStrategy implements BatchingStrategy { + + private final List messages = new ArrayList<>(); + private String exchange; + private String routingKey; + private int currentSize; + + @Override + public MessageBatch addToBatch(String exchange, String routingKey, org.springframework.amqp.core.Message message) { + this.exchange = exchange; + this.routingKey = routingKey; + this.messages.add(message); + currentSize += message.getBody().length + 2; + + MessageBatch batch = null; + if (this.messages.size() == 5) { + batch = this.doReleaseBatch(); + } + + return batch; + } + + @Override + public Date nextRelease() { + return null; + } + + @Override + public Collection releaseBatches() { + MessageBatch batch = this.doReleaseBatch(); + return batch == null ? Collections.emptyList() : Collections.singletonList(batch); + } + + private MessageBatch doReleaseBatch() { + if (this.messages.size() < 1) { + return null; + } + else { + ByteBuffer byteBuffer = ByteBuffer.wrap(new byte[this.currentSize]); + for (org.springframework.amqp.core.Message message: messages) { + byteBuffer.put(message.getBody()).putChar('\n'); + } + MessageBatch messageBatch = new MessageBatch(this.exchange, this.routingKey, + new org.springframework.amqp.core.Message(byteBuffer.array(), new MessageProperties())); + this.messages.clear(); + return messageBatch; + } + } + } + }