diff --git a/connect/mirror/src/test/java/org/apache/kafka/connect/mirror/integration/MirrorConnectorsIntegrationBaseTest.java b/connect/mirror/src/test/java/org/apache/kafka/connect/mirror/integration/MirrorConnectorsIntegrationBaseTest.java index 56d2bf4974093..978c881176cb4 100644 --- a/connect/mirror/src/test/java/org/apache/kafka/connect/mirror/integration/MirrorConnectorsIntegrationBaseTest.java +++ b/connect/mirror/src/test/java/org/apache/kafka/connect/mirror/integration/MirrorConnectorsIntegrationBaseTest.java @@ -28,11 +28,12 @@ import org.apache.kafka.clients.admin.TopicDescription; import org.apache.kafka.clients.consumer.Consumer; import org.apache.kafka.clients.consumer.ConsumerConfig; -import org.apache.kafka.clients.consumer.ConsumerRebalanceListener; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.KafkaConsumer; import org.apache.kafka.clients.consumer.OffsetAndMetadata; +import org.apache.kafka.clients.consumer.RebalanceConsumer; +import org.apache.kafka.clients.consumer.RebalanceListener; import org.apache.kafka.clients.producer.KafkaProducer; import org.apache.kafka.clients.producer.Producer; import org.apache.kafka.clients.producer.ProducerRecord; @@ -1536,14 +1537,14 @@ protected final void warmUpConsumer(Map consumerProps) { private void warmUpConsumer(String clusterName, EmbeddedKafkaCluster kafkaCluster, Map consumerProps, String topic) { AtomicBoolean joinedGroup = new AtomicBoolean(false); - ConsumerRebalanceListener rebalanceListener = new ConsumerRebalanceListener() { + RebalanceListener rebalanceListener = new RebalanceListener() { @Override - public void onPartitionsRevoked(Collection partitions) { + public void onPartitionsRevoked(Collection partitions, RebalanceConsumer consumer) { // no-op } @Override - public void onPartitionsAssigned(Collection partitions) { + public void onPartitionsAssigned(Collection partitions, RebalanceConsumer consumer) { joinedGroup.set(true); } }; diff --git a/connect/runtime/src/main/java/org/apache/kafka/connect/runtime/WorkerSinkTask.java b/connect/runtime/src/main/java/org/apache/kafka/connect/runtime/WorkerSinkTask.java index 1de9ff2d9a56e..d1defc974c258 100644 --- a/connect/runtime/src/main/java/org/apache/kafka/connect/runtime/WorkerSinkTask.java +++ b/connect/runtime/src/main/java/org/apache/kafka/connect/runtime/WorkerSinkTask.java @@ -17,11 +17,12 @@ package org.apache.kafka.connect.runtime; import org.apache.kafka.clients.consumer.Consumer; -import org.apache.kafka.clients.consumer.ConsumerRebalanceListener; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.OffsetAndMetadata; import org.apache.kafka.clients.consumer.OffsetCommitCallback; +import org.apache.kafka.clients.consumer.RebalanceConsumer; +import org.apache.kafka.clients.consumer.RebalanceListener; import org.apache.kafka.common.KafkaException; import org.apache.kafka.common.TopicPartition; import org.apache.kafka.common.errors.WakeupException; @@ -323,14 +324,15 @@ public int commitFailures() { @Override protected void initializeAndStart() { SinkConnectorConfig.validate(taskConfig); + consumer.setRebalanceListener(new HandleRebalance()); if (SinkConnectorConfig.hasTopicsConfig(taskConfig)) { List topics = SinkConnectorConfig.parseTopicsList(taskConfig); - consumer.subscribe(topics, new HandleRebalance()); + consumer.subscribe(topics); log.debug("{} Initializing and starting task for topics {}", this, String.join(", ", topics)); } else { String topicsRegexStr = taskConfig.get(SinkTask.TOPICS_REGEX_CONFIG); Pattern pattern = Pattern.compile(topicsRegexStr); - consumer.subscribe(pattern, new HandleRebalance()); + consumer.subscribe(pattern); log.debug("{} Initializing and starting task for topics regex {}", this, topicsRegexStr); } @@ -729,9 +731,9 @@ long getNextCommit() { return nextCommit; } - private class HandleRebalance implements ConsumerRebalanceListener { + private class HandleRebalance implements RebalanceListener { @Override - public void onPartitionsAssigned(Collection partitions) { + public void onPartitionsAssigned(Collection partitions, RebalanceConsumer rebalanceConsumer) { log.debug("{} Partitions assigned {}", WorkerSinkTask.this, partitions); for (TopicPartition tp : partitions) { @@ -783,12 +785,12 @@ else if (!context.pausedPartitions().isEmpty()) } @Override - public void onPartitionsRevoked(Collection partitions) { + public void onPartitionsRevoked(Collection partitions, RebalanceConsumer rebalanceConsumer) { onPartitionsRemoved(partitions, false); } @Override - public void onPartitionsLost(Collection partitions) { + public void onPartitionsLost(Collection partitions, RebalanceConsumer rebalanceConsumer) { onPartitionsRemoved(partitions, true); } diff --git a/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/ErrorHandlingTaskTest.java b/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/ErrorHandlingTaskTest.java index a9e5f289732e5..b27ac15ebeef8 100644 --- a/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/ErrorHandlingTaskTest.java +++ b/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/ErrorHandlingTaskTest.java @@ -17,11 +17,11 @@ package org.apache.kafka.connect.runtime; import org.apache.kafka.clients.admin.NewTopic; -import org.apache.kafka.clients.consumer.ConsumerRebalanceListener; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.KafkaConsumer; import org.apache.kafka.clients.consumer.OffsetAndMetadata; +import org.apache.kafka.clients.consumer.RebalanceListener; import org.apache.kafka.clients.producer.KafkaProducer; import org.apache.kafka.common.TopicPartition; import org.apache.kafka.common.config.ConfigDef; @@ -389,8 +389,8 @@ private void assertSinkMetricValue(String name, double expected) { private void verifyInitializeSink() { verify(sinkTask).start(TASK_PROPS); verify(sinkTask).initialize(any(WorkerSinkTaskContext.class)); - verify(consumer).subscribe(eq(List.of(TOPIC)), - any(ConsumerRebalanceListener.class)); + verify(consumer).setRebalanceListener(any(RebalanceListener.class)); + verify(consumer).subscribe(eq(List.of(TOPIC))); } private void assertSourceMetricValue(String name, double expected) { diff --git a/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskTest.java b/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskTest.java index 4815b79019c26..3424a49a2d793 100644 --- a/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskTest.java +++ b/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskTest.java @@ -17,13 +17,13 @@ package org.apache.kafka.connect.runtime; import org.apache.kafka.clients.consumer.Consumer; -import org.apache.kafka.clients.consumer.ConsumerRebalanceListener; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.KafkaConsumer; import org.apache.kafka.clients.consumer.MockConsumer; import org.apache.kafka.clients.consumer.OffsetAndMetadata; import org.apache.kafka.clients.consumer.OffsetCommitCallback; +import org.apache.kafka.clients.consumer.RebalanceListener; import org.apache.kafka.clients.consumer.internals.AutoOffsetResetStrategy; import org.apache.kafka.common.MetricName; import org.apache.kafka.common.TopicPartition; @@ -170,7 +170,7 @@ public class WorkerSinkTaskTest { private KafkaConsumer consumer; @Mock private ErrorHandlingMetrics errorHandlingMetrics; - private final ArgumentCaptor rebalanceListener = ArgumentCaptor.forClass(ConsumerRebalanceListener.class); + private final ArgumentCaptor rebalanceListener = ArgumentCaptor.forClass(RebalanceListener.class); private long recordsReturnedTp1; private long recordsReturnedTp3; @@ -366,7 +366,7 @@ public void testShutdown() throws Exception { verify(sinkTask, times(2)).put(anyList()); doAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT, null); return null; }).when(consumer).close(); @@ -494,14 +494,14 @@ public void testPollRedeliveryWithConsumerRebalance() { when(consumer.poll(any(Duration.class))) .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }) .thenAnswer(expectConsumerPoll(1)) // Empty consumer poll (all partitions are paused) with rebalance; one new partition is assigned .thenAnswer(invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(Set.of()); - rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION3)); + rebalanceListener.getValue().onPartitionsRevoked(Set.of(), null); + rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION3), null); return ConsumerRecords.empty(); }) .thenAnswer(expectConsumerPoll(0)) @@ -509,8 +509,8 @@ public void testPollRedeliveryWithConsumerRebalance() { .thenAnswer(invocation -> { ConsumerRecord newRecord = new ConsumerRecord<>(TOPIC, PARTITION3, FIRST_OFFSET, RAW_KEY, RAW_VALUE); - rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT); - rebalanceListener.getValue().onPartitionsAssigned(List.of()); + rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT, null); + rebalanceListener.getValue().onPartitionsAssigned(List.of(), null); return new ConsumerRecords<>(Map.of(TOPIC_PARTITION3, List.of(newRecord)), Map.of(TOPIC_PARTITION3, new OffsetAndMetadata(FIRST_OFFSET + 1, Optional.empty(), ""))); }); @@ -560,7 +560,7 @@ public void testErrorInRebalancePartitionLoss() { expectPollInitialAssignment() .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsLost(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsLost(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }); @@ -584,7 +584,7 @@ public void testErrorInRebalancePartitionRevocation() { expectPollInitialAssignment() .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }); @@ -608,8 +608,8 @@ public void testErrorInRebalancePartitionAssignment() { expectPollInitialAssignment() .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT); - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT, null); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }); @@ -649,22 +649,22 @@ public void testPartialRevocationAndAssignment() { when(consumer.poll(any(Duration.class))) .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }) .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(Set.of(TOPIC_PARTITION)); - rebalanceListener.getValue().onPartitionsAssigned(Set.of()); + rebalanceListener.getValue().onPartitionsRevoked(Set.of(TOPIC_PARTITION), null); + rebalanceListener.getValue().onPartitionsAssigned(Set.of(), null); return ConsumerRecords.empty(); }) .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(Set.of()); - rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION3)); + rebalanceListener.getValue().onPartitionsRevoked(Set.of(), null); + rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION3), null); return ConsumerRecords.empty(); }) .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsLost(Set.of(TOPIC_PARTITION3)); - rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION)); + rebalanceListener.getValue().onPartitionsLost(Set.of(TOPIC_PARTITION3), null); + rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION), null); return ConsumerRecords.empty(); }); @@ -720,21 +720,21 @@ public void testPreCommitFailureAfterPartialRevocationAndAssignment() { // First poll; assignment is [TP1, TP2] when(consumer.poll(any(Duration.class))) .thenAnswer((Answer>) invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }) // Second poll; a single record is delivered from TP1 .thenAnswer(expectConsumerPoll(1)) // Third poll; assignment changes to [TP2] .thenAnswer(invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(Set.of(TOPIC_PARTITION)); - rebalanceListener.getValue().onPartitionsAssigned(Set.of()); + rebalanceListener.getValue().onPartitionsRevoked(Set.of(TOPIC_PARTITION), null); + rebalanceListener.getValue().onPartitionsAssigned(Set.of(), null); return ConsumerRecords.empty(); }) // Fourth poll; assignment changes to [TP2, TP3] .thenAnswer(invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(Set.of()); - rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION3)); + rebalanceListener.getValue().onPartitionsRevoked(Set.of(), null); + rebalanceListener.getValue().onPartitionsAssigned(Set.of(TOPIC_PARTITION3), null); return ConsumerRecords.empty(); }) // Fifth poll; an offset commit takes place @@ -788,8 +788,8 @@ public void testWakeupInCommitSyncCausesRetry() { expectPollInitialAssignment() .thenAnswer(expectConsumerPoll(1)) .thenAnswer(invocation -> { - rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT); - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsRevoked(INITIAL_ASSIGNMENT, null); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }); expectConversionAndTransformation(null, new RecordHeaders()); @@ -1375,7 +1375,7 @@ public void testCommitWithOutOfOrderCallback() { // iter 1 Answer> consumerPollRebalance = invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }; @@ -1425,14 +1425,14 @@ public void testCommitWithOutOfOrderCallback() { final AtomicBoolean rebalanced = new AtomicBoolean(); Answer> consumerPollRebalanced = invocation -> { // Rebalance always begins with revoking current partitions ... - rebalanceListener.getValue().onPartitionsRevoked(originalPartitions); + rebalanceListener.getValue().onPartitionsRevoked(originalPartitions, null); // Respond to the rebalance Map offsets = new HashMap<>(); offsets.put(TOPIC_PARTITION, rebalanceOffsets.get(TOPIC_PARTITION).offset()); offsets.put(TOPIC_PARTITION2, rebalanceOffsets.get(TOPIC_PARTITION2).offset()); offsets.put(TOPIC_PARTITION3, rebalanceOffsets.get(TOPIC_PARTITION3).offset()); sinkTaskContext.getValue().offset(offsets); - rebalanceListener.getValue().onPartitionsAssigned(rebalancedPartitions); + rebalanceListener.getValue().onPartitionsAssigned(rebalancedPartitions, null); rebalanced.set(true); // Run the previous async commit handler @@ -1689,7 +1689,8 @@ public void testTopicsRegex() { ArgumentCaptor topicsRegex = ArgumentCaptor.forClass(Pattern.class); - verify(consumer).subscribe(topicsRegex.capture(), rebalanceListener.capture()); + verify(consumer).setRebalanceListener(rebalanceListener.capture()); + verify(consumer).subscribe(topicsRegex.capture()); assertEquals("te.*", topicsRegex.getValue().pattern()); verify(sinkTask).initialize(sinkTaskContext.capture()); verify(sinkTask).start(props); @@ -1915,7 +1916,8 @@ private void expectRebalanceAssignmentError(RuntimeException e) { } private void verifyInitializeTask() { - verify(consumer).subscribe(eq(List.of(TOPIC)), rebalanceListener.capture()); + verify(consumer).setRebalanceListener(rebalanceListener.capture()); + verify(consumer).subscribe(eq(List.of(TOPIC))); verify(sinkTask).initialize(sinkTaskContext.capture()); verify(sinkTask).start(TASK_PROPS); } @@ -1926,7 +1928,7 @@ private OngoingStubbing> expectPollInitialAssign return when(consumer.poll(any(Duration.class))).thenAnswer( invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); } ); diff --git a/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskThreadedTest.java b/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskThreadedTest.java index 729b5f0436c2b..36a5f3ce0fe34 100644 --- a/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskThreadedTest.java +++ b/connect/runtime/src/test/java/org/apache/kafka/connect/runtime/WorkerSinkTaskThreadedTest.java @@ -16,12 +16,12 @@ */ package org.apache.kafka.connect.runtime; -import org.apache.kafka.clients.consumer.ConsumerRebalanceListener; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.KafkaConsumer; import org.apache.kafka.clients.consumer.OffsetAndMetadata; import org.apache.kafka.clients.consumer.OffsetCommitCallback; +import org.apache.kafka.clients.consumer.RebalanceListener; import org.apache.kafka.common.TopicPartition; import org.apache.kafka.common.header.internals.RecordHeaders; import org.apache.kafka.common.internals.Plugin; @@ -141,7 +141,7 @@ public class WorkerSinkTaskThreadedTest { private WorkerSinkTask workerTask; @Mock private KafkaConsumer consumer; - private final ArgumentCaptor rebalanceListener = ArgumentCaptor.forClass(ConsumerRebalanceListener.class); + private final ArgumentCaptor rebalanceListener = ArgumentCaptor.forClass(RebalanceListener.class); @Mock private TaskStatus.Listener statusListener; @Mock @@ -555,7 +555,8 @@ public void testRewindOnRebalanceDuringPoll() { } private void verifyInitializeTask() { - verify(consumer).subscribe(eq(List.of(TOPIC)), rebalanceListener.capture()); + verify(consumer).setRebalanceListener(rebalanceListener.capture()); + verify(consumer).subscribe(eq(List.of(TOPIC))); verify(sinkTask).initialize(sinkTaskContext.capture()); verify(sinkTask).start(TASK_PROPS); } @@ -592,7 +593,7 @@ private void expectPolls(final long pollDelayMs) { // Stub out all the consumer stream/iterator responses, which we just want to verify occur, // but don't care about the exact details here. when(consumer.poll(any(Duration.class))).thenAnswer(invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }).thenAnswer((Answer>) invocation -> { // "Sleep" so time will progress @@ -618,14 +619,14 @@ private void expectRebalanceDuringPoll(long startOffset) { offsets.put(TOPIC_PARTITION, startOffset); when(consumer.poll(any(Duration.class))).thenAnswer(invocation -> { - rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT); + rebalanceListener.getValue().onPartitionsAssigned(INITIAL_ASSIGNMENT, null); return ConsumerRecords.empty(); }).thenAnswer((Answer>) invocation -> { // "Sleep" so time will progress time.sleep(1L); sinkTaskContext.getValue().offset(offsets); - rebalanceListener.getValue().onPartitionsAssigned(partitions); + rebalanceListener.getValue().onPartitionsAssigned(partitions, null); TopicPartition topicPartition = new TopicPartition(TOPIC, PARTITION); ConsumerRecord consumerRecord = new ConsumerRecord<>( diff --git a/connect/runtime/src/testFixtures/java/org/apache/kafka/connect/util/clusters/EmbeddedKafkaCluster.java b/connect/runtime/src/testFixtures/java/org/apache/kafka/connect/util/clusters/EmbeddedKafkaCluster.java index 7913d60fc2837..f11817dd28d7b 100644 --- a/connect/runtime/src/testFixtures/java/org/apache/kafka/connect/util/clusters/EmbeddedKafkaCluster.java +++ b/connect/runtime/src/testFixtures/java/org/apache/kafka/connect/util/clusters/EmbeddedKafkaCluster.java @@ -29,11 +29,11 @@ import org.apache.kafka.clients.admin.OffsetSpec; import org.apache.kafka.clients.admin.TopicDescription; import org.apache.kafka.clients.consumer.Consumer; -import org.apache.kafka.clients.consumer.ConsumerRebalanceListener; import org.apache.kafka.clients.consumer.ConsumerRecord; import org.apache.kafka.clients.consumer.ConsumerRecords; import org.apache.kafka.clients.consumer.KafkaConsumer; import org.apache.kafka.clients.consumer.OffsetAndMetadata; +import org.apache.kafka.clients.consumer.RebalanceListener; import org.apache.kafka.clients.producer.KafkaProducer; import org.apache.kafka.clients.producer.ProducerConfig; import org.apache.kafka.clients.producer.ProducerRecord; @@ -658,13 +658,12 @@ public KafkaConsumer createConsumerAndSubscribeTo(Map createConsumerAndSubscribeTo(Map consumerProps, ConsumerRebalanceListener rebalanceListener, String... topics) { + public KafkaConsumer createConsumerAndSubscribeTo(Map consumerProps, RebalanceListener rebalanceListener, String... topics) { KafkaConsumer consumer = createConsumer(consumerProps); if (rebalanceListener != null) { - consumer.subscribe(List.of(topics), rebalanceListener); - } else { - consumer.subscribe(List.of(topics)); + consumer.setRebalanceListener(rebalanceListener); } + consumer.subscribe(List.of(topics)); return consumer; }