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 {

Reply via email to