This is an automated email from the ASF dual-hosted git repository.

mjsax pushed a commit to branch trunk
in repository https://gitbox.apache.org/repos/asf/kafka.git


The following commit(s) were added to refs/heads/trunk by this push:
     new 89bc8808f90 KAFKA-20116: Make task-offset-sum available to client 
background thread (1/N) (#22595)
89bc8808f90 is described below

commit 89bc8808f9038491bd7e746e2e8936f3d9a897bd
Author: Matthias J. Sax <[email protected]>
AuthorDate: Mon Jun 22 14:42:29 2026 -0700

    KAFKA-20116: Make task-offset-sum available to client background thread 
(1/N) (#22595)
    
    This PR adds base logic to make task-offset-sum available to the client
    background thread, allowing us to add task-offset-sum to the
    StreamsHeartbeatRequest in a follow up PR. We are getting the
    task-offset-sum from StateUpdater thread, and thus the supplier must be
    thread-safe.
    
    Part of KIP-1071.
    
    Reviewers: Lucas Brutschy <[email protected]>
---
 .../consumer/internals/StreamsRebalanceData.java   | 11 +++-
 .../consumer/internals/AsyncKafkaConsumerTest.java | 10 ++--
 .../consumer/internals/RequestManagersTest.java    |  2 +-
 .../StreamsGroupHeartbeatRequestManagerTest.java   |  3 +-
 .../internals/StreamsRebalanceDataTest.java        | 52 +++++++++++-------
 .../kafka/api/AuthorizerIntegrationTest.scala      |  3 +-
 .../kafka/api/IntegrationTestHarness.scala         |  3 +-
 .../processor/internals/StateDirectory.java        |  4 ++
 .../streams/processor/internals/StreamThread.java  | 24 +++++++--
 .../streams/processor/internals/TaskManager.java   | 44 ++++++++++++++++
 .../DefaultStreamsRebalanceListenerTest.java       | 18 ++++---
 .../processor/internals/StreamThreadTest.java      | 61 +++++++++++++---------
 .../processor/internals/TaskManagerTest.java       | 39 ++++++++++++++
 13 files changed, 206 insertions(+), 68 deletions(-)

diff --git 
a/clients/src/main/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceData.java
 
b/clients/src/main/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceData.java
index f4699825b91..286bcdb540c 100644
--- 
a/clients/src/main/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceData.java
+++ 
b/clients/src/main/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceData.java
@@ -33,6 +33,7 @@ import java.util.concurrent.atomic.AtomicBoolean;
 import java.util.concurrent.atomic.AtomicInteger;
 import java.util.concurrent.atomic.AtomicLong;
 import java.util.concurrent.atomic.AtomicReference;
+import java.util.function.Supplier;
 
 /**
  * This class holds the data that is needed to participate in the Streams 
rebalance protocol.
@@ -337,6 +338,8 @@ public class StreamsRebalanceData {
 
     private final Map<String, Subtopology> subtopologies;
 
+    private final Supplier<Map<TaskId, Long>> taskOffsetSum;
+
     private final AtomicReference<Assignment> reconciledAssignment = new 
AtomicReference<>(Assignment.EMPTY);
 
     private final AtomicReference<Map<HostInfo, EndpointPartitions>> 
partitionsByHost = new AtomicReference<>(Collections.emptyMap());
@@ -355,12 +358,14 @@ public class StreamsRebalanceData {
                                 final Optional<HostInfo> endpoint,
                                 final Optional<String> rackId,
                                 final Map<String, Subtopology> subtopologies,
-                                final Map<String, String> clientTags) {
+                                final Map<String, String> clientTags,
+                                final Supplier<Map<TaskId, Long>> 
taskOffsetSum) {
         this.processId = Objects.requireNonNull(processId, "Process ID cannot 
be null");
         this.endpoint = Objects.requireNonNull(endpoint, "Endpoint cannot be 
null");
         this.rackId = Objects.requireNonNull(rackId, "Rack ID cannot be null");
         this.subtopologies = Map.copyOf(Objects.requireNonNull(subtopologies, 
"Subtopologies cannot be null"));
         this.clientTags = Map.copyOf(Objects.requireNonNull(clientTags, 
"Client tags cannot be null"));
+        this.taskOffsetSum = Objects.requireNonNull(taskOffsetSum, "Task 
offset sum supplier cannot be null");
     }
 
     public UUID processId() {
@@ -383,6 +388,10 @@ public class StreamsRebalanceData {
         return subtopologies;
     }
 
+    public Map<TaskId, Long> taskOffsetSum() {
+        return taskOffsetSum.get();
+    }
+
     public int topologyEpoch() {
         return 0;
     }
diff --git 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/AsyncKafkaConsumerTest.java
 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/AsyncKafkaConsumerTest.java
index 6f9ef98a66f..59442cb3205 100644
--- 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/AsyncKafkaConsumerTest.java
+++ 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/AsyncKafkaConsumerTest.java
@@ -1664,7 +1664,7 @@ public class AsyncKafkaConsumerTest {
     public void testStreamRebalanceData() {
         final String groupId = "consumerGroupA";
         try (final MockedStatic<RequestManagers> requestManagers = 
mockStatic(RequestManagers.class)) {
-            StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of());
+            StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of);
             consumer = 
newConsumerWithStreamRebalanceData(requiredConsumerConfigAndGroupId(groupId), 
streamsRebalanceData);
             final Optional<StreamsRebalanceData> groupMetadataUpdateListener = 
captureStreamRebalanceData(requestManagers);
             assertTrue(groupMetadataUpdateListener.isPresent());
@@ -2588,7 +2588,7 @@ public class AsyncKafkaConsumerTest {
     @Test
     public void 
testCloseInvokesStreamsRebalanceListenerOnTasksRevokedWhenMemberEpochPositive() 
{
         final String groupId = "streamsGroup";
-        final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of());
+        final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of);
         
         try (final MockedStatic<RequestManagers> requestManagers = 
mockStatic(RequestManagers.class)) {
             consumer = 
newConsumerWithStreamRebalanceData(requiredConsumerConfigAndGroupId(groupId), 
streamsRebalanceData);
@@ -2608,7 +2608,7 @@ public class AsyncKafkaConsumerTest {
     @Test
     public void 
testCloseInvokesStreamsRebalanceListenerOnAllTasksLostWhenMemberEpochZeroOrNegative()
 {
         final String groupId = "streamsGroup";
-        final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of());
+        final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of);
         
         try (final MockedStatic<RequestManagers> requestManagers = 
mockStatic(RequestManagers.class)) {
             consumer = 
newConsumerWithStreamRebalanceData(requiredConsumerConfigAndGroupId(groupId), 
streamsRebalanceData);
@@ -2628,7 +2628,7 @@ public class AsyncKafkaConsumerTest {
     @Test
     public void testCloseWrapsStreamsRebalanceListenerException() {
         final String groupId = "streamsGroup";
-        final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of());
+        final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of);
         
         try (final MockedStatic<RequestManagers> requestManagers = 
mockStatic(RequestManagers.class)) {
             consumer = 
newConsumerWithStreamRebalanceData(requiredConsumerConfigAndGroupId(groupId), 
streamsRebalanceData);
@@ -2693,7 +2693,7 @@ public class AsyncKafkaConsumerTest {
     @Test
     public void 
testStreamsTasksAssignedEventSendsErrorWhenApplyAssignmentFails() {
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-            UUID.randomUUID(), Optional.empty(), Optional.empty(), Map.of(), 
Map.of());
+            UUID.randomUUID(), Optional.empty(), Optional.empty(), Map.of(), 
Map.of(), Map::of);
         final InterruptException applyAssignmentError = new 
InterruptException("Thread was interrupted");
 
         consumer = newConsumerWithStreamRebalanceData(
diff --git 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/RequestManagersTest.java
 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/RequestManagersTest.java
index 78cb459eb1d..3b4eda0cc3f 100644
--- 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/RequestManagersTest.java
+++ 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/RequestManagersTest.java
@@ -117,7 +117,7 @@ public class RequestManagersTest {
             new Metrics(),
             mock(OffsetCommitCallbackInvoker.class),
             listener,
-            Optional.of(new StreamsRebalanceData(UUID.randomUUID(), 
Optional.empty(), Optional.empty(), Map.of(), Map.of())),
+            Optional.of(new StreamsRebalanceData(UUID.randomUUID(), 
Optional.empty(), Optional.empty(), Map.of(), Map.of(), Map::of)),
             new PositionsValidator(logContext, time, subscriptions, metadata)
         ).get();
         assertTrue(requestManagers.streamsMembershipManager.isPresent());
diff --git 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsGroupHeartbeatRequestManagerTest.java
 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsGroupHeartbeatRequestManagerTest.java
index 3ede3565c11..214d824ed5e 100644
--- 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsGroupHeartbeatRequestManagerTest.java
+++ 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsGroupHeartbeatRequestManagerTest.java
@@ -166,7 +166,8 @@ class StreamsGroupHeartbeatRequestManagerTest {
         Optional.of(ENDPOINT),
         Optional.of(RACK_ID),
         SUBTOPOLOGIES,
-        CLIENT_TAGS
+        CLIENT_TAGS,
+        Map::of
     );
 
     private final Time time = new MockTime();
diff --git 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceDataTest.java
 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceDataTest.java
index d383b043e46..92126a224d7 100644
--- 
a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceDataTest.java
+++ 
b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/StreamsRebalanceDataTest.java
@@ -298,7 +298,8 @@ public class StreamsRebalanceDataTest {
             endpoint,
             Optional.empty(),
             subtopologies,
-            clientTags
+            clientTags,
+            Map::of
         );
 
         assertThrows(
@@ -330,7 +331,8 @@ public class StreamsRebalanceDataTest {
                 endpoint,
                 Optional.empty(),
                 subtopologies,
-                clientTags
+                clientTags,
+                Map::of
             )
         );
         assertEquals("Process ID cannot be null", exception.getMessage());
@@ -349,7 +351,8 @@ public class StreamsRebalanceDataTest {
                 null,
                 Optional.empty(),
                 subtopologies,
-                clientTags
+                clientTags,
+                Map::of
             )
         );
         assertEquals("Endpoint cannot be null", exception.getMessage());
@@ -368,7 +371,8 @@ public class StreamsRebalanceDataTest {
                 endpoint,
                 Optional.empty(),
                 null,
-                clientTags
+                clientTags,
+                Map::of
             )
         );
         assertEquals("Subtopologies cannot be null", exception.getMessage());
@@ -388,7 +392,8 @@ public class StreamsRebalanceDataTest {
                 endpoint,
                 null,
                 subtopologies,
-                clientTags
+                clientTags,
+                Map::of
             )
         );
         assertEquals("Rack ID cannot be null", exception.getMessage());
@@ -407,7 +412,8 @@ public class StreamsRebalanceDataTest {
                 endpoint,
                 Optional.empty(),
                 subtopologies,
-                null
+                null,
+                Map::of
             )
         );
         assertEquals("Client tags cannot be null", exception.getMessage());
@@ -424,7 +430,8 @@ public class StreamsRebalanceDataTest {
             endpoint,
             Optional.empty(),
             subtopologies,
-            clientTags
+            clientTags,
+            Map::of
         );
 
         assertEquals(StreamsRebalanceData.Assignment.EMPTY, 
streamsRebalanceData.reconciledAssignment());
@@ -441,7 +448,8 @@ public class StreamsRebalanceDataTest {
             endpoint,
             Optional.empty(),
             subtopologies,
-            clientTags
+            clientTags,
+            Map::of
         );
 
         assertTrue(streamsRebalanceData.partitionsByHost().isEmpty());
@@ -458,7 +466,8 @@ public class StreamsRebalanceDataTest {
             endpoint,
             Optional.empty(),
             subtopologies,
-            clientTags
+            clientTags,
+            Map::of
         );
 
         assertFalse(streamsRebalanceData.shutdownRequested());
@@ -475,7 +484,8 @@ public class StreamsRebalanceDataTest {
             endpoint,
             Optional.empty(),
             subtopologies,
-            clientTags
+            clientTags,
+            Map::of
         );
 
         assertTrue(streamsRebalanceData.statuses().isEmpty());
@@ -490,11 +500,12 @@ public class StreamsRebalanceDataTest {
         final Map<String, String> clientTags = Map.of("clientTag1",
                 "clientTagValue1");
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-                processId,
-                endpoint,
-                Optional.empty(),
-                subtopologies,
-                clientTags
+            processId,
+            endpoint,
+            Optional.empty(),
+            subtopologies,
+            clientTags,
+            Map::of
         );
 
         assertEquals(-1, streamsRebalanceData.heartbeatIntervalMs());
@@ -509,11 +520,12 @@ public class StreamsRebalanceDataTest {
         final Map<String, String> clientTags = Map.of("clientTag1",
                 "clientTagValue1");
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-                processId,
-                endpoint,
-                Optional.empty(),
-                subtopologies,
-                clientTags
+            processId,
+            endpoint,
+            Optional.empty(),
+            subtopologies,
+            clientTags,
+            Map::of
         );
 
         streamsRebalanceData.setHeartbeatIntervalMs(1000);
diff --git 
a/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala 
b/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala
index 92e89d9df62..393985d25c9 100644
--- a/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala
+++ b/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala
@@ -3946,7 +3946,8 @@ class AuthorizerIntegrationTest extends 
AbstractAuthorizerIntegrationTest {
           else util.Map.of(),
           util.Set.of()
         )),
-      Map.empty[String, String].asJava
+      Map.empty[String, String].asJava,
+      () => util.Map.of[StreamsRebalanceData.TaskId, java.lang.Long]()
     ))
     consumer.subscribe(
       if (topicAsSourceTopic || topicAsRepartitionSourceTopic) 
util.Set.of(sourceTopic, topic) else util.Set.of(sourceTopic),
diff --git 
a/core/src/test/scala/integration/kafka/api/IntegrationTestHarness.scala 
b/core/src/test/scala/integration/kafka/api/IntegrationTestHarness.scala
index d8a39f5f317..8c88a7281ac 100644
--- a/core/src/test/scala/integration/kafka/api/IntegrationTestHarness.scala
+++ b/core/src/test/scala/integration/kafka/api/IntegrationTestHarness.scala
@@ -266,7 +266,8 @@ abstract class IntegrationTestHarness extends 
KafkaServerTestHarness {
           changelogTopics.map(c => (c, new 
StreamsRebalanceData.TopicInfo(Optional.empty(), boxed, 
util.Map.of()))).toMap.asJava,
           util.Set.of()
         )),
-      Map.empty[String, String].asJava
+      Map.empty[String, String].asJava,
+      () => util.Map.of[StreamsRebalanceData.TaskId, java.lang.Long]()
     )
 
     val consumer = createStreamsConsumer(
diff --git 
a/streams/src/main/java/org/apache/kafka/streams/processor/internals/StateDirectory.java
 
b/streams/src/main/java/org/apache/kafka/streams/processor/internals/StateDirectory.java
index 403cefed11f..3bb71e68ad4 100644
--- 
a/streams/src/main/java/org/apache/kafka/streams/processor/internals/StateDirectory.java
+++ 
b/streams/src/main/java/org/apache/kafka/streams/processor/internals/StateDirectory.java
@@ -310,6 +310,10 @@ public class StateDirectory implements AutoCloseable {
                 .collect(Collectors.toMap(Map.Entry::getKey, 
Map.Entry::getValue));
     }
 
+    public Map<TaskId, Long> taskOffsetSums() {
+        return Collections.unmodifiableMap(taskOffsetSums);
+    }
+
     public void updateTaskOffsets(final TaskId taskId, final 
Map<TopicPartition, Long> changelogOffsets) {
         if (!changelogOffsets.isEmpty()) {
             taskOffsetSums.put(taskId, sumOfChangelogOffsets(taskId, 
changelogOffsets));
diff --git 
a/streams/src/main/java/org/apache/kafka/streams/processor/internals/StreamThread.java
 
b/streams/src/main/java/org/apache/kafka/streams/processor/internals/StreamThread.java
index c5c7e3b87cb..0433d7fb625 100644
--- 
a/streams/src/main/java/org/apache/kafka/streams/processor/internals/StreamThread.java
+++ 
b/streams/src/main/java/org/apache/kafka/streams/processor/internals/StreamThread.java
@@ -98,6 +98,7 @@ import java.util.concurrent.atomic.AtomicInteger;
 import java.util.concurrent.atomic.AtomicLong;
 import java.util.concurrent.atomic.AtomicReference;
 import java.util.function.BiConsumer;
+import java.util.function.Supplier;
 import java.util.stream.Collectors;
 
 import static 
org.apache.kafka.clients.consumer.CloseOptions.GroupMembershipOperation.DEFAULT;
@@ -513,7 +514,14 @@ public class StreamThread extends Thread implements 
ProcessingThread {
             consumerConfigs.put(ConsumerConfig.AUTO_OFFSET_RESET_CONFIG, 
"none");
         }
 
-        final MainConsumerSetup mainConsumerSetup = 
setupMainConsumer(topologyMetadata, config, clientSupplier, processId, 
consumerConfigs);
+        final MainConsumerSetup mainConsumerSetup = setupMainConsumer(
+            topologyMetadata,
+            config,
+            clientSupplier,
+            processId,
+            consumerConfigs,
+            taskManager::taskOffsetSumSnapshot
+        );
 
         taskManager.setMainConsumer(mainConsumerSetup.mainConsumer);
         referenceContainer.mainConsumer = mainConsumerSetup.mainConsumer;
@@ -555,7 +563,8 @@ public class StreamThread extends Thread implements 
ProcessingThread {
                                                        final StreamsConfig 
config,
                                                        final 
KafkaClientSupplier clientSupplier,
                                                        final UUID processId,
-                                                       final Map<String, 
Object> consumerConfigs) {
+                                                       final Map<String, 
Object> consumerConfigs,
+                                                       final 
Supplier<Map<StreamsRebalanceData.TaskId, Long>> taskOffsetSum) {
         if 
(config.getString(StreamsConfig.GROUP_PROTOCOL_CONFIG).equalsIgnoreCase(GroupProtocol.STREAMS.name))
 {
             if (topologyMetadata.hasNamedTopologies()) {
                 throw new IllegalStateException("Named topologies and the 
STREAMS protocol cannot be used at the same time.");
@@ -566,7 +575,8 @@ public class StreamThread extends Thread implements 
ProcessingThread {
                     config,
                     
parseHostInfo(config.getString(StreamsConfig.APPLICATION_SERVER_CONFIG)),
                     parseRackId((String) 
config.originals().get(CommonClientConfigs.CLIENT_RACK_CONFIG)),
-                    topologyMetadata
+                    topologyMetadata,
+                    taskOffsetSum
                 )
             );
             final ByteArrayDeserializer keyDeserializer = new 
ByteArrayDeserializer();
@@ -690,7 +700,8 @@ public class StreamThread extends Thread implements 
ProcessingThread {
                                                                  final 
StreamsConfig config,
                                                                  final 
Optional<StreamsRebalanceData.HostInfo> endpoint,
                                                                  final 
Optional<String> rackId,
-                                                                 final 
TopologyMetadata topologyMetadata) {
+                                                                 final 
TopologyMetadata topologyMetadata,
+                                                                 final 
Supplier<Map<StreamsRebalanceData.TaskId, Long>> taskOffsetSum) {
         final InternalTopologyBuilder internalTopologyBuilder = 
topologyMetadata.lookupBuilderForNamedTopology(null);
 
         final Map<String, StreamsRebalanceData.Subtopology> subtopologies = 
initBrokerTopology(config, internalTopologyBuilder);
@@ -700,7 +711,8 @@ public class StreamThread extends Thread implements 
ProcessingThread {
             endpoint,
             rackId,
             subtopologies,
-            config.getClientTags()
+            config.getClientTags(),
+            taskOffsetSum
         );
     }
 
@@ -1239,6 +1251,7 @@ public class StreamThread extends Thread implements 
ProcessingThread {
         if (isStartingRunningOrPartitionAssigned()) {
 
             taskManager.updateLags();
+            taskManager.maybeUpdateTaskOffsetSumSnapshot();
 
             /*
              * Within an iteration, after processing up to N (N initialized as 
1 upon start up) records for each applicable tasks, check the current time:
@@ -1386,6 +1399,7 @@ public class StreamThread extends Thread implements 
ProcessingThread {
         if (isRunning()) {
 
             taskManager.updateLags();
+            taskManager.maybeUpdateTaskOffsetSumSnapshot();
 
             checkStateUpdater();
 
diff --git 
a/streams/src/main/java/org/apache/kafka/streams/processor/internals/TaskManager.java
 
b/streams/src/main/java/org/apache/kafka/streams/processor/internals/TaskManager.java
index 0e2723d07af..3d3732eba68 100644
--- 
a/streams/src/main/java/org/apache/kafka/streams/processor/internals/TaskManager.java
+++ 
b/streams/src/main/java/org/apache/kafka/streams/processor/internals/TaskManager.java
@@ -23,6 +23,7 @@ import org.apache.kafka.clients.consumer.Consumer;
 import org.apache.kafka.clients.consumer.ConsumerGroupMetadata;
 import org.apache.kafka.clients.consumer.ConsumerRecords;
 import org.apache.kafka.clients.consumer.OffsetAndMetadata;
+import org.apache.kafka.clients.consumer.internals.StreamsRebalanceData;
 import org.apache.kafka.common.KafkaException;
 import org.apache.kafka.common.Metric;
 import org.apache.kafka.common.MetricName;
@@ -106,6 +107,11 @@ public class TaskManager {
 
     private final Map<TaskId, BackoffRecord> taskIdToBackoffRecord = new 
HashMap<>();
 
+    // Published by the stream thread (via maybeUpdateTaskOffsetSumSnapshot) 
and read lock-free by the
+    // streams-protocol heartbeat thread. Covers standby, warmup, 
restoring-active and dormant tasks; running-active
+    // tasks are omitted (their assignment is not offset-driven, and they are 
caught up by definition).
+    private final AtomicReference<Map<StreamsRebalanceData.TaskId, Long>> 
taskOffsetSumSnapshot = new AtomicReference<>(Map.of());
+
     private final ActiveTaskCreator activeTaskCreator;
     private final StandbyTaskCreator standbyTaskCreator;
     private final StateUpdater stateUpdater;
@@ -1249,6 +1255,44 @@ public class TaskManager {
         }
     }
 
+    /**
+     * Returns the per-task changelog offset-sum snapshot for the 
streams-protocol rebalance. Safe to invoke from any
+     * thread; the snapshot is refreshed by the stream thread via {@link 
#maybeUpdateTaskOffsetSumSnapshot()}.
+     */
+    public Map<StreamsRebalanceData.TaskId, Long> taskOffsetSumSnapshot() {
+        return taskOffsetSumSnapshot.get();
+    }
+
+    /**
+     * Recomputes the offset-sum snapshot reported to the streams-group 
coordinator from the offset sums maintained in
+     * the {@link StateDirectory}, which already cover all stateful tasks with 
state on disk (standby, warmup,
+     * restoring-active and dormant) using a conservative (per-partition 
lower-bound) sum. Running-active tasks are
+     * excluded: their assignment is not offset-driven and they are caught up 
by definition.
+     */
+    public void maybeUpdateTaskOffsetSumSnapshot() {
+        final Set<TaskId> runningActiveTasks = new HashSet<>();
+        for (final Task task : allTasks().values()) {
+            if (task.isActive() && task.state() == State.RUNNING) {
+                runningActiveTasks.add(task.id());
+            }
+        }
+
+        final Map<TaskId, Long> offsetSums = stateDirectory.taskOffsetSums();
+        final Map<StreamsRebalanceData.TaskId, Long> snapshot = new 
HashMap<>(offsetSums.size());
+        for (final Map.Entry<TaskId, Long> entry : offsetSums.entrySet()) {
+            final TaskId taskId = entry.getKey();
+            if (runningActiveTasks.contains(taskId)) {
+                continue;
+            }
+            snapshot.put(
+                new 
StreamsRebalanceData.TaskId(String.valueOf(taskId.subtopology()), 
taskId.partition()),
+                entry.getValue()
+            );
+        }
+
+        taskOffsetSumSnapshot.set(Collections.unmodifiableMap(snapshot));
+    }
+
     /**
      * Compute the offset total summed across all stores in a task. Includes 
offset sum for any tasks we own the
      * lock for, which includes assigned and unassigned tasks we locked in 
{@link #tryToLockAllNonEmptyTaskDirectories()}.
diff --git 
a/streams/src/test/java/org/apache/kafka/streams/processor/internals/DefaultStreamsRebalanceListenerTest.java
 
b/streams/src/test/java/org/apache/kafka/streams/processor/internals/DefaultStreamsRebalanceListenerTest.java
index f94277fa3df..7bd51b7710f 100644
--- 
a/streams/src/test/java/org/apache/kafka/streams/processor/internals/DefaultStreamsRebalanceListenerTest.java
+++ 
b/streams/src/test/java/org/apache/kafka/streams/processor/internals/DefaultStreamsRebalanceListenerTest.java
@@ -106,7 +106,8 @@ public class DefaultStreamsRebalanceListenerTest {
                     Set.of()
                 )
             ),
-            Map.of()
+            Map.of(),
+            Map::of
         ));
         when(streamThread.state()).thenReturn(state);
 
@@ -131,7 +132,7 @@ public class DefaultStreamsRebalanceListenerTest {
         final Exception exception = new RuntimeException("sample exception");
         doThrow(exception).when(taskManager).handleRevocation(any());
 
-        createRebalanceListenerWithRebalanceData(new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of()));
+        createRebalanceListenerWithRebalanceData(new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of));
 
         final Exception actualException = assertThrows(RuntimeException.class, 
() -> defaultStreamsRebalanceListener.onTasksRevoked(Set.of()));
 
@@ -266,7 +267,8 @@ public class DefaultStreamsRebalanceListenerTest {
                     Set.of()
                 )
             ),
-            Map.of()
+            Map.of(),
+            Map::of
         ));
 
         defaultStreamsRebalanceListener.onTasksRevoked(
@@ -301,7 +303,8 @@ public class DefaultStreamsRebalanceListenerTest {
                     Set.of()
                 )
             ),
-            Map.of()
+            Map.of(),
+            Map::of
         ));
 
         defaultStreamsRebalanceListener.onTasksAssigned(
@@ -331,7 +334,7 @@ public class DefaultStreamsRebalanceListenerTest {
             return null;
         }).when(taskManager).handleLostAll();
 
-        createRebalanceListenerWithRebalanceData(new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of()));
+        createRebalanceListenerWithRebalanceData(new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of));
 
         defaultStreamsRebalanceListener.onAllTasksLost();
 
@@ -362,7 +365,8 @@ public class DefaultStreamsRebalanceListenerTest {
                     Set.of()
                 )
             ),
-            Map.of()
+            Map.of(),
+            Map::of
         ));
 
         assertThrows(RuntimeException.class, () -> 
defaultStreamsRebalanceListener.onTasksRevoked(
@@ -382,7 +386,7 @@ public class DefaultStreamsRebalanceListenerTest {
             throw exception;
         }).when(taskManager).handleAssignment(any(), any());
 
-        createRebalanceListenerWithRebalanceData(new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of()));
+        createRebalanceListenerWithRebalanceData(new 
StreamsRebalanceData(UUID.randomUUID(), Optional.empty(), Optional.empty(), 
Map.of(), Map.of(), Map::of));
 
         assertThrows(RuntimeException.class, () -> 
defaultStreamsRebalanceListener.onTasksAssigned(
             new StreamsRebalanceData.Assignment(Set.of(), Set.of(), Set.of(), 
false)
diff --git 
a/streams/src/test/java/org/apache/kafka/streams/processor/internals/StreamThreadTest.java
 
b/streams/src/test/java/org/apache/kafka/streams/processor/internals/StreamThreadTest.java
index 95292c15899..dbff3d6933f 100644
--- 
a/streams/src/test/java/org/apache/kafka/streams/processor/internals/StreamThreadTest.java
+++ 
b/streams/src/test/java/org/apache/kafka/streams/processor/internals/StreamThreadTest.java
@@ -1703,7 +1703,7 @@ public class StreamThreadTest {
         when(mainConsumer.groupMetadata()).thenReturn(consumerGroupMetadata);
         
when(consumerGroupMetadata.groupInstanceId()).thenReturn(Optional.empty());
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-            UUID.randomUUID(), Optional.empty(), Optional.empty(), Map.of(), 
Map.of());
+            UUID.randomUUID(), Optional.empty(), Optional.empty(), Map.of(), 
Map.of(), Map::of);
         thread = new StreamThread(
             mockTime, config, null,
             mainConsumer, consumer,
@@ -1736,7 +1736,7 @@ public class StreamThreadTest {
         when(mainConsumer.groupMetadata()).thenReturn(consumerGroupMetadata);
         
when(consumerGroupMetadata.groupInstanceId()).thenReturn(Optional.empty());
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-            UUID.randomUUID(), Optional.empty(), Optional.empty(), Map.of(), 
Map.of());
+            UUID.randomUUID(), Optional.empty(), Optional.empty(), Map.of(), 
Map.of(), Map::of);
         thread = new StreamThread(
             mockTime, config, null,
             mainConsumer, consumer,
@@ -2981,7 +2981,8 @@ public class StreamThreadTest {
             Optional.empty(),
             Optional.empty(),
             Map.of(),
-            Map.of()
+            Map.of(),
+            Map::of
         );
 
         final StreamsMetricsImpl streamsMetrics =
@@ -3446,6 +3447,7 @@ public class StreamThreadTest {
         final InOrder inOrder = Mockito.inOrder(mainConsumer, 
thread.taskManager());
         inOrder.verify(mainConsumer).poll(Mockito.any());
         inOrder.verify(thread.taskManager()).updateLags();
+        
inOrder.verify(thread.taskManager()).maybeUpdateTaskOffsetSumSnapshot();
     }
 
 
@@ -3932,7 +3934,8 @@ public class StreamThreadTest {
             Optional.empty(),
             Optional.empty(),
             Map.of(),
-            Map.of()
+            Map.of(),
+            Map::of
         );
         final Runnable shutdownErrorHook = mock(Runnable.class);
 
@@ -3993,7 +3996,8 @@ public class StreamThreadTest {
             Optional.empty(),
             Optional.empty(),
             Map.of(),
-            Map.of()
+            Map.of(),
+            Map::of
         );
 
         final Properties props = configProps(false, false);
@@ -4091,11 +4095,12 @@ public class StreamThreadTest {
         when(mainConsumer.poll(Mockito.any(Duration.class))).thenReturn(new 
ConsumerRecords<>(Map.of(), Map.of()));
         when(mainConsumer.groupMetadata()).thenReturn(consumerGroupMetadata);
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-                UUID.randomUUID(),
-                Optional.empty(),
-                Optional.empty(),
-                Map.of(),
-                Map.of()
+            UUID.randomUUID(),
+            Optional.empty(),
+            Optional.empty(),
+            Map.of(),
+            Map.of(),
+            Map::of
         );
         final Runnable shutdownErrorHook = mock(Runnable.class);
 
@@ -4163,11 +4168,12 @@ public class StreamThreadTest {
         when(mainConsumer.poll(Mockito.any(Duration.class))).thenReturn(new 
ConsumerRecords<>(Map.of(), Map.of()));
         when(mainConsumer.groupMetadata()).thenReturn(consumerGroupMetadata);
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-                UUID.randomUUID(),
-                Optional.empty(),
-                Optional.empty(),
-                Map.of(),
-                Map.of()
+            UUID.randomUUID(),
+            Optional.empty(),
+            Optional.empty(),
+            Map.of(),
+            Map.of(),
+            Map::of
         );
         final Runnable shutdownErrorHook = mock(Runnable.class);
 
@@ -4231,7 +4237,8 @@ public class StreamThreadTest {
             Optional.empty(),
             Optional.empty(),
             Map.of(),
-            Map.of()
+            Map.of(),
+            Map::of
         );
 
         final Properties props = configProps(false, false);
@@ -4289,11 +4296,12 @@ public class StreamThreadTest {
         when(mainConsumer.poll(Mockito.any(Duration.class))).thenReturn(new 
ConsumerRecords<>(Map.of(), Map.of()));
         when(mainConsumer.groupMetadata()).thenReturn(consumerGroupMetadata);
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-                UUID.randomUUID(),
-                Optional.empty(),
-                Optional.empty(),
-                Map.of(),
-                Map.of()
+            UUID.randomUUID(),
+            Optional.empty(),
+            Optional.empty(),
+            Map.of(),
+            Map.of(),
+            Map::of
         );
 
         final Properties props = configProps(false, false);
@@ -4361,11 +4369,12 @@ public class StreamThreadTest {
         when(mainConsumer.poll(Mockito.any(Duration.class))).thenReturn(new 
ConsumerRecords<>(Map.of(), Map.of()));
         when(mainConsumer.groupMetadata()).thenReturn(consumerGroupMetadata);
         final StreamsRebalanceData streamsRebalanceData = new 
StreamsRebalanceData(
-                UUID.randomUUID(),
-                Optional.empty(),
-                Optional.empty(),
-                Map.of(),
-                Map.of()
+            UUID.randomUUID(),
+            Optional.empty(),
+            Optional.empty(),
+            Map.of(),
+            Map.of(),
+            Map::of
         );
 
         final Properties props = configProps(false, false);
diff --git 
a/streams/src/test/java/org/apache/kafka/streams/processor/internals/TaskManagerTest.java
 
b/streams/src/test/java/org/apache/kafka/streams/processor/internals/TaskManagerTest.java
index 3a84f779e42..553a84d4a37 100644
--- 
a/streams/src/test/java/org/apache/kafka/streams/processor/internals/TaskManagerTest.java
+++ 
b/streams/src/test/java/org/apache/kafka/streams/processor/internals/TaskManagerTest.java
@@ -24,6 +24,7 @@ import 
org.apache.kafka.clients.consumer.CommitFailedException;
 import org.apache.kafka.clients.consumer.Consumer;
 import org.apache.kafka.clients.consumer.ConsumerGroupMetadata;
 import org.apache.kafka.clients.consumer.OffsetAndMetadata;
+import org.apache.kafka.clients.consumer.internals.StreamsRebalanceData;
 import org.apache.kafka.common.KafkaException;
 import org.apache.kafka.common.KafkaFuture;
 import org.apache.kafka.common.Metric;
@@ -1948,6 +1949,44 @@ public class TaskManagerTest {
         );
     }
 
+    @Test
+    public void shouldReturnEmptyTaskOffsetSumSnapshotBeforeRefresh() {
+        final TasksRegistry tasks = mock(TasksRegistry.class);
+        final TaskManager taskManager = 
setUpTaskManager(ProcessingMode.AT_LEAST_ONCE, tasks);
+
+        assertThat(taskManager.taskOffsetSumSnapshot(), 
is(Collections.emptyMap()));
+    }
+
+    @Test
+    public void 
shouldPublishTaskOffsetSumSnapshotFromStateDirectoryExcludingRunningActiveTasks()
 {
+        final StreamTask runningActiveTask = statefulTask(taskId00, 
taskId00ChangelogPartitions).inState(State.RUNNING).build();
+        final StreamTask restoringActiveTask = statefulTask(taskId01, 
taskId01ChangelogPartitions).inState(State.RESTORING).build();
+        final StandbyTask standbyTask = standbyTask(taskId02, 
taskId02ChangelogPartitions).inState(State.RUNNING).build();
+
+        final TasksRegistry tasks = mock(TasksRegistry.class);
+        final TaskManager taskManager = 
setUpTaskManager(ProcessingMode.AT_LEAST_ONCE, tasks);
+        // running-active tasks are owned by the stream thread; 
restoring-active and standby tasks live in the state updater
+        
when(tasks.allInitializedTasksPerId()).thenReturn(mkMap(mkEntry(taskId00, 
runningActiveTask)));
+        when(stateUpdater.tasks()).thenReturn(Set.of(restoringActiveTask, 
standbyTask));
+        // StateDirectory holds sums for every stateful task with state on 
disk, including a dormant task (taskId03)
+        // that is not currently assigned (not in allTasks()).
+        when(stateDirectory.taskOffsetSums()).thenReturn(mkMap(
+            mkEntry(taskId00, 10L),
+            mkEntry(taskId01, 20L),
+            mkEntry(taskId02, 30L),
+            mkEntry(taskId03, 40L)
+        ));
+
+        taskManager.maybeUpdateTaskOffsetSumSnapshot();
+
+        // running-active taskId00 is omitted; restoring-active, standby, and 
dormant tasks are reported with their sums
+        assertThat(taskManager.taskOffsetSumSnapshot(), is(mkMap(
+            mkEntry(new StreamsRebalanceData.TaskId("0", 1), 20L),
+            mkEntry(new StreamsRebalanceData.TaskId("0", 2), 30L),
+            mkEntry(new StreamsRebalanceData.TaskId("0", 3), 40L)
+        )));
+    }
+
     @Test
     public void shouldSkipUnknownOffsetsWhenComputingOffsetSum() {
         final StreamTask restoringStatefulTask = statefulTask(taskId01, 
taskId01ChangelogPartitions)


Reply via email to