diff --git a/client/src/main/java/org/apache/rocketmq/client/producer/DefaultMQProducer.java b/client/src/main/java/org/apache/rocketmq/client/producer/DefaultMQProducer.java
index 2091bbabbff..2c73d0e3bc9 100644
--- a/client/src/main/java/org/apache/rocketmq/client/producer/DefaultMQProducer.java
+++ b/client/src/main/java/org/apache/rocketmq/client/producer/DefaultMQProducer.java
@@ -54,6 +54,7 @@
import java.util.Set;
import java.util.concurrent.CopyOnWriteArraySet;
import java.util.concurrent.ExecutorService;
+import java.util.concurrent.atomic.AtomicBoolean;
/**
* This class is the entry point for applications intending to send messages.
@@ -162,6 +163,7 @@ public class DefaultMQProducer extends ClientConfig implements MQProducer {
* Instance for batching message automatically
*/
private ProduceAccumulator produceAccumulator = null;
+ private final AtomicBoolean produceAccumulatorStarted = new AtomicBoolean(false);
/**
* Indicate whether to block message when asynchronous sending traffic is too heavy.
@@ -374,7 +376,7 @@ public DefaultMQProducer(final String namespace, final String producerGroup, RPC
public void start() throws MQClientException {
this.setProducerGroup(withNamespace(this.producerGroup));
this.defaultMQProducerImpl.start();
- if (this.produceAccumulator != null) {
+ if (this.produceAccumulator != null && this.produceAccumulatorStarted.compareAndSet(false, true)) {
this.produceAccumulator.start();
}
if (enableTrace) {
@@ -411,7 +413,7 @@ public void start() throws MQClientException {
@Override
public void shutdown() {
this.defaultMQProducerImpl.shutdown();
- if (this.produceAccumulator != null) {
+ if (this.produceAccumulator != null && this.produceAccumulatorStarted.compareAndSet(true, false)) {
this.produceAccumulator.shutdown();
}
if (null != traceDispatcher) {
diff --git a/client/src/main/java/org/apache/rocketmq/client/producer/ProduceAccumulator.java b/client/src/main/java/org/apache/rocketmq/client/producer/ProduceAccumulator.java
index 809830e4641..7b9234485e3 100644
--- a/client/src/main/java/org/apache/rocketmq/client/producer/ProduceAccumulator.java
+++ b/client/src/main/java/org/apache/rocketmq/client/producer/ProduceAccumulator.java
@@ -56,6 +56,8 @@ public class ProduceAccumulator {
private final Map asyncSendBatchs = new ConcurrentHashMap();
private final AtomicLong currentlyHoldSize = new AtomicLong(0);
private final String instanceName;
+ // Number of started producers sharing this accumulator through the same client ID.
+ private int producerCount;
public ProduceAccumulator(String instanceName) {
this.instanceName = instanceName;
@@ -155,12 +157,17 @@ private void doWork() throws Exception {
}
}
- void start() {
- guardThreadForSyncSend.start();
- guardThreadForAsyncSend.start();
+ synchronized void start() {
+ if (producerCount++ == 0) {
+ guardThreadForSyncSend.start();
+ guardThreadForAsyncSend.start();
+ }
}
- void shutdown() {
+ synchronized void shutdown() {
+ if (producerCount == 0 || --producerCount > 0) {
+ return;
+ }
guardThreadForSyncSend.shutdown();
guardThreadForAsyncSend.shutdown();
}
diff --git a/client/src/test/java/org/apache/rocketmq/client/producer/ProduceAccumulatorTest.java b/client/src/test/java/org/apache/rocketmq/client/producer/ProduceAccumulatorTest.java
index 8e76238d47f..49aed041845 100644
--- a/client/src/test/java/org/apache/rocketmq/client/producer/ProduceAccumulatorTest.java
+++ b/client/src/test/java/org/apache/rocketmq/client/producer/ProduceAccumulatorTest.java
@@ -17,6 +17,7 @@
package org.apache.rocketmq.client.producer;
+import java.lang.reflect.Field;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
@@ -24,6 +25,7 @@
import java.util.concurrent.TimeUnit;
import org.apache.rocketmq.client.exception.MQBrokerException;
import org.apache.rocketmq.client.exception.MQClientException;
+import org.apache.rocketmq.client.impl.producer.DefaultMQProducerImpl;
import org.apache.rocketmq.common.message.Message;
import org.apache.rocketmq.common.message.MessageBatch;
import org.apache.rocketmq.common.message.MessageQueue;
@@ -31,8 +33,17 @@
import org.junit.Test;
import static org.assertj.core.api.Assertions.assertThat;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.times;
+import static org.mockito.Mockito.verify;
public class ProduceAccumulatorTest {
+ private void setProducerField(DefaultMQProducer producer, String fieldName, Object value) throws Exception {
+ Field field = DefaultMQProducer.class.getDeclaredField(fieldName);
+ field.setAccessible(true);
+ field.set(producer, value);
+ }
+
private boolean compareMessageBatch(MessageBatch a, MessageBatch b) {
if (!a.getTopic().equals(b.getTopic())) {
return false;
@@ -101,6 +112,60 @@ public void onException(Throwable e) {
assertThat(compareMessageBatch(messageBatch1, messageBatch2)).isTrue();
}
+ @Test
+ public void testSharedAccumulatorRemainsRunningUntilLastProducerShutdown() throws Exception {
+ MockMQProducer mockMQProducer = new MockMQProducer();
+ ProduceAccumulator produceAccumulator = new ProduceAccumulator("shared-client");
+ produceAccumulator.batchMaxDelayMs(50);
+ DefaultMQProducer producerA = new DefaultMQProducer("producer-a");
+ DefaultMQProducer producerB = new DefaultMQProducer("producer-b");
+ producerA.setInstanceName("shared-client");
+ producerB.setInstanceName("shared-client");
+ setProducerField(producerA, "defaultMQProducerImpl", mock(DefaultMQProducerImpl.class));
+ setProducerField(producerB, "defaultMQProducerImpl", mock(DefaultMQProducerImpl.class));
+ setProducerField(producerA, "produceAccumulator", produceAccumulator);
+ setProducerField(producerB, "produceAccumulator", produceAccumulator);
+
+ // Two producers with the same client ID share this accumulator.
+ producerA.start();
+ producerB.start();
+ producerA.shutdown();
+
+ CountDownLatch countDownLatch = new CountDownLatch(1);
+ try {
+ produceAccumulator.send(new Message("testTopic", "1".getBytes()), new SendCallback() {
+ @Override
+ public void onSuccess(SendResult sendResult) {
+ countDownLatch.countDown();
+ }
+
+ @Override
+ public void onException(Throwable e) {
+ countDownLatch.countDown();
+ }
+ }, mockMQProducer);
+
+ assertThat(countDownLatch.await(3, TimeUnit.SECONDS)).isTrue();
+ } finally {
+ producerB.shutdown();
+ }
+ }
+
+ @Test
+ public void testProducerReleasesSharedAccumulatorOnlyOnce() throws Exception {
+ DefaultMQProducer producer = new DefaultMQProducer("testProducerGroup");
+ ProduceAccumulator produceAccumulator = mock(ProduceAccumulator.class);
+ setProducerField(producer, "defaultMQProducerImpl", mock(DefaultMQProducerImpl.class));
+ setProducerField(producer, "produceAccumulator", produceAccumulator);
+
+ producer.start();
+ producer.shutdown();
+ producer.shutdown();
+
+ verify(produceAccumulator, times(1)).start();
+ verify(produceAccumulator, times(1)).shutdown();
+ }
+
@Test
public void testProduceAccumulator_sync() throws MQBrokerException, RemotingException, InterruptedException, MQClientException {
final MockMQProducer mockMQProducer = new MockMQProducer();