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)