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

1996fanrui pushed a commit to branch release-2.3
in repository https://gitbox.apache.org/repos/asf/flink.git

commit 5e5e73d09663425ec15e4b6eacb6b94c11be1e16
Author: Rui Fan <[email protected]>
AuthorDate: Fri Jul 31 14:47:05 2026 +0200

    [FLINK-40269][runtime] Fix channel state assignment for duplicate 
connections
    
    Co-authored-by: Roman Khachatryan <[email protected]>
---
 .../runtime/checkpoint/TaskStateAssignment.java    | 106 +++++++++++++--------
 .../checkpoint/StateAssignmentOperationTest.java   |  66 +++++++++++++
 2 files changed, 133 insertions(+), 39 deletions(-)

diff --git 
a/flink-runtime/src/main/java/org/apache/flink/runtime/checkpoint/TaskStateAssignment.java
 
b/flink-runtime/src/main/java/org/apache/flink/runtime/checkpoint/TaskStateAssignment.java
index a6db5e4837f..40b58a31101 100644
--- 
a/flink-runtime/src/main/java/org/apache/flink/runtime/checkpoint/TaskStateAssignment.java
+++ 
b/flink-runtime/src/main/java/org/apache/flink/runtime/checkpoint/TaskStateAssignment.java
@@ -53,7 +53,6 @@ import java.util.Map;
 import java.util.Objects;
 import java.util.Optional;
 import java.util.Set;
-import java.util.function.BiFunction;
 import java.util.function.Function;
 import java.util.stream.Collectors;
 import java.util.stream.IntStream;
@@ -230,11 +229,6 @@ class TaskStateAssignment {
         return downstreamAssignments;
     }
 
