diff --git a/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/inbound/KafkaMessageSource.java b/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/inbound/KafkaMessageSource.java index 6e2510a417..d238582c85 100644 --- a/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/inbound/KafkaMessageSource.java +++ b/spring-integration-kafka/src/main/java/org/springframework/integration/kafka/inbound/KafkaMessageSource.java @@ -23,6 +23,7 @@ import java.util.Collection; import java.util.Collections; import java.util.HashMap; import java.util.Iterator; +import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Set; @@ -120,6 +121,8 @@ public class KafkaMessageSource extends AbstractMessageSource impl private final ConsumerProperties consumerProperties; + private final Collection assignedPartitions = new LinkedHashSet<>(); + private Duration pollTimeout; private RecordMessageConverter messageConverter = new MessagingMessageConverter(); @@ -136,8 +139,6 @@ public class KafkaMessageSource extends AbstractMessageSource impl private volatile Consumer consumer; - private volatile Collection assignedPartitions = new ArrayList<>(); - private volatile boolean pausing; private volatile boolean paused; @@ -233,6 +234,24 @@ public class KafkaMessageSource extends AbstractMessageSource impl this.commitTimeout = consumerProperties.getSyncCommitTimeout(); } + /** + * Return the currently assigned partitions. + * @return the partitions. + * @since 3.2.2 + */ + public Collection getAssignedPartitions() { + return Collections.unmodifiableCollection(this.assignedPartitions); + } + + /** + * Return true if the source is currently paused. + * @return true if paused. + * @since 3.2.2 + */ + public boolean isPaused() { + return this.paused; + } + @Override protected void onInit() { if (!StringUtils.hasText(this.consumerProperties.getClientId())) { @@ -478,7 +497,7 @@ public class KafkaMessageSource extends AbstractMessageSource impl @Override public void onPartitionsRevoked(Collection partitions) { - KafkaMessageSource.this.assignedPartitions.clear(); + KafkaMessageSource.this.assignedPartitions.removeAll(partitions); if (KafkaMessageSource.this.logger.isInfoEnabled()) { KafkaMessageSource.this.logger .info("Partitions revoked: " + partitions); @@ -495,9 +514,27 @@ public class KafkaMessageSource extends AbstractMessageSource impl } + @Override + public void onPartitionsLost(Collection partitions) { + if (providedRebalanceListener != null) { + if (isConsumerAware) { + ((ConsumerAwareRebalanceListener) providedRebalanceListener).onPartitionsLost(partitions); + } + else { + providedRebalanceListener.onPartitionsLost(partitions); + } + } + onPartitionsRevoked(partitions); + } + @Override public void onPartitionsAssigned(Collection partitions) { - KafkaMessageSource.this.assignedPartitions = new ArrayList<>(partitions); + KafkaMessageSource.this.assignedPartitions.addAll(partitions); + if (KafkaMessageSource.this.paused) { + KafkaMessageSource.this.consumer.pause(KafkaMessageSource.this.assignedPartitions); + KafkaMessageSource.this.logger.warn("Paused consumer resumed by Kafka due to rebalance; " + + "consumer paused again, so the initial poll() will never return any records"); + } if (KafkaMessageSource.this.logger.isInfoEnabled()) { KafkaMessageSource.this.logger .info("Partitions assigned: " + partitions); @@ -524,7 +561,7 @@ public class KafkaMessageSource extends AbstractMessageSource impl .map(TopicPartitionOffset::getTopicPartition) .collect(Collectors.toList()); this.consumer.assign(topicPartitionsToAssign); - this.assignedPartitions = new ArrayList<>(topicPartitionsToAssign); + this.assignedPartitions.addAll(topicPartitionsToAssign); TopicPartitionOffset[] partitions = this.consumerProperties.getTopicPartitions(); diff --git a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageSourceTests.java b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageSourceTests.java index 2f90beb287..764d487265 100644 --- a/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageSourceTests.java +++ b/spring-integration-kafka/src/test/java/org/springframework/integration/kafka/inbound/MessageSourceTests.java @@ -35,11 +35,13 @@ import static org.mockito.Mockito.times; import java.time.Duration; import java.time.temporal.ChronoUnit; +import java.util.ArrayList; import java.util.Arrays; import java.util.Collection; import java.util.Collections; import java.util.HashSet; import java.util.LinkedHashMap; +import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Set; @@ -157,8 +159,10 @@ class MessageSourceTests { @Test void testRebalanceListener() { Consumer consumer = mock(Consumer.class); - TopicPartition topicPartition = new TopicPartition("foo", 0); - List assigned = Collections.singletonList(topicPartition); + TopicPartition topicPartition1 = new TopicPartition("foo", 0); + List assigned1 = new ArrayList<>(Collections.singletonList(topicPartition1)); + TopicPartition topicPartition2 = new TopicPartition("foo", 1); + List assigned2 = new ArrayList<>(Collections.singletonList(topicPartition2)); AtomicReference listener = new AtomicReference<>(); willAnswer(i -> { listener.set(i.getArgument(1)); @@ -190,11 +194,27 @@ class MessageSourceTests { source.receive(); - listener.get().onPartitionsAssigned(assigned); + listener.get().onPartitionsAssigned(assigned1); assertThat(partitionsAssignedCalled.get()).isTrue(); + assertThat(new ArrayList<>(source.getAssignedPartitions())).isEqualTo(assigned1); + listener.get().onPartitionsAssigned(assigned2); + List temp = new ArrayList<>(assigned1); + temp.addAll(assigned2); + assertThat(new ArrayList<>(source.getAssignedPartitions())).isEqualTo(temp); - listener.get().onPartitionsRevoked(assigned); + listener.get().onPartitionsRevoked(assigned1); assertThat(partitionsRevokedCalled.get()).isTrue(); + assertThat(new ArrayList<>(source.getAssignedPartitions())).isEqualTo(assigned2); + + source.pause(); + assertThat(source.isPaused()).isFalse(); + InOrder inOrder = inOrder(consumer); + source.receive(); + assertThat(source.isPaused()).isTrue(); + inOrder.verify(consumer).pause(new LinkedHashSet<>(assigned2)); + inOrder.verify(consumer).poll(any()); + listener.get().onPartitionsAssigned(assigned1); + inOrder.verify(consumer).pause(new LinkedHashSet<>(temp)); } @Test