diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java index 2339430464c7..64c7543b6ecb 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorker.java @@ -405,28 +405,25 @@ private StreamingWorkerHarnessFactoryOutput createFanOutStreamingEngineWorkerHar .setBytes(MAX_GET_WORK_FETCH_BYTES) .build(), windmillStreamFactory, - (workItem, + (computationState, + workItem, serializedWorkItemSize, watermarks, processingContext, drainMode, appliedFinalizeIds, - getWorkStreamLatencies) -> - checkNotNull(computationStateCache) - .get(processingContext.computationId()) - .ifPresent( - computationState -> { - memoryMonitor.waitForResources("GetWork"); - streamingWorkScheduler.queueAppliedFinalizeIds(appliedFinalizeIds); - streamingWorkScheduler.scheduleWork( - computationState, - workItem, - serializedWorkItemSize, - watermarks, - processingContext, - drainMode, - getWorkStreamLatencies); - }), + getWorkStreamLatencies) -> { + memoryMonitor.waitForResources("GetWork"); + streamingWorkScheduler.queueAppliedFinalizeIds(appliedFinalizeIds); + streamingWorkScheduler.scheduleWork( + computationState, + workItem, + serializedWorkItemSize, + watermarks, + processingContext, + drainMode, + getWorkStreamLatencies); + }, ChannelCachingRemoteStubFactory.create(options.getGcpCredential(), channelCache), GetWorkBudgetDistributors.distributeEvenly(), checkNotNull(dispatcherClient), @@ -441,7 +438,8 @@ private StreamingWorkerHarnessFactoryOutput createFanOutStreamingEngineWorkerHar .setCommitWorkStreamFactory( () -> CloseableStream.create(commitWorkStream, () -> {})) .build(), - getDataMetricTracker); + getDataMetricTracker, + checkNotNull(this.computationStateCache)::get); ChannelzServlet channelzServlet = createChannelzServlet( options, fanOutStreamingEngineWorkerHarness::currentWindmillEndpoints); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkState.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkState.java index de4082581293..519b2b2948a1 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkState.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkState.java @@ -28,18 +28,17 @@ import java.util.Optional; import java.util.Queue; import java.util.function.BiConsumer; -import javax.annotation.Nullable; import javax.annotation.concurrent.GuardedBy; import javax.annotation.concurrent.ThreadSafe; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItem; import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache; -import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache.ForComputation; import org.apache.beam.runners.dataflow.worker.windmill.work.budget.GetWorkBudget; import org.apache.beam.sdk.annotations.Internal; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.annotations.VisibleForTesting; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap; +import org.checkerframework.checker.nullness.qual.Nullable; import org.joda.time.Duration; import org.joda.time.Instant; import org.slf4j.Logger; @@ -78,7 +77,7 @@ public final class ActiveWorkState { private ActiveWorkState( Map> activeWork, - ForComputation computationStateCache) { + WindmillStateCache.ForComputation computationStateCache) { this.activeWork = activeWork; this.computationStateCache = computationStateCache; this.activeGetWorkBudget = GetWorkBudget.noBudget(); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationState.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationState.java index 5e850d4312ea..a03091824104 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationState.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationState.java @@ -22,7 +22,6 @@ import java.util.Map; import java.util.Optional; import java.util.concurrent.ConcurrentLinkedQueue; -import javax.annotation.Nullable; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache; import org.apache.beam.runners.dataflow.worker.windmill.work.budget.GetWorkBudget; @@ -30,6 +29,7 @@ import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap; +import org.checkerframework.checker.nullness.qual.Nullable; import org.joda.time.Instant; /** diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/Work.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/Work.java index 4541a1c313a2..5759be7cecf6 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/Work.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/Work.java @@ -144,22 +144,26 @@ public static Work create( } public static ProcessingContext createProcessingContext( - String computationId, + ComputationState computationState, GetDataClient getDataClient, Consumer workCommitter, HeartbeatSender heartbeatSender) { return ProcessingContext.create( - computationId, getDataClient, workCommitter, heartbeatSender, /* backendWorkerToken= */ ""); + computationState, + getDataClient, + workCommitter, + heartbeatSender, + /* backendWorkerToken= */ ""); } public static ProcessingContext createProcessingContext( - String computationId, + ComputationState computationState, GetDataClient getDataClient, Consumer workCommitter, HeartbeatSender heartbeatSender, String backendWorkerToken) { return ProcessingContext.create( - computationId, getDataClient, workCommitter, heartbeatSender, backendWorkerToken); + computationState, getDataClient, workCommitter, heartbeatSender, backendWorkerToken); } private static LatencyAttribution.Builder createLatencyAttributionWithActiveLatencyBreakdown( @@ -207,6 +211,10 @@ public long getSerializedWorkItemSize() { return serializedWorkItemSize; } + public ComputationState getComputationState() { + return processingContext.computationState(); + } + public String getComputationId() { return processingContext.computationId(); } @@ -457,17 +465,21 @@ public KeyGroup getKeyGroup() { public abstract static class ProcessingContext { private static ProcessingContext create( - String computationId, + ComputationState computationState, GetDataClient getDataClient, Consumer workCommitter, HeartbeatSender heartbeatSender, String backendWorkerToken) { return new AutoValue_Work_ProcessingContext( - computationId, getDataClient, heartbeatSender, workCommitter, backendWorkerToken); + computationState, getDataClient, heartbeatSender, workCommitter, backendWorkerToken); } /** Computation that the {@link Work} belongs to. */ - public abstract String computationId(); + public abstract ComputationState computationState(); + + public String computationId() { + return computationState().getComputationId(); + } /** Handles GetData requests to streaming backend. */ public abstract GetDataClient getDataClient(); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarness.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarness.java index f3262c17b698..a81e7537d077 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarness.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarness.java @@ -40,6 +40,7 @@ import javax.annotation.concurrent.GuardedBy; import javax.annotation.concurrent.ThreadSafe; import org.apache.beam.repackaged.core.org.apache.commons.lang3.tuple.Pair; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc.CloudWindmillServiceV1Alpha1Stub; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GetWorkRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.JobHeader; @@ -96,6 +97,7 @@ public final class FanOutStreamingEngineWorkerHarness implements StreamingWorker private final GetWorkBudget totalGetWorkBudget; private final Function workCommitterFactory; private final ThrottlingGetDataMetricTracker getDataMetricTracker; + private final Function> computationStateFetcher; private final ExecutorService windmillStreamManager; private final ExecutorService workerMetadataConsumer; private final Object metadataLock = new Object(); @@ -131,7 +133,8 @@ private FanOutStreamingEngineWorkerHarness( GrpcDispatcherClient dispatcherClient, Function workCommitterFactory, ThrottlingGetDataMetricTracker getDataMetricTracker, - ExecutorService workerMetadataConsumer) { + ExecutorService workerMetadataConsumer, + Function> computationStateFetcher) { this.jobHeader = jobHeader; this.getDataMetricTracker = getDataMetricTracker; this.started = false; @@ -150,6 +153,7 @@ private FanOutStreamingEngineWorkerHarness( this.activeMetadataType = WindmillEndpoints.Type.UNKNOWN; this.pendingMetadataType = WindmillEndpoints.Type.UNKNOWN; this.workCommitterFactory = workCommitterFactory; + this.computationStateFetcher = computationStateFetcher; } /** @@ -166,7 +170,8 @@ public static FanOutStreamingEngineWorkerHarness create( GetWorkBudgetDistributor getWorkBudgetDistributor, GrpcDispatcherClient dispatcherClient, Function workCommitterFactory, - ThrottlingGetDataMetricTracker getDataMetricTracker) { + ThrottlingGetDataMetricTracker getDataMetricTracker, + Function> computationStateFetcher) { return new FanOutStreamingEngineWorkerHarness( jobHeader, totalGetWorkBudget, @@ -178,9 +183,8 @@ public static FanOutStreamingEngineWorkerHarness create( workCommitterFactory, getDataMetricTracker, Executors.newSingleThreadExecutor( - new ThreadFactoryBuilder() - .setNameFormat(WORKER_METADATA_CONSUMER_THREAD_NAME) - .build())); + new ThreadFactoryBuilder().setNameFormat(WORKER_METADATA_CONSUMER_THREAD_NAME).build()), + computationStateFetcher); } @VisibleForTesting @@ -193,7 +197,8 @@ static FanOutStreamingEngineWorkerHarness forTesting( GetWorkBudgetDistributor getWorkBudgetDistributor, GrpcDispatcherClient dispatcherClient, Function workCommitterFactory, - ThrottlingGetDataMetricTracker getDataMetricTracker) { + ThrottlingGetDataMetricTracker getDataMetricTracker, + Function> computationStateFetcher) { FanOutStreamingEngineWorkerHarness fanOutStreamingEngineWorkProvider = new FanOutStreamingEngineWorkerHarness( jobHeader, @@ -210,7 +215,8 @@ static FanOutStreamingEngineWorkerHarness forTesting( // blocked by the consumeWorkerMetadata() task. Test suites run in different // environments and non-determinism has lead to past flakiness. See // https://github.com/apache/beam/issues/28957. - MoreExecutors.newDirectExecutorService()); + MoreExecutors.newDirectExecutorService(), + computationStateFetcher); fanOutStreamingEngineWorkProvider.start(); return fanOutStreamingEngineWorkProvider; } @@ -448,7 +454,8 @@ private WindmillStreamSender createAndStartWindmillStreamSender(Endpoint endpoin getDataStream -> StreamGetDataClient.create( getDataStream, this::getGlobalDataStream, getDataMetricTracker), - workCommitterFactory); + workCommitterFactory, + computationStateFetcher); windmillStreamSender.start(); return windmillStreamSender; } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java index f41223310385..00c949009206 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/SingleSourceWorkerHarness.java @@ -179,7 +179,7 @@ private void streamingEngineDispatchLoop( .setOutputDataWatermark(workItem.getOutputDataWatermark()) .build(), Work.createProcessingContext( - computationId, + computationState, getDataClient, workCommitter::commit, heartbeatSender), @@ -250,7 +250,7 @@ private void applianceDispatchLoop(Supplier getWorkFn) workItem.getSerializedSize(), watermarks.setOutputDataWatermark(workItem.getOutputDataWatermark()).build(), Work.createProcessingContext( - computationId, getDataClient, workCommitter::commit, heartbeatSender), + computationState, getDataClient, workCommitter::commit, heartbeatSender), computationWork.getDrainMode(), /* getWorkStreamLatencies= */ ImmutableList.of()); } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSender.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSender.java index d150ee6bf1d1..5abe93f234a1 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSender.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSender.java @@ -19,6 +19,7 @@ import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkState; +import java.util.Optional; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; @@ -27,6 +28,7 @@ import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; import javax.annotation.concurrent.ThreadSafe; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GetWorkRequest; import org.apache.beam.runners.dataflow.worker.windmill.WindmillConnection; import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.CommitWorkStream; @@ -75,7 +77,8 @@ private WindmillStreamSender( GrpcWindmillStreamFactory streamingEngineStreamFactory, WorkItemScheduler workItemScheduler, Function getDataClientFactory, - Function workCommitterFactory) { + Function workCommitterFactory, + Function> computationStateFetcher) { this.started = new AtomicBoolean(false); this.getWorkBudget = getWorkBudget; @@ -91,7 +94,8 @@ private WindmillStreamSender( FixedStreamHeartbeatSender.create(getDataStream), getDataClientFactory.apply(getDataStream), workCommitter, - workItemScheduler); + workItemScheduler, + computationStateFetcher); // 3 threads, 1 for each stream type (GetWork, GetData, CommitWork). this.streamStarter = Executors.newFixedThreadPool( @@ -105,7 +109,8 @@ static WindmillStreamSender create( GrpcWindmillStreamFactory streamingEngineStreamFactory, WorkItemScheduler workItemScheduler, Function getDataClientFactory, - Function workCommitterFactory) { + Function workCommitterFactory, + Function> computationStateFetcher) { return new WindmillStreamSender( connection, getWorkRequest, @@ -113,7 +118,8 @@ static WindmillStreamSender create( streamingEngineStreamFactory, workItemScheduler, getDataClientFactory, - workCommitterFactory); + workCommitterFactory, + computationStateFetcher); } private static GetWorkRequest withRequestBudget(GetWorkRequest request, GetWorkBudget budget) { diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java index 9eb9a37b1b76..e11bb587bf0d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java @@ -395,12 +395,21 @@ BoundedQueueExecutorWorkHandleImpl createBudgetHandle(Work work, long bytes) { if (keyGroupWorkQueue == null) { return null; } - @Nullable QueuedWork queuedWork = keyGroupWorkQueue.pollWork(computationId, keyGroup); - if (queuedWork == null) { - return null; + while (true) { + @Nullable QueuedWork queuedWork = keyGroupWorkQueue.pollWork(computationId, keyGroup); + if (queuedWork == null) { + return null; + } + Work work = queuedWork.getWork().work(); + if (work.isFailed()) { + queuedWork.getHandle().close(); + work.getComputationState() + .completeWorkAndScheduleNextWorkForKey(work.getShardedKey(), work.id()); + continue; + } + internalHandle.merge(queuedWork.getHandle()); + return queuedWork.getWork(); } - internalHandle.merge(queuedWork.getHandle()); - return queuedWork.getWork(); } private void decrementCounters(int elements, long bytes) { diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStream.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStream.java index de8ebf14b709..546a957a4861 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStream.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStream.java @@ -21,6 +21,7 @@ import java.io.PrintWriter; import java.time.Duration; +import java.util.Optional; import java.util.Set; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; @@ -29,6 +30,7 @@ import java.util.function.Function; import javax.annotation.concurrent.GuardedBy; import net.jcip.annotations.ThreadSafe; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; @@ -81,6 +83,7 @@ final class GrpcDirectGetWorkStream private final HeartbeatSender heartbeatSender; private final WorkCommitter workCommitter; private final GetDataClient getDataClient; + private final Function> computationStateFetcher; private final AtomicReference lastRequest; private final boolean requestBatchedGetWorkResponse; @@ -102,7 +105,8 @@ private GrpcDirectGetWorkStream( WorkCommitter workCommitter, WorkItemScheduler workItemScheduler, Duration halfClosePhysicalStreamAfter, - ScheduledExecutorService executorService) { + ScheduledExecutorService executorService, + Function> computationStateFetcher) { super( LOG, startGetWorkRpcFn, @@ -118,6 +122,7 @@ private GrpcDirectGetWorkStream( this.heartbeatSender = heartbeatSender; this.workCommitter = workCommitter; this.getDataClient = getDataClient; + this.computationStateFetcher = computationStateFetcher; this.lastRequest = new AtomicReference<>(); this.budgetTracker = new GetWorkBudgetTracker( @@ -145,7 +150,8 @@ static GrpcDirectGetWorkStream create( WorkCommitter workCommitter, WorkItemScheduler workItemScheduler, Duration halfClosePhysicalStreamAfter, - ScheduledExecutorService executor) { + ScheduledExecutorService executor, + Function> computationStateFetcher) { return new GrpcDirectGetWorkStream( backendWorkerToken, startGetWorkRpcFn, @@ -160,7 +166,8 @@ static GrpcDirectGetWorkStream create( workCommitter, workItemScheduler, halfClosePhysicalStreamAfter, - executor); + executor, + computationStateFetcher); } private static Watermarks createWatermarks( @@ -273,25 +280,37 @@ protected void sendHealthCheck() throws WindmillStreamShutdownException { } private void consumeAssembledWorkItem(AssembledWorkItem assembledWorkItem) { - WorkItem workItem = assembledWorkItem.workItem(); GetWorkResponseChunkAssembler.ComputationMetadata metadata = assembledWorkItem.computationMetadata(); - workItemScheduler.scheduleWork( - workItem, - assembledWorkItem.bufferedSize(), - createWatermarks(workItem, metadata), - createProcessingContext(metadata.computationId()), - metadata.drainMode(), - assembledWorkItem.appliedFinalizeIds(), - assembledWorkItem.latencyAttributions()); + Optional maybeComputationState = + computationStateFetcher.apply(metadata.computationId()); + if (maybeComputationState.isPresent()) { + ComputationState computationState = maybeComputationState.get(); + WorkItem workItem = assembledWorkItem.workItem(); + workItemScheduler.scheduleWork( + computationState, + workItem, + assembledWorkItem.bufferedSize(), + createWatermarks(workItem, metadata), + createProcessingContext(computationState), + metadata.drainMode(), + assembledWorkItem.appliedFinalizeIds(), + assembledWorkItem.latencyAttributions()); + } else { + LOG.warn("Received work for unknown computation: {}", metadata.computationId()); + } budgetTracker.recordBudgetReceived(assembledWorkItem.bufferedSize()); GetWorkBudget extension = budgetTracker.computeBudgetExtension(); maybeSendRequestExtension(extension); } - private Work.ProcessingContext createProcessingContext(String computationId) { + private Work.ProcessingContext createProcessingContext(ComputationState computationState) { return Work.createProcessingContext( - computationId, getDataClient, workCommitter::commit, heartbeatSender, backendWorkerToken()); + computationState, + getDataClient, + workCommitter::commit, + heartbeatSender, + backendWorkerToken()); } @Override diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcWindmillStreamFactory.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcWindmillStreamFactory.java index 97ca3c4e83d7..f4465413f43d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcWindmillStreamFactory.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcWindmillStreamFactory.java @@ -23,6 +23,7 @@ import java.io.PrintWriter; import java.util.Collection; import java.util.List; +import java.util.Optional; import java.util.Set; import java.util.Timer; import java.util.TimerTask; @@ -37,6 +38,7 @@ import java.util.function.Supplier; import javax.annotation.concurrent.ThreadSafe; import org.apache.beam.runners.dataflow.worker.status.StatusDataProvider; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillMetadataServiceV1Alpha1Grpc.CloudWindmillMetadataServiceV1Alpha1Stub; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc.CloudWindmillServiceV1Alpha1Stub; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationHeartbeatResponse; @@ -287,7 +289,8 @@ public GetWorkStream createDirectGetWorkStream( HeartbeatSender heartbeatSender, GetDataClient getDataClient, WorkCommitter workCommitter, - WorkItemScheduler workItemScheduler) { + WorkItemScheduler workItemScheduler, + Function> computationStateFetcher) { return GrpcDirectGetWorkStream.create( connection.backendWorkerToken(), responseObserver -> @@ -303,7 +306,8 @@ public GetWorkStream createDirectGetWorkStream( workCommitter, workItemScheduler, directStreamingRpcPhysicalStreamHalfCloseAfter, - executorForDirectStreams(connection.backendWorkerToken(), "GetWork")); + executorForDirectStreams(connection.backendWorkerToken(), "GetWork"), + computationStateFetcher); } public GetDataStream createGetDataStream(CloudWindmillServiceV1Alpha1Stub stub) { diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/WorkItemScheduler.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/WorkItemScheduler.java index a2dfa50a0d63..be4c5562031f 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/WorkItemScheduler.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/WorkItemScheduler.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.windmill.work; import javax.annotation.CheckReturnValue; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution; @@ -32,6 +33,7 @@ public interface WorkItemScheduler { /** * Schedule {@link WorkItem}(s). * + * @param computationState {@link ComputationState} for the workItem. * @param workItem {@link WorkItem} to be processed. * @param watermarks processing watermarks for the workItem. * @param processingContext for processing the workItem. @@ -41,6 +43,7 @@ public interface WorkItemScheduler { * back to Streaming Engine backend. */ void scheduleWork( + ComputationState computationState, WorkItem workItem, long serializedWorkItemSize, Watermarks watermarks, diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java index 9ed705550bc6..055890d0f6af 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java @@ -20,6 +20,7 @@ import static org.apache.beam.runners.dataflow.util.Structs.addObject; import static org.apache.beam.runners.dataflow.util.Structs.addString; import static org.apache.beam.runners.dataflow.worker.counters.DataflowCounterUpdateExtractor.splitIntToLong; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.hamcrest.MatcherAssert.assertThat; import static org.hamcrest.Matchers.both; import static org.hamcrest.Matchers.contains; @@ -382,7 +383,10 @@ private static ExecutableWork createMockWork( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), + createMockComputationState(computationId), + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()), @@ -1524,13 +1528,144 @@ public void testCompleteCommit_retryableFailureTriggersReExecution() throws Exce worker.stop(); } + @Test + public void testMultiKeyCommit_queuedWorkItemFailsAndSubsequentWorkItemPickedUp() + throws Exception { + if (!streamingEngine) { + return; + } + BlockingKvDoFn.reset(); + StreamingDataflowWorker worker = makeMultiKeyEnabledWorker(new BlockingKvDoFn()); + worker.start(); + + String batchInputText1 = + "work {" + + " computation_id: \"" + + DEFAULT_COMPUTATION_ID + + "\"" + + " input_data_watermark: 0" + + " work {" + + " key: \"key1\"" + + " sharding_key: 1" + + " work_token: 1" + + " cache_token: 2" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data1\"" + + " }" + + " }" + + " }" + + " work {" + + " key: \"key2\"" + + " sharding_key: 2" + + " work_token: 2" + + " cache_token: 3" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data2\"" + + " }" + + " }" + + " }" + + "}"; + + String batchInputText2 = + "work {" + + " computation_id: \"" + + DEFAULT_COMPUTATION_ID + + "\"" + + " input_data_watermark: 0" + + " work {" + + " key: \"key2\"" + + " sharding_key: 2" + + " work_token: 3" + + " cache_token: 4" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data3\"" + + " }" + + " }" + + " }" + + "}"; + Windmill.GetWorkResponse batchInput1 = + buildInput( + batchInputText1, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + Windmill.GetWorkResponse batchInput2 = + buildInput( + batchInputText2, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + + server.whenGetDataCalled().answerByDefault(StreamingDataflowWorkerTest::emptyDataResponder); + + server.whenGetWorkCalled().thenReturn(batchInput1).thenReturn(batchInput2); + server.waitForEmptyWorkQueue(); + + // Wait for key1 to start processing and block on BlockingKvDoFn. + BlockingKvDoFn.counter.get().acquire(1); + + // Fail key2 (work token 2) via failed heartbeat while key1 is still processing. + ComputationHeartbeatResponse.Builder failedHeartbeat = + ComputationHeartbeatResponse.newBuilder(); + failedHeartbeat + .setComputationId(DEFAULT_COMPUTATION_ID) + .addHeartbeatResponsesBuilder() + .setCacheToken(3) + .setWorkToken(2) + .setShardingKey(2) + .setFailed(true); + server.sendFailedHeartbeats(Collections.singletonList(failedHeartbeat.build())); + + // Unblock key1 to allow bundle to poll key2 (token 2 -> failed, skipped) and key2 (token 3). + BlockingKvDoFn.blocker.get().countDown(); + + Map result = server.waitForAndGetCommits(2); + + assertTrue(result.containsKey(1L)); + assertTrue(result.containsKey(3L)); + assertFalse(result.containsKey(2L)); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(1, multiKeyCommits.size()); + Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = multiKeyCommits.get(0); + assertEquals(2, multiKeyCommit.getRequestsCount()); + assertEquals(1, multiKeyCommit.getRequests(0).getWorkToken()); + assertEquals(3, multiKeyCommit.getRequests(1).getWorkToken()); + + worker.stop(); + } + private StreamingDataflowWorker makeMultiKeyEnabledWorker() { + return makeMultiKeyEnabledWorker(new WorkDoFn()); + } + + private StreamingDataflowWorker makeMultiKeyEnabledWorker( + DoFn, KV> doFn) { KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); List instructions = Arrays.asList( makeSourceInstruction(kvCoder), - makeDoFnInstruction(new WorkDoFn(), 0, kvCoder), + makeDoFnInstruction(doFn, 0, kvCoder), makeSinkInstruction(kvCoder, 1)); StreamingDataflowWorker worker = @@ -3966,7 +4101,7 @@ public void testLatencyAttributionProtobufsPopulated() { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - "computationId", + createMockComputationState("computationId"), new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), @@ -4837,6 +4972,25 @@ public void processElement(ProcessContext c, @StateId("state") ValueState, KV> { + public static final AtomicReference blocker = + new AtomicReference<>(new CountDownLatch(1)); + public static final AtomicReference counter = + new AtomicReference<>(new Semaphore(0)); + + @ProcessElement + public void processElement(ProcessContext c) throws InterruptedException { + counter.get().release(); + blocker.get().await(); + c.output(c.element()); + } + + public static void reset() { + blocker.set(new CountDownLatch(1)); + counter.set(new Semaphore(0)); + } + } + static class LargeCommitFn extends DoFn, KV> { @ProcessElement diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java index c5efcea4e47c..5766b8196516 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java @@ -19,6 +19,7 @@ import static org.apache.beam.runners.dataflow.worker.counters.DataflowCounterUpdateExtractor.longToSplitInt; import static org.apache.beam.runners.dataflow.worker.counters.DataflowCounterUpdateExtractor.splitIntToLong; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.hamcrest.MatcherAssert.assertThat; import static org.hamcrest.Matchers.containsInAnyOrder; import static org.hamcrest.Matchers.equalTo; @@ -175,7 +176,10 @@ private static Work createMockWork(Windmill.WorkItem workItem, Watermarks waterm workItem.getSerializedSize(), watermarks, Work.createProcessingContext( - COMPUTATION_ID, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), + createMockComputationState(COMPUTATION_ID), + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()); diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindmillReaderIteratorBaseTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindmillReaderIteratorBaseTest.java index be77da540889..7c6301fd0411 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindmillReaderIteratorBaseTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindmillReaderIteratorBaseTest.java @@ -17,6 +17,7 @@ */ package org.apache.beam.runners.dataflow.worker; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; @@ -261,7 +262,7 @@ private static Work createMockWork(Windmill.WorkItem workItem) { .setInputDataWatermark(new org.joda.time.Instant(1000)) .build(), Work.createProcessingContext( - "computationId", + createMockComputationState("computationId"), mock( org.apache.beam.runners.dataflow.worker.windmill.client.getdata.GetDataClient .class), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindowingWindmillReaderTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindowingWindmillReaderTest.java index 3c778650eb3e..fa95e64fc90e 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindowingWindmillReaderTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindowingWindmillReaderTest.java @@ -17,6 +17,7 @@ */ package org.apache.beam.runners.dataflow.worker; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; @@ -91,7 +92,10 @@ private static Work createMockWork(Windmill.WorkItem workItem) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(new Instant(1000)).build(), Work.createProcessingContext( - "computationId", new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), + createMockComputationState("computationId"), + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()); diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java index 679227a11dc0..2a0096b1ee18 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java @@ -26,6 +26,7 @@ import static org.apache.beam.runners.dataflow.worker.SourceTranslationUtils.readerProgressToCloudProgress; import static org.apache.beam.runners.dataflow.worker.WorkerCustomSources.BoundedReaderIterator.getReaderProgress; import static org.apache.beam.runners.dataflow.worker.WorkerCustomSources.BoundedReaderIterator.longToParallelism; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.apache.beam.sdk.testing.ExpectedLogs.verifyLogged; import static org.apache.beam.sdk.testing.SourceTestUtils.readFromSource; import static org.apache.beam.sdk.util.CoderUtils.encodeToByteArray; @@ -207,7 +208,10 @@ private static Work createMockWork(Windmill.WorkItem workItem, Watermarks waterm workItem.getSerializedSize(), watermarks, Work.createProcessingContext( - COMPUTATION_ID, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), + createMockComputationState(COMPUTATION_ID), + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()); @@ -1046,7 +1050,7 @@ public void testFailedWorkItemsAbort() throws Exception { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(new Instant(0)).build(), Work.createProcessingContext( - COMPUTATION_ID, + createMockComputationState(COMPUTATION_ID), new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkStateTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkStateTest.java index aa0eae0d159f..9e65bc57119c 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkStateTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ActiveWorkStateTest.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.streaming; import static com.google.common.truth.Truth.assertThat; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertNotNull; @@ -94,7 +95,10 @@ private static ExecutableWork expiredWork(Windmill.WorkItem workItem) { private static Work.ProcessingContext createWorkProcessingContext() { return Work.createProcessingContext( - "computationId", new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)); + createMockComputationState("computationId"), + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)); } private static WorkId workId(long workToken, long cacheToken) { diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateCacheTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateCacheTest.java index f57e20d4b5fb..0ea47e4037f7 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateCacheTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateCacheTest.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.streaming; import static com.google.common.truth.Truth.assertThat; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.eq; @@ -72,7 +73,7 @@ private static ExecutableWork createWork(ShardedKey shardedKey, long workToken, workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - "computationId", + createMockComputationState("computationId"), new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTest.java index 22ddc8e4de5b..6184560670c7 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTest.java @@ -58,7 +58,7 @@ private ExecutableWork createWork(Windmill.WorkItem workItem) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - "computationId", new FakeGetDataClient(), ignored -> {}, mockHeartbeatSender), + computationState, new FakeGetDataClient(), ignored -> {}, mockHeartbeatSender), false, Instant::now, ImmutableList.of()), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTestUtils.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTestUtils.java new file mode 100644 index 000000000000..bfe3e0c87a21 --- /dev/null +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTestUtils.java @@ -0,0 +1,33 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.beam.runners.dataflow.worker.streaming; + +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +/** Test utilities for creating and manipulating {@link ComputationState} objects in unit tests. */ +public final class ComputationStateTestUtils { + + private ComputationStateTestUtils() {} + + public static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } +} diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/WorkTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/WorkTest.java index 61e52ddd61bd..ad7ed86e2495 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/WorkTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/WorkTest.java @@ -17,6 +17,7 @@ */ package org.apache.beam.runners.dataflow.worker.streaming; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; @@ -51,7 +52,7 @@ private static Work createTestWork() { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - "comp", + createMockComputationState("comp"), mock( org.apache.beam.runners.dataflow.worker.windmill.client.getdata.GetDataClient .class), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarnessTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarnessTest.java index cf6df7f0e478..03944de29d9d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarnessTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/FanOutStreamingEngineWorkerHarnessTest.java @@ -32,11 +32,13 @@ import java.io.IOException; import java.util.ArrayList; import java.util.HashSet; +import java.util.Optional; import java.util.Set; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.stream.Collectors; import javax.annotation.Nullable; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.util.MemoryMonitor; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillMetadataServiceV1Alpha1Grpc; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc; @@ -122,7 +124,8 @@ public class FanOutStreamingEngineWorkerHarnessTest { private FanOutStreamingEngineWorkerHarness fanOutStreamingEngineWorkProvider; private static WorkItemScheduler noOpProcessWorkItemFn() { - return (workItem, + return (computationState, + workItem, serializedWorkItemSize, watermarks, processingContext, @@ -195,7 +198,8 @@ private FanOutStreamingEngineWorkerHarness newFanOutStreamingEngineWorkerHarness getWorkBudgetDistributor, dispatcherClient, ignored -> mock(WorkCommitter.class), - new ThrottlingGetDataMetricTracker(mock(MemoryMonitor.class))); + new ThrottlingGetDataMetricTracker(mock(MemoryMonitor.class)), + ignored -> Optional.of(mock(ComputationState.class))); getWorkerMetadataReady.await(); return harness; } @@ -246,7 +250,8 @@ public void testStreamsStartCorrectly() throws InterruptedException { any(), any(), any(), - eq(noOpProcessWorkItemFn())); + eq(noOpProcessWorkItemFn()), + any()); verify(streamFactory, times(1)) .createDirectGetWorkStream( any(), @@ -254,7 +259,8 @@ public void testStreamsStartCorrectly() throws InterruptedException { any(), any(), any(), - eq(noOpProcessWorkItemFn())); + eq(noOpProcessWorkItemFn()), + any()); verify(streamFactory, times(2)).createDirectGetDataStream(any()); verify(streamFactory, times(2)).createDirectCommitWorkStream(any()); diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSenderTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSenderTest.java index 457f75593e23..ae8c02a91278 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSenderTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/harness/WindmillStreamSenderTest.java @@ -26,6 +26,8 @@ import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; +import java.util.Optional; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GetWorkRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.JobHeader; @@ -64,7 +66,8 @@ public class WindmillStreamSenderTest { .build()) .build()); private final WorkItemScheduler workItemScheduler = - (workItem, + (computationState, + workItem, serializedWorkItemSize, watermarks, processingContext, @@ -115,7 +118,8 @@ public void testStartStream_startsAllStreams() { any(), any(), any(), - eq(workItemScheduler)); + eq(workItemScheduler), + any()); verify(streamFactory).createDirectGetDataStream(eq(connection)); verify(streamFactory).createDirectCommitWorkStream(eq(connection)); @@ -146,7 +150,8 @@ public void testStartStream_onlyStartsStreamsOnce() { any(), any(), any(), - eq(workItemScheduler)); + eq(workItemScheduler), + any()); verify(streamFactory, times(1)).createDirectGetDataStream(eq(connection)); verify(streamFactory, times(1)).createDirectCommitWorkStream(eq(connection)); @@ -180,7 +185,8 @@ public void testStartStream_onlyStartsStreamsOnceConcurrent() throws Interrupted any(), any(), any(), - eq(workItemScheduler)); + eq(workItemScheduler), + any()); verify(streamFactory, times(1)).createDirectGetDataStream(eq(connection)); verify(streamFactory, times(1)).createDirectCommitWorkStream(eq(connection)); @@ -203,7 +209,8 @@ public void testCloseAllStreams_closesAllStreams() { any(), any(), any(), - eq(workItemScheduler))) + eq(workItemScheduler), + any())) .thenReturn(mockGetWorkStream); when(mockStreamFactory.createDirectGetDataStream(eq(connection))).thenReturn(mockGetDataStream); @@ -240,7 +247,8 @@ public void testCloseAllStreams_doesNotStartStreamsAfterClose() { any(), any(), any(), - eq(workItemScheduler))) + eq(workItemScheduler), + any())) .thenReturn(mockGetWorkStream); when(mockStreamFactory.createDirectGetDataStream(eq(connection))).thenReturn(mockGetDataStream); @@ -289,6 +297,7 @@ private WindmillStreamSender newWindmillStreamSender( streamFactory, workItemScheduler, ignored -> mock(GetDataClient.class), - ignored -> mock(WorkCommitter.class)); + ignored -> mock(WorkCommitter.class), + ignored -> Optional.of(mock(ComputationState.class))); } } diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutorTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutorTest.java index 0e75fa01f4f0..51b13d1218fa 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutorTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutorTest.java @@ -17,6 +17,7 @@ */ package org.apache.beam.runners.dataflow.worker.util; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.hamcrest.Matchers.greaterThan; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; @@ -25,6 +26,7 @@ import static org.junit.Assert.assertThat; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; import java.util.Arrays; import java.util.Collection; @@ -34,6 +36,7 @@ import java.util.function.BiConsumer; import java.util.function.Consumer; import org.apache.beam.runners.dataflow.worker.streaming.BoundedQueueExecutorWorkHandle; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; @@ -90,6 +93,39 @@ private static ExecutableWork createWorkWithCompIdAndKeyGroup( computationId, keyGroup, (work, handle) -> executeWorkFn.accept(work)); } + private static ExecutableWork createWorkWithComputationStateAndKeyGroup( + ComputationState computationState, + Work.KeyGroup keyGroup, + long workToken, + Consumer executeWorkFn) { + WorkItem workItem = + WorkItem.newBuilder() + .setKey(ByteString.EMPTY) + .setShardingKey(1) + .setWorkToken(workToken) + .setCacheToken(1) + .setKeyGroup( + Windmill.Uint128Proto.newBuilder() + .setHigh(keyGroup.high()) + .setLow(keyGroup.low()) + .build()) + .build(); + return ExecutableWork.create( + Work.create( + workItem, + workItem.getSerializedSize(), + Watermarks.builder().setInputDataWatermark(Instant.now()).build(), + Work.createProcessingContext( + computationState, + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)), + false, + Instant::now, + ImmutableList.of()), + (work, handle) -> executeWorkFn.accept(work)); + } + private static ExecutableWork createWorkWithHandle( String computationId, Work.KeyGroup keyGroup, @@ -112,7 +148,10 @@ private static ExecutableWork createWorkWithHandle( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), + createMockComputationState(computationId), + new FakeGetDataClient(), + ignored -> {}, + mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()), @@ -575,4 +614,79 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { blockerStop.countDown(); testExecutor.shutdown(); } + + @Test + public void testPollWork_skipsFailedWorkAndCallsCompleteWorkAndScheduleNextWorkForKey() + throws Exception { + BoundedQueueExecutor testExecutor = + new BoundedQueueExecutor( + 1, + 60, + TimeUnit.SECONDS, + 100, + 10000000, + new ThreadFactoryBuilder().setNameFormat("testPollWork-%d").setDaemon(true).build(), + useFairMonitor, + /* useKeyGroupWorkQueue= */ true); + + CountDownLatch blockerStart = new CountDownLatch(1); + CountDownLatch blockerStop = new CountDownLatch(1); + AtomicReference blockerHandleRef = new AtomicReference<>(); + ExecutableWork blockerWork = + createWorkWithHandle( + "compA", + DEFAULT_KEY_GROUP, + (work, handle) -> { + blockerHandleRef.set(handle); + blockerStart.countDown(); + try { + blockerStop.await(); + } catch (InterruptedException e) { + throw new RuntimeException(e); + } + }); + + testExecutor.execute(blockerWork, 10); + blockerStart.await(); + BoundedQueueExecutorWorkHandleImpl stealHandle = + (BoundedQueueExecutorWorkHandleImpl) blockerHandleRef.get(); + assertNotNull(stealHandle); + + Work.KeyGroup keyGroup = Work.KeyGroup.create(1, 1); + ComputationState mockCompState = createMockComputationState("compA"); + + ExecutableWork work1 = + createWorkWithComputationStateAndKeyGroup(mockCompState, keyGroup, 101, ignored -> {}); + ExecutableWork work2 = + createWorkWithComputationStateAndKeyGroup(mockCompState, keyGroup, 102, ignored -> {}); + + // Enqueue both tasks (they will wait in the queue because the thread is blocked). + testExecutor.execute(work1, 100); + testExecutor.execute(work2, 150); + + assertEquals(3, testExecutor.elementsOutstanding()); + assertEquals(260, testExecutor.bytesOutstanding()); + + // Mark work1 as failed while waiting in the queue. + work1.work().setFailed(); + + // pollWork should skip work1, close work1's handle, invoke + // completeWorkAndScheduleNextWorkForKey on mockCompState, + // and return work2. + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup, stealHandle); + assertNotNull(stolen); + assertEquals(work2, stolen); + + verify(mockCompState) + .completeWorkAndScheduleNextWorkForKey(work1.work().getShardedKey(), work1.work().id()); + + // Verify stealHandle merged (blockerWork: 10 bytes, work2: 150 bytes). + assertEquals(160, stealHandle.bytes()); + + // Polling again should return null since no more tasks exist for keyGroup. + assertNull(testExecutor.pollWork("compA", keyGroup, stealHandle)); + + blockerStop.countDown(); + testExecutor.shutdown(); + } } diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueueTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueueTest.java index 77fcb0597586..6815df61cd04 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueueTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueueTest.java @@ -17,6 +17,7 @@ */ package org.apache.beam.runners.dataflow.worker.util; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertNotNull; @@ -109,7 +110,7 @@ private QueuedWork createQueuedWork( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - computationId, + createMockComputationState(computationId), new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java index b0ca89ac4c2b..c34f7b07616d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.windmill.client.commits; import static com.google.common.truth.Truth.assertThat; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertNotNull; import static org.mockito.Mockito.mock; @@ -37,6 +38,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItem; import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; +import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache; import org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender; import org.apache.beam.vendor.grpc.v1p69p0.com.google.protobuf.ByteString; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList; @@ -69,7 +71,7 @@ private static Work createMockWork(long workToken) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - "computationId", + createMockComputationState("computationId"), new FakeGetDataClient(), ignored -> { throw new UnsupportedOperationException(); @@ -86,7 +88,7 @@ private static ComputationState createComputationState(String computationId) { new MapTask().setSystemName("system").setStageName("stage"), mock(BoundedQueueExecutor.class), ImmutableMap.of(), - null); + mock(WindmillStateCache.ForComputation.class)); } private StreamingApplianceWorkCommitter createWorkCommitter( diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java index 4c2b8a9f44fb..3961e4c02886 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.windmill.client.commits; import static com.google.common.truth.Truth.assertThat; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus.OK; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; @@ -61,6 +62,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.CommitWorkStream; import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStreamPool; import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; +import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache; import org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender; import org.apache.beam.vendor.grpc.v1p69p0.com.google.protobuf.ByteString; import org.apache.beam.vendor.grpc.v1p69p0.io.grpc.testing.GrpcCleanupRule; @@ -118,7 +120,7 @@ private static Work createMockWork(long workToken) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - "computationId", + createMockComputationState("computationId"), new FakeGetDataClient(), ignored -> { throw new UnsupportedOperationException(); @@ -135,7 +137,7 @@ private static ComputationState createComputationState(String computationId) { new MapTask().setSystemName("system").setStageName("stage"), mock(BoundedQueueExecutor.class), ImmutableMap.of(), - null); + mock(WindmillStateCache.ForComputation.class)); } private static CompleteCommit asCompleteCommit( diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStreamTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStreamTest.java index 71e1300d90cf..c53cfa3dc326 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStreamTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStreamTest.java @@ -29,11 +29,13 @@ import java.util.Collections; import java.util.HashSet; import java.util.List; +import java.util.Optional; import java.util.Set; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.stream.Collectors; import javax.annotation.Nullable; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationWorkItemMetadata; @@ -66,7 +68,8 @@ public class GrpcDirectGetWorkStreamTest { private static final WorkItemScheduler NO_OP_WORK_ITEM_SCHEDULER = - (workItem, + (computationState, + workItem, serializedWorkItemSize, watermarks, processingContext, @@ -161,7 +164,8 @@ private GrpcDirectGetWorkStream createGetWorkStream( mock(HeartbeatSender.class), mock(GetDataClient.class), mock(WorkCommitter.class), - workItemScheduler); + workItemScheduler, + ignored -> Optional.of(mock(ComputationState.class))); getWorkStream.start(); return getWorkStream; } @@ -281,7 +285,8 @@ public void testConsumedWorkItem_computesAndSendsCorrectExtension() throws Inter createGetWorkStream( testStub, initialBudget, - (work, + (computationState, + work, serializedWorkItemSize, watermarks, processingContext, @@ -331,7 +336,8 @@ public void testConsumedWorkItem_doesNotSendExtensionIfOutstandingBudgetHigh() createGetWorkStream( testStub, initialBudget, - (work, + (computationState, + work, serializedWorkItemSize, watermarks, processingContext, @@ -370,7 +376,8 @@ public void testConsumedWorkItems() throws InterruptedException { createGetWorkStream( testStub, initialBudget, - (work, + (computationState, + work, serializedWorkItemSize, watermarks, processingContext, @@ -415,7 +422,8 @@ public void testConsumedWorkItems_itemsSplitAcrossResponses() throws Interrupted createGetWorkStream( testStub, initialBudget, - (work, + (computationState, + work, serializedWorkItemSize, watermarks, processingContext, diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessorTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessorTest.java index 89f3aa0c0d98..ac069bfcf178 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessorTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessorTest.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures; import static com.google.common.truth.Truth.assertThat; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.assertThrows; import static org.mockito.Mockito.mock; @@ -93,7 +94,7 @@ private static ExecutableWork createWork(Supplier clock, Consumer workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - "computationId", + createMockComputationState("computationId"), new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/refresh/ActiveWorkRefresherTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/refresh/ActiveWorkRefresherTest.java index caa25bf83090..04d54b61aeb3 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/refresh/ActiveWorkRefresherTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/refresh/ActiveWorkRefresherTest.java @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.windmill.work.refresh; import static com.google.common.truth.Truth.assertThat; +import static org.apache.beam.runners.dataflow.worker.streaming.ComputationStateTestUtils.createMockComputationState; import static org.junit.Assert.*; import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.eq; @@ -136,7 +137,10 @@ private ExecutableWork createOldWork( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - "computationId", new FakeGetDataClient(), ignored -> {}, heartbeatSender), + createMockComputationState("computationId"), + new FakeGetDataClient(), + ignored -> {}, + heartbeatSender), false, ActiveWorkRefresherTest::aLongTimeAgo, ImmutableList.of()),