This is an automated email from the ASF dual-hosted git repository.
github-merge-queue[bot] pushed a commit to branch dev
in repository https://gitbox.apache.org/repos/asf/seatunnel.git
The following commit(s) were added to refs/heads/dev by this push:
new ae9132fe07 [Fix][Zeta] Checkpoint restore remap duplicates subtask
state for parallelism > 1 (#12195)
ae9132fe07 is described below
commit ae9132fe0767eee612259c293dc28b73a48dfae4
Author: Daniel <[email protected]>
AuthorDate: Mon Sep 14 14:37:01 2026 +0000
[Fix][Zeta] Checkpoint restore remap duplicates subtask state for
parallelism > 1 (#12195)
Co-authored-by: Claude Fable 5.1 <[email protected]>
---
.../server/checkpoint/CheckpointCoordinator.java | 49 ++++-
.../checkpoint/CheckpointCoordinatorTest.java | 198 +++++++++++++++++++++
2 files changed, 246 insertions(+), 1 deletion(-)
diff --git
a/seatunnel-engine/seatunnel-engine-server/src/main/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinator.java
b/seatunnel-engine/seatunnel-engine-server/src/main/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinator.java
index 0fd4cf039d..d4dfa45436 100644
---
a/seatunnel-engine/seatunnel-engine-server/src/main/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinator.java
+++
b/seatunnel-engine/seatunnel-engine-server/src/main/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinator.java
@@ -112,6 +112,18 @@ public class CheckpointCoordinator {
*/
private final Map<Long, Integer> pipelineTasks;
+ /**
+ * Current parallelism of every action in this pipeline, keyed by {@link
ActionStateKey}.
+ *
+ * <p>Used exclusively by the checkpoint-state remap in {@link
#restoreTaskState(TaskLocation)}.
+ * It is intentionally kept separate from {@link #pipelineTasks}: that map
is keyed by {@link
+ * TaskLocation#getTaskVertexId()}, which encodes the subtask's own
parallelism index and is
+ * therefore unique per subtask (see the {@link TaskLocation}
constructor), so grouping by it
+ * can never recover how many subtasks currently run a given action. See
{@link
+ * #getActionParallelism(Map)} for how this is derived instead.
+ */
+ private final Map<ActionStateKey, Integer> actionParallelism;
+
private final Map<Long, SeaTunnelTaskState> pipelineTaskStatus;
private final CheckpointPlan plan;
@@ -238,6 +250,7 @@ public class CheckpointCoordinator {
this.scheduler = MDCTracer.tracing(scheduler);
this.serializer = new ProtoStuffSerializer();
this.pipelineTasks = getPipelineTasks(plan.getPipelineSubtasks());
+ this.actionParallelism =
getActionParallelism(plan.getSubtaskActions());
this.pipelineTaskStatus = new ConcurrentHashMap<>();
this.checkpointIdCounter = checkpointIdCounter;
this.readyToCloseStartingTask = new CopyOnWriteArraySet<>();
@@ -378,7 +391,6 @@ public class CheckpointCoordinator {
if (!latestCompletedCheckpoint.isRestored()) {
latestCompletedCheckpoint.setRestored(true);
}
- final Integer currentParallelism =
pipelineTasks.get(taskLocation.getTaskVertexId());
plan.getSubtaskActions()
.get(taskLocation)
.forEach(
@@ -396,6 +408,12 @@ public class CheckpointCoordinator {
states.add(actionState.getCoordinatorState());
return;
}
+ // The remap step must be the CURRENT
parallelism of this
+ // specific action, not of the task/vertex: a
single subtask can
+ // carry several chained actions, and
pipelineTasks is keyed by
+ // TaskLocation#getTaskVertexId(), which is
unique per subtask (see
+ // actionParallelism's javadoc), so it can
never be reused here.
+ final int currentParallelism =
actionParallelism.get(tuple.f0());
for (int i = tuple.f1();
i < actionState.getParallelism();
i += currentParallelism) {
@@ -829,6 +847,35 @@ public class CheckpointCoordinator {
.collect(Collectors.toMap(Map.Entry::getKey, entry ->
entry.getValue().size()));
}
+ /**
+ * Derives the current parallelism of every action from the
subtask-to-action mapping of the
+ * checkpoint plan.
+ *
+ * <p>{@code
+ *
org.apache.seatunnel.engine.server.dag.physical.PhysicalPlanGenerator#fillCheckpointPlan}
+ * (and the enumerator/committer task wiring around it) records, for every
subtask, one {@code
+ * (ActionStateKey, index)} tuple per action that subtask participates in,
where {@code index}
+ * is that subtask's own parallelism index (or {@link
CheckpointPlan#COORDINATOR_INDEX} for
+ * coordinator-only tasks such as the split enumerator). Counting, per
action, how many tuples
+ * carry a real (non-coordinator) index therefore yields exactly the
number of subtasks
+ * currently running that action, i.e. its current parallelism.
+ *
+ * <p>This must NOT be approximated via {@link #getPipelineTasks(Set)}:
that map is keyed by
+ * {@link TaskLocation#getTaskVertexId()}, which folds in the subtask's
own parallelism index
+ * and is therefore unique per subtask (see the {@link TaskLocation}
constructor), so grouping
+ * by it always yields a group size of 1 regardless of the action's real
parallelism.
+ *
+ * @param subtaskActions the subtask-to-action mapping of the checkpoint
plan
+ * @return the current parallelism of every action key found in {@code
subtaskActions}
+ */
+ public static Map<ActionStateKey, Integer> getActionParallelism(
+ Map<TaskLocation, Set<Tuple2<ActionStateKey, Integer>>>
subtaskActions) {
+ return subtaskActions.values().stream()
+ .flatMap(Set::stream)
+ .filter(tuple -> !COORDINATOR_INDEX.equals(tuple.f1()))
+ .collect(Collectors.groupingBy(Tuple2::f0,
Collectors.summingInt(tuple -> 1)));
+ }
+
@SneakyThrows
public PassiveCompletableFuture<CompletedCheckpoint> startSavepoint() {
LOG.info("start save point for job id: {}.", jobId);
diff --git
a/seatunnel-engine/seatunnel-engine-server/src/test/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinatorTest.java
b/seatunnel-engine/seatunnel-engine-server/src/test/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinatorTest.java
index cf33840e91..b6a2e396f9 100644
---
a/seatunnel-engine/seatunnel-engine-server/src/test/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinatorTest.java
+++
b/seatunnel-engine/seatunnel-engine-server/src/test/java/org/apache/seatunnel/engine/server/checkpoint/CheckpointCoordinatorTest.java
@@ -18,6 +18,7 @@
package org.apache.seatunnel.engine.server.checkpoint;
import org.apache.seatunnel.common.utils.ReflectionUtils;
+import org.apache.seatunnel.engine.checkpoint.storage.PipelineState;
import org.apache.seatunnel.engine.checkpoint.storage.api.CheckpointStorage;
import org.apache.seatunnel.engine.common.config.server.CheckpointConfig;
import
org.apache.seatunnel.engine.common.config.server.CheckpointStorageConfig;
@@ -25,8 +26,10 @@ import
org.apache.seatunnel.engine.common.utils.concurrent.CompletableFuture;
import org.apache.seatunnel.engine.core.checkpoint.CheckpointIDCounter;
import org.apache.seatunnel.engine.core.checkpoint.CheckpointType;
import org.apache.seatunnel.engine.core.job.RestoreMode;
+import org.apache.seatunnel.engine.serializer.protobuf.ProtoStuffSerializer;
import org.apache.seatunnel.engine.server.AbstractSeaTunnelServerTest;
import
org.apache.seatunnel.engine.server.checkpoint.monitor.CheckpointMonitorService;
+import
org.apache.seatunnel.engine.server.checkpoint.operation.NotifyTaskRestoreOperation;
import
org.apache.seatunnel.engine.server.checkpoint.operation.TaskAcknowledgeOperation;
import org.apache.seatunnel.engine.server.common.SeaTunnelEngineContext;
import org.apache.seatunnel.engine.server.execution.TaskGroupLocation;
@@ -37,6 +40,7 @@ import
org.apache.seatunnel.engine.server.task.statemachine.SeaTunnelTaskState;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Test;
+import org.mockito.ArgumentCaptor;
import org.mockito.MockedStatic;
import org.mockito.Mockito;
@@ -45,6 +49,7 @@ import com.hazelcast.map.IMap;
import com.hazelcast.spi.impl.NodeEngine;
import com.hazelcast.spi.impl.operationservice.impl.InvocationFuture;
+import java.nio.charset.StandardCharsets;
import java.time.Instant;
import java.util.ArrayList;
import java.util.Arrays;
@@ -65,6 +70,7 @@ import java.util.concurrent.TimeoutException;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicLong;
+import java.util.stream.Collectors;
import static
org.apache.seatunnel.engine.common.Constant.IMAP_RUNNING_JOB_STATE;
@@ -1176,6 +1182,198 @@ public class CheckpointCoordinatorTest
executorService.shutdownNow();
}
}
+
+ /**
+ * Regression for the per-subtask state remap in {@code
CheckpointCoordinator#restoreTaskState}:
+ * for every old/new parallelism pair in 1..4 x 1..4, every subtask state
recorded in the
+ * checkpoint must be handed to exactly one subtask of the restored plan
(no duplicate, no
+ * drop), a subtask must only receive checkpointed indexes congruent to
its own index modulo the
+ * new parallelism, and the coordinator task must receive exactly the
coordinator state.
+ *
+ * <p>Before the fix the remap step was looked up through {@code
+ * TaskLocation#getTaskVertexId()}, which is unique per subtask, so the
step degenerated to 1
+ * and subtask j received every checkpointed state with index >= j,
duplicating splits after any
+ * restore with parallelism > 1.
+ */
+ @Test
+ void
testRestoreTaskStateDeliversEveryCheckpointedSubtaskStateExactlyOnce() {
+ ExecutorService executorService = Executors.newCachedThreadPool();
+ try {
+ for (int oldParallelism = 1; oldParallelism <= 4;
oldParallelism++) {
+ for (int newParallelism = 1; newParallelism <= 4;
newParallelism++) {
+ assertRestoreRemap(executorService, oldParallelism,
newParallelism);
+ }
+ }
+ } finally {
+ executorService.shutdownNow();
+ }
+ }
+
+ /**
+ * Restores a checkpoint taken at {@code oldParallelism} into a plan with
{@code newParallelism}
+ * subtasks through the real {@code restoreTaskState} path and asserts the
remap contract on the
+ * {@link NotifyTaskRestoreOperation}s sent to the member nodes.
+ */
+ private void assertRestoreRemap(
+ ExecutorService executorService, int oldParallelism, int
newParallelism) {
+ String scenario = "oldParallelism=" + oldParallelism + ",
newParallelism=" + newParallelism;
+ ActionStateKey actionKey = new ActionStateKey("ActionStateKey -
remap-source");
+
+ // Checkpoint taken at the old parallelism: one state per old subtask
plus the coordinator
+ // (split enumerator) state, indexed exactly as
PendingCheckpoint#acknowledgeTask records
+ // them (subtask index -> subtaskStates, COORDINATOR_INDEX ->
coordinatorState).
+ ActionState actionState = new ActionState(actionKey, oldParallelism);
+ actionState.reportState(
+ CheckpointPlan.COORDINATOR_INDEX,
+ new ActionSubtaskState(
+ actionKey,
+ CheckpointPlan.COORDINATOR_INDEX,
+
Collections.singletonList("coordinator".getBytes(StandardCharsets.UTF_8))));
+ for (int index = 0; index < oldParallelism; index++) {
+ actionState.reportState(
+ index,
+ new ActionSubtaskState(
+ actionKey,
+ index,
+ Collections.singletonList(
+
String.valueOf(index).getBytes(StandardCharsets.UTF_8))));
+ }
+ Map<ActionStateKey, ActionState> taskStates = new HashMap<>();
+ taskStates.put(actionKey, actionState);
+ long now = System.currentTimeMillis();
+ CompletedCheckpoint completedCheckpoint =
+ new CompletedCheckpoint(
+ 1L,
+ 1,
+ 1L,
+ now,
+ CheckpointType.SAVEPOINT_TYPE,
+ now,
+ taskStates,
+ new HashMap<>());
+ PipelineState pipelineState =
+ PipelineState.builder()
+ .jobId("1")
+ .pipelineId(1)
+ .checkpointId(1L)
+ .states(new
ProtoStuffSerializer().serialize(completedCheckpoint))
+ .build();
+
+ // Plan of the restored job at the new parallelism, shaped like the
output of
+ // PhysicalPlanGenerator: the coordinator task is registered with
COORDINATOR_INDEX and
+ // every parallelism index gets its own task group (hence its own task
id) registered
+ // with that index.
+ TaskLocation coordinatorTask = new TaskLocation(new
TaskGroupLocation(1L, 1, 1), 0, 0);
+ Map<TaskLocation, Set<Tuple2<ActionStateKey, Integer>>> subtaskActions
= new HashMap<>();
+ subtaskActions.put(
+ coordinatorTask,
+ Collections.singleton(Tuple2.tuple2(actionKey,
CheckpointPlan.COORDINATOR_INDEX)));
+ List<TaskLocation> subtasks = new ArrayList<>();
+ for (int index = 0; index < newParallelism; index++) {
+ TaskLocation subtask =
+ new TaskLocation(new TaskGroupLocation(1L, 1, 2 + index),
0, index);
+ subtasks.add(subtask);
+ subtaskActions.put(subtask,
Collections.singleton(Tuple2.tuple2(actionKey, index)));
+ }
+ Map<ActionStateKey, Integer> pipelineActions = new HashMap<>();
+ pipelineActions.put(actionKey, newParallelism);
+ Set<TaskLocation> pipelineSubtasks = new HashSet<>(subtasks);
+ pipelineSubtasks.add(coordinatorTask);
+ CheckpointPlan plan =
+ CheckpointPlan.builder()
+ .pipelineId(1)
+ .pipelineSubtasks(pipelineSubtasks)
+
.startingSubtasks(Collections.singleton(coordinatorTask))
+ .pipelineActions(pipelineActions)
+ .subtaskActions(subtaskActions)
+ .build();
+
+ int derivedParallelism =
+
CheckpointCoordinator.getActionParallelism(plan.getSubtaskActions()).get(actionKey);
+ Assertions.assertEquals(
+ newParallelism,
+ derivedParallelism,
+ scenario + ": the remap step must be the per-action
parallelism");
+
+ CheckpointConfig checkpointConfig = new CheckpointConfig();
+ checkpointConfig.setStorage(new CheckpointStorageConfig());
+ CheckpointManager mockManager = Mockito.mock(CheckpointManager.class);
+ Mockito.doReturn(Mockito.mock(InvocationFuture.class))
+ .when(mockManager)
+ .sendOperationToMemberNode(Mockito.any(TaskOperation.class));
+ @SuppressWarnings("unchecked")
+ IMap<Object, Object> mockIMap = Mockito.mock(IMap.class);
+ CheckpointCoordinator coordinator =
+ new CheckpointCoordinator(
+ mockManager,
+ Mockito.mock(CheckpointStorage.class),
+ checkpointConfig,
+ 1L,
+ plan,
+ Mockito.mock(CheckpointIDCounter.class),
+ pipelineState,
+ executorService,
+ mockIMap,
+ true,
+ null);
+
+ ReflectionUtils.invoke(coordinator, "restoreTaskState",
coordinatorTask);
+ for (TaskLocation subtask : subtasks) {
+ ReflectionUtils.invoke(coordinator, "restoreTaskState", subtask);
+ }
+
+ ArgumentCaptor<TaskOperation> captor =
ArgumentCaptor.forClass(TaskOperation.class);
+ Mockito.verify(mockManager, Mockito.times(newParallelism + 1))
+ .sendOperationToMemberNode(captor.capture());
+ Map<TaskLocation, List<Integer>> restoredIndexes = new HashMap<>();
+ for (TaskOperation operation : captor.getAllValues()) {
+ Assertions.assertInstanceOf(NotifyTaskRestoreOperation.class,
operation, scenario);
+ @SuppressWarnings("unchecked")
+ List<ActionSubtaskState> restored =
+ (List<ActionSubtaskState>)
+ ReflectionUtils.getField(operation,
"restoredState")
+ .orElseThrow(
+ () ->
+ new IllegalStateException(
+ "restoredState
field not found"));
+ restoredIndexes.put(
+ operation.getTaskLocation(),
+ restored.stream()
+ .map(ActionSubtaskState::getIndex)
+ .collect(Collectors.toList()));
+ }
+
+ Assertions.assertEquals(
+ Collections.singletonList(CheckpointPlan.COORDINATOR_INDEX),
+ restoredIndexes.get(coordinatorTask),
+ scenario + ": the coordinator task must receive exactly the
coordinator state");
+
+ List<Integer> delivered = new ArrayList<>();
+ for (TaskLocation subtask : subtasks) {
+ List<Integer> indexes = restoredIndexes.get(subtask);
+ Assertions.assertNotNull(indexes, scenario + ": subtask " +
subtask.getTaskIndex());
+ for (Integer index : indexes) {
+ Assertions.assertEquals(
+ subtask.getTaskIndex(),
+ index % newParallelism,
+ scenario
+ + ": subtask "
+ + subtask.getTaskIndex()
+ + " received checkpointed index "
+ + index);
+ }
+ delivered.addAll(indexes);
+ }
+ Collections.sort(delivered);
+ List<Integer> expected = new ArrayList<>();
+ for (int index = 0; index < oldParallelism; index++) {
+ expected.add(index);
+ }
+ Assertions.assertEquals(
+ expected,
+ delivered,
+ scenario + ": every checkpointed subtask state must be
restored exactly once");
+ }
}
class TestCheckpointManager extends CheckpointManager {