This is an automated email from the ASF dual-hosted git repository. 1996fanrui pushed a commit to branch release-2.0 in repository https://gitbox.apache.org/repos/asf/flink.git
commit 654354dcdaba9eafca1948c8d8a5baff90769e0c 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 | 109 ++++++++++++++------- .../checkpoint/StateAssignmentOperationTest.java | 66 +++++++++++++ 2 files changed, 138 insertions(+), 37 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 9a314918f4c..bd2c9f5689c 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 @@ -47,7 +47,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; @@ -173,11 +172,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 = @@ -210,12 +204,6 @@ class TaskStateAssignment { instanceID, inputOperatorID, getUpstreamAssignments(), - (assignment, recompute) -> { - int assignmentIndex = - getAssignmentIndex( - assignment.getDownstreamAssignments(), this); - return assignment.getOutputMapping(assignmentIndex, recompute); - }, inputSubtaskMappings, this::getInputMapping, true)) @@ -224,12 +212,6 @@ class TaskStateAssignment { instanceID, outputOperatorID, getDownstreamAssignments(), - (assignment, recompute) -> { - int assignmentIndex = - getAssignmentIndex( - assignment.getUpstreamAssignments(), this); - return assignment.getInputMapping(assignmentIndex, recompute); - }, outputSubtaskMappings, this::getOutputMapping, false)) @@ -279,7 +261,6 @@ class TaskStateAssignment { OperatorInstanceID instanceID, OperatorID expectedOperatorID, TaskStateAssignment[] connectedAssignments, - BiFunction<TaskStateAssignment, Boolean, SubtasksRescaleMapping> mappingRetriever, Map<Integer, SubtasksRescaleMapping> subtaskGateOrPartitionMappings, Function<Integer, SubtasksRescaleMapping> subtaskMappingCalculator, boolean isInput) { @@ -288,8 +269,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 @@ -302,7 +286,6 @@ class TaskStateAssignment { createGateOrPartitionRescalingDescriptors( instanceID, connectedAssignments, - assignment -> mappingRetriever.apply(assignment, true), subtaskGateOrPartitionMappings, subtaskMappingCalculator, rescaledChannelsMappings, @@ -322,7 +305,6 @@ class TaskStateAssignment { createGateOrPartitionRescalingDescriptors( OperatorInstanceID instanceID, TaskStateAssignment[] connectedAssignments, - Function<TaskStateAssignment, SubtasksRescaleMapping> mappingCalculator, Map<Integer, SubtasksRescaleMapping> subtaskGateOrPartitionMappings, Function<Integer, SubtasksRescaleMapping> subtaskMappingCalculator, SubtasksRescaleMapping[] rescaledChannelsMappings, @@ -339,8 +321,11 @@ class TaskStateAssignment { Optional.ofNullable(rescaledChannelsMappings[partition]) .orElseGet( () -> - mappingCalculator.apply( - connectedAssignment)); + getConnectedMapping( + isInput, + partition, + connectedAssignment, + true)); SubtasksRescaleMapping subtaskMapping = Optional.ofNullable( subtaskGateOrPartitionMappings.get(partition)) @@ -409,6 +394,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) { @@ -418,6 +408,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]; @@ -471,12 +486,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; @@ -495,16 +506,40 @@ 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; } + 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 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 public String toString() { return "TaskStateAssignment for " + executionJobVertex.getName(); 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 fc393090d5b..0efcfbfec9d 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 @@ -554,6 +554,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) + .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,
