diff --git a/spring-pulsar/src/main/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainer.java b/spring-pulsar/src/main/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainer.java index 2b5a25f7..1be0eb14 100644 --- a/spring-pulsar/src/main/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainer.java +++ b/spring-pulsar/src/main/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainer.java @@ -30,6 +30,9 @@ import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.locks.Condition; +import java.util.concurrent.locks.Lock; +import java.util.concurrent.locks.ReentrantLock; import java.util.stream.Collectors; import java.util.stream.Stream; import java.util.stream.StreamSupport; @@ -91,6 +94,10 @@ public class DefaultPulsarMessageListenerContainer extends AbstractPulsarMess private final AtomicBoolean receiveInProgress = new AtomicBoolean(); + private final Lock lockOnPause = new ReentrantLock(); + + private final Condition pausedCondition = this.lockOnPause.newCondition(); + public DefaultPulsarMessageListenerContainer(PulsarConsumerFactory pulsarConsumerFactory, PulsarContainerProperties pulsarContainerProperties) { this(pulsarConsumerFactory, pulsarContainerProperties, null); @@ -207,6 +214,14 @@ public class DefaultPulsarMessageListenerContainer extends AbstractPulsarMess consumer.resume(); } setPaused(false); + this.lockOnPause.lock(); + try { + // signal the lock's condition to continue. + this.pausedCondition.signal(); + } + finally { + this.lockOnPause.unlock(); + } } private final class Listener implements SchedulingAwareRunnable { @@ -352,6 +367,7 @@ public class DefaultPulsarMessageListenerContainer extends AbstractPulsarMess Messages messages = null; List> messageList = null; while (isRunning()) { + checkIfPausedAndHandleAccordingly(); // Always receive messages in batch mode. try { if (!inRetryMode.get() && !messagesPendingInBatch.get()) { @@ -442,6 +458,23 @@ public class DefaultPulsarMessageListenerContainer extends AbstractPulsarMess } } + private void checkIfPausedAndHandleAccordingly() { + if (isPaused()) { + // try acquiring the lock. + DefaultPulsarMessageListenerContainer.this.lockOnPause.lock(); + try { + // Waiting on lock's condition. + DefaultPulsarMessageListenerContainer.this.pausedCondition.await(); + } + catch (InterruptedException e) { + throw new IllegalStateException("Exception occurred trying to wake up the paused listener thread."); + } + finally { + DefaultPulsarMessageListenerContainer.this.lockOnPause.unlock(); + } + } + } + private Observation newObservation(Message message) { if (this.observationRegistry == null) { return Observation.NOOP; diff --git a/spring-pulsar/src/test/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainerTests.java b/spring-pulsar/src/test/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainerTests.java index c8033d88..13f59f74 100644 --- a/spring-pulsar/src/test/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainerTests.java +++ b/spring-pulsar/src/test/java/org/springframework/pulsar/listener/DefaultPulsarMessageListenerContainerTests.java @@ -19,6 +19,7 @@ package org.springframework.pulsar.listener; import static org.assertj.core.api.Assertions.assertThat; import static org.awaitility.Awaitility.await; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.doAnswer; import static org.mockito.Mockito.spy; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; @@ -30,8 +31,12 @@ import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.Set; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; +import java.util.concurrent.locks.Condition; +import java.util.concurrent.locks.Lock; +import java.util.concurrent.locks.ReentrantLock; import org.apache.pulsar.client.api.Consumer; import org.apache.pulsar.client.api.DeadLetterPolicy; @@ -42,6 +47,7 @@ import org.apache.pulsar.client.api.Schema; import org.apache.pulsar.client.api.SubscriptionInitialPosition; import org.apache.pulsar.client.api.SubscriptionType; import org.apache.pulsar.client.impl.MultiplierRedeliveryBackoff; +import org.awaitility.Awaitility; import org.junit.jupiter.api.Test; import org.springframework.pulsar.core.ConsumerTestUtils; @@ -49,6 +55,7 @@ import org.springframework.pulsar.core.DefaultPulsarConsumerFactory; import org.springframework.pulsar.core.DefaultPulsarProducerFactory; import org.springframework.pulsar.core.PulsarTemplate; import org.springframework.pulsar.test.support.PulsarTestContainerSupport; +import org.springframework.test.util.ReflectionTestUtils; /** * @author Soby Chacko @@ -76,6 +83,7 @@ class DefaultPulsarMessageListenerContainerTests implements PulsarTestContainerS DefaultPulsarMessageListenerContainer container = new DefaultPulsarMessageListenerContainer<>( pulsarConsumerFactory, pulsarContainerProperties); container.start(); + Map prodConfig = new HashMap<>(); prodConfig.put("topicName", "dpmlct-012"); DefaultPulsarProducerFactory pulsarProducerFactory = new DefaultPulsarProducerFactory<>(pulsarClient, @@ -87,6 +95,75 @@ class DefaultPulsarMessageListenerContainerTests implements PulsarTestContainerS pulsarClient.close(); } + @Test + void containerPauseAndResumeFeatureUsingWaitAndNotify() throws Exception { + Set topics = Collections.singleton("containerPauseResumeWaitNotify-topic"); + Map config = Map.of("topicNames", topics, "subscriptionName", + "containerPauseResumeWaitNotify-sub"); + PulsarClient pulsarClient = PulsarClient.builder().serviceUrl(PulsarTestContainerSupport.getPulsarBrokerUrl()) + .build(); + DefaultPulsarConsumerFactory pulsarConsumerFactory = new DefaultPulsarConsumerFactory<>(pulsarClient, + config); + PulsarContainerProperties pulsarContainerProperties = new PulsarContainerProperties(); + pulsarContainerProperties.setMessageListener((PulsarRecordMessageListener) (consumer, msg) -> { + }); + pulsarContainerProperties.setSchema(Schema.STRING); + DefaultPulsarMessageListenerContainer container = new DefaultPulsarMessageListenerContainer<>( + pulsarConsumerFactory, pulsarContainerProperties); + + Lock reentrantLock = new ReentrantLock(); + Condition lockCondition = reentrantLock.newCondition(); + + Lock spyLock = spy(reentrantLock); + Condition spyCondition = spy(lockCondition); + + ReflectionTestUtils.setField(container, "lockOnPause", spyLock); + ReflectionTestUtils.setField(container, "pausedCondition", spyCondition); + + CountDownLatch latchOnLockInvocation = new CountDownLatch(2); + CountDownLatch latchOnUnlockInvocation = new CountDownLatch(2); + CountDownLatch latchOnAwaitInvocation = new CountDownLatch(1); + CountDownLatch latchOnSignalInvocation = new CountDownLatch(1); + + doAnswer(invocation -> { + latchOnLockInvocation.countDown(); + return invocation.callRealMethod(); + }).when(spyLock).lock(); + + doAnswer(invocation -> { + latchOnUnlockInvocation.countDown(); + return invocation.callRealMethod(); + }).when(spyLock).unlock(); + + doAnswer(invocation -> { + latchOnAwaitInvocation.countDown(); + return invocation.callRealMethod(); + }).when(spyCondition).await(); + + doAnswer(invocation -> { + latchOnSignalInvocation.countDown(); + return invocation.callRealMethod(); + }).when(spyCondition).signal(); + + container.start(); + + container.pause(); + + Awaitility.await().until(container::isPaused); + + container.resume(); + + Awaitility.await().until(() -> !container.isPaused()); + + assertThat(latchOnLockInvocation.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(latchOnUnlockInvocation.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(latchOnAwaitInvocation.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(latchOnSignalInvocation.await(10, TimeUnit.SECONDS)).isTrue(); + + container.stop(); + pulsarClient.close(); + } + @Test void subscriptionInitialPositionEarliest() throws Exception { Map config = new HashMap<>();