diff --git a/flink-tests/src/test/java/org/apache/flink/test/checkpointing/FileMergingChannelStateITCase.java b/flink-tests/src/test/java/org/apache/flink/test/checkpointing/FileMergingChannelStateITCase.java index 234e452d5ea82b..75d8fa8235a6f7 100644 --- a/flink-tests/src/test/java/org/apache/flink/test/checkpointing/FileMergingChannelStateITCase.java +++ b/flink-tests/src/test/java/org/apache/flink/test/checkpointing/FileMergingChannelStateITCase.java @@ -19,6 +19,7 @@ package org.apache.flink.test.checkpointing; import org.apache.flink.api.common.JobExecutionResult; +import org.apache.flink.api.common.JobID; import org.apache.flink.api.common.functions.RichMapFunction; import org.apache.flink.api.common.state.ListState; import org.apache.flink.api.common.state.ListStateDescriptor; @@ -30,9 +31,14 @@ import org.apache.flink.configuration.StateRecoveryOptions; import org.apache.flink.core.execution.CheckpointingMode; import org.apache.flink.core.execution.JobClient; +import org.apache.flink.runtime.checkpoint.AbstractCheckpointStats; +import org.apache.flink.runtime.checkpoint.CheckpointStatsSnapshot; +import org.apache.flink.runtime.checkpoint.CompletedCheckpointStats; import org.apache.flink.runtime.checkpoint.OperatorState; import org.apache.flink.runtime.checkpoint.OperatorSubtaskState; import org.apache.flink.runtime.checkpoint.metadata.CheckpointMetadata; +import org.apache.flink.runtime.executiongraph.AccessExecutionGraph; +import org.apache.flink.runtime.executiongraph.ErrorInfo; import org.apache.flink.runtime.minicluster.MiniCluster; import org.apache.flink.runtime.state.FunctionInitializationContext; import org.apache.flink.runtime.state.FunctionSnapshotContext; @@ -54,16 +60,23 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.RegisterExtension; import org.junit.jupiter.api.io.TempDir; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import javax.annotation.Nullable; +import java.io.IOException; import java.time.Duration; import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; +import java.util.HashSet; import java.util.List; +import java.util.Set; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicLongArray; +import java.util.concurrent.atomic.AtomicReference; import java.util.stream.Collectors; import java.util.stream.Stream; @@ -73,9 +86,10 @@ /** Tests recovery of file-merged channel state after a job is restarted from a checkpoint. */ class FileMergingChannelStateITCase { + private static final Logger LOG = LoggerFactory.getLogger(FileMergingChannelStateITCase.class); + private static final int TASK_MANAGER_COUNT = 3; private static final int WORD_COUNT = 16; - private static final int INITIAL_CHECKPOINTS_TO_WAIT = 2; private static final long RECORD_COUNT = 160_000L; private static final long EXPECTED_COUNT_PER_WORD = RECORD_COUNT / WORD_COUNT; private static final String SLOW_MAPPER_UID = "slow-word-mapper"; @@ -135,11 +149,9 @@ void testRestoreFileMergedChannelState(@InjectMiniCluster MiniCluster miniCluste try { CommonTestUtils.waitForAllTaskRunning(miniCluster, initialJobClient.getJobID(), true); - // The first periodic checkpoint can start before the slow mapper has accumulated input - // channel state. checkpointPath = - CommonTestUtils.waitForCheckpointWithInflightBuffers( - initialJobClient.getJobID(), miniCluster, INITIAL_CHECKPOINTS_TO_WAIT); + waitForCheckpointWithSlowMapperChannelState( + initialJobClient.getJobID(), miniCluster); assertFileMergedChannelState(TestUtils.loadCheckpointMetadata(checkpointPath)); } finally { try { @@ -216,37 +228,135 @@ private StreamExecutionEnvironment createEnvironment( return env; } + /** + * Returns the path of a completed checkpoint that carries in-flight input channel state for the + * slow mapper. + * + *

Whether a particular checkpoint contains in-flight data for a particular subtask depends + * on where the barriers happen to be when the checkpoint is triggered, so waiting for a fixed + * number of checkpoints - or for the latest checkpoint that persisted any in-flight + * data anywhere in the job - also accepts checkpoints that do not exercise channel state + * recovery for the mapper at all. Inspect the metadata of every completed checkpoint instead + * and return the first one that really contains the state this test is about. + */ + private static String waitForCheckpointWithSlowMapperChannelState( + JobID jobID, MiniCluster miniCluster) throws Exception { + final Set inspectedCheckpoints = new HashSet<>(); + final AtomicReference restorePath = new AtomicReference<>(); + CommonTestUtils.waitUntilCondition( + () -> { + final AccessExecutionGraph graph = miniCluster.getExecutionGraph(jobID).get(); + for (CompletedCheckpointStats checkpoint : + checkpointsNotInspectedYet(graph, inspectedCheckpoints)) { + if (carriesSlowMapperChannelState(checkpoint)) { + restorePath.set(checkpoint.getExternalPath()); + return true; + } + } + failIfJobStoppedCheckpointing(graph, inspectedCheckpoints); + return false; + }); + return restorePath.get(); + } + + /** + * Returns the retained checkpoints that persisted in-flight data and have not been looked at by + * an earlier call, oldest first: the earliest usable checkpoint is the one that leaves the most + * records for the restored job to replay. + */ + private static List checkpointsNotInspectedYet( + AccessExecutionGraph graph, Set inspectedCheckpoints) { + final CheckpointStatsSnapshot snapshot = graph.getCheckpointStatsSnapshot(); + if (snapshot == null) { + return Collections.emptyList(); + } + // The history is ordered from the newest to the oldest checkpoint. + final List history = + new ArrayList<>(snapshot.getHistory().getCheckpoints()); + Collections.reverse(history); + return history.stream() + .filter(CompletedCheckpointStats.class::isInstance) + .map(CompletedCheckpointStats.class::cast) + .filter(checkpoint -> checkpoint.getPersistedData() > 0L) + .filter(checkpoint -> checkpoint.getExternalPath() != null) + .filter(checkpoint -> inspectedCheckpoints.add(checkpoint.getCheckpointId())) + .collect(Collectors.toList()); + } + + /** + * Returns whether restoring from the given checkpoint would exercise file-merged channel state + * recovery, i.e. whether it holds in-flight input channel state for the slow mapper. + */ + private static boolean carriesSlowMapperChannelState(CompletedCheckpointStats checkpoint) { + try { + final CheckpointMetadata metadata = + TestUtils.loadCheckpointMetadata(checkpoint.getExternalPath()); + return !collectChannelStateDelegates(metadata).slowMapperInputChannelState.isEmpty(); + } catch (IOException e) { + // The checkpoint was subsumed and cleaned up while it was being inspected. + LOG.debug("Skipping checkpoint {}.", checkpoint.getExternalPath(), e); + return false; + } + } + + /** + * Stops the wait with the job's own failure cause once the job has reached a terminal state, as + * no further checkpoint can complete from then on. + */ + private static void failIfJobStoppedCheckpointing( + AccessExecutionGraph graph, Set inspectedCheckpoints) { + if (!graph.getState().isGloballyTerminalState()) { + return; + } + final ErrorInfo failureInfo = graph.getFailureInfo(); + throw new IllegalStateException( + String.format( + "Job reached the terminal state %s before completing a checkpoint with " + + "in-flight input channel state for %s. Inspected checkpoints: %s.", + graph.getState(), SLOW_MAPPER_UID, inspectedCheckpoints), + failureInfo == null ? null : failureInfo.getException()); + } + private static void assertFileMergedChannelState(CheckpointMetadata metadata) { - final List channelStateDelegates = new ArrayList<>(); - final List slowMapperChannelStateDelegates = new ArrayList<>(); + final ChannelStateDelegates delegates = collectChannelStateDelegates(metadata); + + assertThat(delegates.all) + .as("channel state delegates in the checkpoint") + .isNotEmpty() + .allSatisfy( + handle -> assertThat(handle).isInstanceOf(SegmentFileStateHandle.class)); + assertThat(delegates.all.stream().mapToLong(StreamStateHandle::getStateSize).sum()) + .isPositive(); + assertThat(delegates.slowMapperInputChannelState) + .as("channel state delegates belonging to the stateless slow mapper") + .isNotEmpty(); + } + + private static ChannelStateDelegates collectChannelStateDelegates(CheckpointMetadata metadata) { + final ChannelStateDelegates delegates = new ChannelStateDelegates(); for (OperatorState operatorState : metadata.getOperatorStates()) { for (OperatorSubtaskState subtaskState : operatorState.getStates()) { - final List subtaskChannelStateDelegates = - collectUniqueDisposableInChannelState( - Stream.of( - subtaskState.getInputChannelState(), - subtaskState.getUpstreamOutputBufferState(), - subtaskState.getResultSubpartitionState())) - .collect(Collectors.toList()); - channelStateDelegates.addAll(subtaskChannelStateDelegates); + collectUniqueDisposableInChannelState( + Stream.of( + subtaskState.getInputChannelState(), + subtaskState.getUpstreamOutputBufferState(), + subtaskState.getResultSubpartitionState())) + .forEach(delegates.all::add); if (operatorState.getOperatorUid().filter(SLOW_MAPPER_UID::equals).isPresent()) { collectUniqueDisposableInChannelState( Stream.of(subtaskState.getInputChannelState())) - .forEach(slowMapperChannelStateDelegates::add); + .forEach(delegates.slowMapperInputChannelState::add); } } } + return delegates; + } - assertThat(channelStateDelegates) - .as("channel state delegates in the checkpoint") - .isNotEmpty() - .allSatisfy( - handle -> assertThat(handle).isInstanceOf(SegmentFileStateHandle.class)); - assertThat(channelStateDelegates.stream().mapToLong(StreamStateHandle::getStateSize).sum()) - .isPositive(); - assertThat(slowMapperChannelStateDelegates) - .as("channel state delegates belonging to the stateless slow mapper") - .isNotEmpty(); + /** The channel state delegates found in a checkpoint, split by what the test asserts on. */ + private static final class ChannelStateDelegates { + + private final List all = new ArrayList<>(); + private final List slowMapperInputChannelState = new ArrayList<>(); } private static final class SlowWordMapper extends RichMapFunction> {