-    private static int getAssignmentIndex(
-            TaskStateAssignment[] assignments, TaskStateAssignment assignment) 
{
-        return Arrays.asList(assignments).indexOf(assignment);
-    }
-
     public TaskStateAssignment[] getUpstreamAssignments() {
         if (upstreamAssignments == null) {
             upstreamAssignments =
@@ -272,12 +266,6 @@ class TaskStateAssignment {
                                 instanceID,
                                 inputOperatorID,
                                 getUpstreamAssignments(),
-                                (assignment, recompute) -> {
-                                    int assignmentIndex =
-                                            getAssignmentIndex(
-                                                    
assignment.getDownstreamAssignments(), this);
-                                    return 
assignment.getOutputMapping(assignmentIndex, recompute);
-                                },
                                 inputSubtaskMappings,
                                 this::getInputMapping,
                                 true))
@@ -320,11 +308,6 @@ class TaskStateAssignment {
                 instanceID,
                 outputOperatorID,
                 getDownstreamAssignments(),
-                (downstreamAssignment, recompute) -> {
-                    int assignmentIndex =
-                            
getAssignmentIndex(downstreamAssignment.getUpstreamAssignments(), this);
-                    return 
downstreamAssignment.getInputMapping(assignmentIndex, recompute);
-                },
                 outputSubtaskMappings,
                 this::getOutputMapping,
                 false);
@@ -355,7 +338,6 @@ class TaskStateAssignment {
             OperatorInstanceID instanceID,
             OperatorID expectedOperatorID,
             TaskStateAssignment[] connectedAssignments,
-            BiFunction<TaskStateAssignment, Boolean, SubtasksRescaleMapping> 
mappingRetriever,
             Map<Integer, SubtasksRescaleMapping> 
subtaskGateOrPartitionMappings,
             Function<Integer, SubtasksRescaleMapping> subtaskMappingCalculator,
             boolean isInput) {
@@ -364,8 +346,11 @@ class TaskStateAssignment {
         }
 
         SubtasksRescaleMapping[] rescaledChannelsMappings =
-                Arrays.stream(connectedAssignments)
-                        .map(assignment -> mappingRetriever.apply(assignment, 
false))
+                IntStream.range(0, connectedAssignments.length)
+                        .mapToObj(
+                                index ->
+                                        getConnectedMapping(
+                                                isInput, index, 
connectedAssignments[index], false))
                         .toArray(SubtasksRescaleMapping[]::new);
 
         // no state on input and output, especially for any aligned checkpoint
@@ -378,7 +363,6 @@ class TaskStateAssignment {
                 createGateOrPartitionRescalingDescriptors(
                         instanceID,
                         connectedAssignments,
-                        assignment -> mappingRetriever.apply(assignment, true),
                         subtaskGateOrPartitionMappings,
                         subtaskMappingCalculator,
                         rescaledChannelsMappings,
@@ -398,7 +382,6 @@ class TaskStateAssignment {
             createGateOrPartitionRescalingDescriptors(
                     OperatorInstanceID instanceID,
                     TaskStateAssignment[] connectedAssignments,
-                    Function<TaskStateAssignment, SubtasksRescaleMapping> 
mappingCalculator,
                     Map<Integer, SubtasksRescaleMapping> 
subtaskGateOrPartitionMappings,
                     Function<Integer, SubtasksRescaleMapping> 
subtaskMappingCalculator,
                     SubtasksRescaleMapping[] rescaledChannelsMappings,
@@ -415,8 +398,11 @@ class TaskStateAssignment {
                                     
Optional.ofNullable(rescaledChannelsMappings[partition])
                                             .orElseGet(
                                                     () ->
-                                                            
mappingCalculator.apply(
-                                                                    
connectedAssignment));
+                                                            
getConnectedMapping(
+                                                                    isInput,
+                                                                    partition,
+                                                                    
connectedAssignment,
+                                                                    true));
                             SubtasksRescaleMapping subtaskMapping =
                                     Optional.ofNullable(
                                                     
subtaskGateOrPartitionMappings.get(partition))
@@ -485,6 +471,11 @@ class TaskStateAssignment {
         }
     }
 
+    private SubtasksRescaleMapping getOutputMapping(
+            IntermediateDataSetID resultId, boolean recompute) {
+        return getOutputMapping(findResultPartitionIndex(resultId), recompute);
+    }
+
     private SubtasksRescaleMapping getInputMapping(int assignmentIndex, 
boolean recompute) {
         SubtasksRescaleMapping mapping = 
inputSubtaskMappings.get(assignmentIndex);
         if (recompute && mapping == null) {
@@ -494,6 +485,31 @@ class TaskStateAssignment {
         }
     }
 
+    private SubtasksRescaleMapping getInputMapping(
+            IntermediateDataSetID resultId, boolean recompute) {
+        return getInputMapping(findInputGateIndex(resultId), recompute);
+    }
+
+    /**
+     * Resolves the mapping on {@code connectedAssignment} that corresponds to 
{@code index} on
+     * {@code this} assignment, disambiguating by {@link 
IntermediateDataSetID} rather than by array
+     * position (multiple edges can connect the same pair of job vertices).
+     */
+    private SubtasksRescaleMapping getConnectedMapping(
+            boolean isInput,
+            int index,
+            TaskStateAssignment connectedAssignment,
+            boolean recompute) {
+        if (isInput) {
+            IntermediateDataSetID resultId = 
executionJobVertex.getInputs().get(index).getId();
+            return connectedAssignment.getOutputMapping(resultId, recompute);
+        } else {
+            IntermediateDataSetID resultId =
+                    executionJobVertex.getProducedDataSets()[index].getId();
+            return connectedAssignment.getInputMapping(resultId, recompute);
+        }
+    }
+
     public SubtasksRescaleMapping getOutputMapping(int partitionIndex) {
         final TaskStateAssignment downstreamAssignment = 
getDownstreamAssignments()[partitionIndex];
         final IntermediateResult output = 
executionJobVertex.getProducedDataSets()[partitionIndex];
@@ -547,12 +563,8 @@ class TaskStateAssignment {
         if (upstreamAssignment != null && upstreamAssignment.hasOutputState()) 
{
             IntermediateResult inputResult = 
executionJobVertex.getInputs().get(gateIndex);
             IntermediateDataSetID resultId = inputResult.getId();
-            IntermediateResult[] producedDataSets = 
inputResult.getProducer().getProducedDataSets();
-            for (int i = 0; i < producedDataSets.length; i++) {
-                if (producedDataSets[i].getId().equals(resultId)) {
-                    return 
upstreamAssignment.outputStatePartitions.contains(i);
-                }
-            }
+            return upstreamAssignment.outputStatePartitions.contains(
+                    upstreamAssignment.findResultPartitionIndex(resultId));
         }
 
         return false;
@@ -571,12 +583,8 @@ class TaskStateAssignment {
             IntermediateResult producedResult =
                     executionJobVertex.getProducedDataSets()[partitionIndex];
             IntermediateDataSetID resultId = producedResult.getId();
-            List<IntermediateResult> inputs = 
downstreamAssignment.executionJobVertex.getInputs();
-            for (int i = 0; i < inputs.size(); i++) {
-                if (inputs.get(i).getId().equals(resultId)) {
-                    return downstreamAssignment.inputStateGates.contains(i);
-                }
-            }
+            return downstreamAssignment.inputStateGates.contains(
+                    downstreamAssignment.findInputGateIndex(resultId));
         }
         return false;
     }
@@ -642,15 +650,35 @@ class TaskStateAssignment {
 
         IntermediateResult producedResult =
                 executionJobVertex.getProducedDataSets()[partitionIndex];
-        IntermediateDataSetID resultId = producedResult.getId();
-        List<IntermediateResult> inputs = 
downstreamAssignment.executionJobVertex.getInputs();
+        return downstreamAssignment.findInputGateIndex(producedResult.getId());
+    }
+
+    private int findInputGateIndex(IntermediateDataSetID resultId) {
+        List<IntermediateResult> inputs = executionJobVertex.getInputs();
         for (int i = 0; i < inputs.size(); i++) {
             if (inputs.get(i).getId().equals(resultId)) {
                 return i;
             }
         }
         throw new IllegalArgumentException(
-                "No channel rescaler found during rescaling of channel state");
+                "No input gate found for intermediate data set "
+                        + resultId
+                        + " in "
+                        + executionJobVertex.getName());
+    }
+
+    private int findResultPartitionIndex(IntermediateDataSetID resultId) {
+        IntermediateResult[] producedDataSets = 
executionJobVertex.getProducedDataSets();
+        for (int i = 0; i < producedDataSets.length; i++) {
+            if (producedDataSets[i].getId().equals(resultId)) {
+                return i;
+            }
+        }
+        throw new IllegalArgumentException(
+                "No result partition found for intermediate data set "
+                        + resultId
+                        + " in "
+                        + executionJobVertex.getName());
     }
 
     @Override
diff --git 
a/flink-runtime/src/test/java/org/apache/flink/runtime/checkpoint/StateAssignmentOperationTest.java
 
b/flink-runtime/src/test/java/org/apache/flink/runtime/checkpoint/StateAssignmentOperationTest.java
index 711c3ef5cf3..b6e454fb93e 100644
--- 
a/flink-runtime/src/test/java/org/apache/flink/runtime/checkpoint/StateAssignmentOperationTest.java
+++ 
b/flink-runtime/src/test/java/org/apache/flink/runtime/checkpoint/StateAssignmentOperationTest.java
@@ -579,6 +579,72 @@ class StateAssignmentOperationTest {
                                                 RESCALING))));
     }
 
+    @Test
+    void 
testChannelStateAssignmentUsesResultIdForDuplicateJobVertexConnections()
+            throws JobException, JobExecutionException {
+        int oldParallelism = 3;
+        int newParallelism = 2;
+        JobVertex upstream = createJobVertex(new OperatorID(), newParallelism);
+        JobVertex downstream = createJobVertex(new OperatorID(), 
newParallelism);
+        OperatorID upstreamOperator = 
upstream.getOperatorIDs().get(0).getGeneratedOperatorID();
+        OperatorID downstreamOperator = 
downstream.getOperatorIDs().get(0).getGeneratedOperatorID();
+        Random random = new Random();
+
+        OperatorState upstreamState =
+                new OperatorState("", "", upstreamOperator, oldParallelism, 
MAX_P);
+        OperatorState downstreamState =
+                new OperatorState("", "", downstreamOperator, oldParallelism, 
MAX_P);
+        for (int i = 0; i < oldParallelism; i++) {
+            upstreamState.putState(
+                    i,
+                    OperatorSubtaskState.builder()
+                            .setResultSubpartitionState(
+                                    new StateObjectCollection<>(
+                                            asList(
+                                                    
createNewResultSubpartitionStateHandle(
+                                                            10, 0, random),
+                                                    
createNewResultSubpartitionStateHandle(
+                                                            10, 1, random))))
+                            .build());
+            downstreamState.putState(
+                    i,
+                    OperatorSubtaskState.builder()
+                            .setInputChannelState(
+                                    new StateObjectCollection<>(
+                                            asList(
+                                                    
createNewInputChannelStateHandle(10, 0, random),
+                                                    
createNewInputChannelStateHandle(
+                                                            10, 1, random))))
+                            .build());
+        }
+        Map<OperatorID, OperatorState> states = new HashMap<>();
+        states.put(upstreamOperator, upstreamState);
+        states.put(downstreamOperator, downstreamState);
+
+        connectVertices(upstream, downstream, RANGE, RANGE);
+        connectVertices(upstream, downstream, ROUND_ROBIN, ROUND_ROBIN);
+
+        Map<OperatorID, ExecutionJobVertex> vertices = 
toExecutionVertices(upstream, downstream);
+
+        new StateAssignmentOperation(0, new HashSet<>(vertices.values()), 
states, false, false)
+                .assignStates();
+
+        InflightDataRescalingDescriptor outputDescriptor =
+                getAssignedState(vertices.get(upstreamOperator), 
upstreamOperator, 0)
+                        .getOutputRescalingDescriptor();
+        InflightDataRescalingDescriptor inputDescriptor =
+                getAssignedState(vertices.get(downstreamOperator), 
downstreamOperator, 0)
+                        .getInputRescalingDescriptor();
+        assertThat(outputDescriptor.getChannelMapping(0))
+                .isEqualTo(RANGE.getNewToOldSubtasksMapping(oldParallelism, 
newParallelism));
+        assertThat(outputDescriptor.getChannelMapping(1))
+                
.isEqualTo(ROUND_ROBIN.getNewToOldSubtasksMapping(oldParallelism, 
newParallelism));
+        assertThat(inputDescriptor.getChannelMapping(0))
+                .isEqualTo(RANGE.getNewToOldSubtasksMapping(oldParallelism, 
newParallelism));
+        assertThat(inputDescriptor.getChannelMapping(1))
+                
.isEqualTo(ROUND_ROBIN.getNewToOldSubtasksMapping(oldParallelism, 
newParallelism));
+    }
+
     private InflightDataGateOrPartitionRescalingDescriptor gate(
             int[] oldIndices,
             RescaleMappings rescaleMapping,

Reply via email to