From 2c1278636a603fb5e06a5f5bf6f179b0fbf41d7c Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 11 Jun 2026 13:27:09 +0000 Subject: [PATCH 1/8] Drop failed work in BoundedQueueExecutor::pollWork --- .../worker/util/BoundedQueueExecutor.java | 16 +++-- .../worker/util/BoundedQueueExecutorTest.java | 67 +++++++++++++++++++ 2 files changed, 78 insertions(+), 5 deletions(-) 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 8964246c1160..d9f4ae96476b 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 @@ -391,12 +391,18 @@ BoundedQueueExecutorWorkHandleImpl createBudgetHandle(int elements, 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; + } + if (queuedWork.getWork().work().isFailed()) { + queuedWork.getHandle().close(); + } else { + 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/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 a98102751fb2..9106133cec24 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 @@ -553,4 +553,71 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { blockerStop.countDown(); testExecutor.shutdown(); } + + @Test + public void testPollWorkDropsFailedWork() throws Exception { + BoundedQueueExecutor testExecutor = + new BoundedQueueExecutor( + /* initialMaximumPoolSize= */ 1, + /* keepAliveTime= */ 60, + /* unit= */ TimeUnit.SECONDS, + /* maximumElementsOutstanding= */ 100, + /* maximumBytesOutstanding= */ 10000000, + new ThreadFactoryBuilder().setNameFormat("testStealing-%d").setDaemon(true).build(), + useFairMonitor, + /*useKeyGroupWorkQueue=*/ true); + + // Create blocker task to occupy the worker thread + CountDownLatch blockerStart = new CountDownLatch(1); + CountDownLatch blockerStop = new CountDownLatch(1); + ExecutableWork blockerWork = + createWorkWithCompIdAndKeyGroup( + "blockerComp", + DEFAULT_KEY_GROUP, + ignored -> { + blockerStart.countDown(); + try { + blockerStop.await(); + } catch (InterruptedException e) { + throw new RuntimeException(e); + } + }); + + testExecutor.execute(blockerWork, 0); + blockerStart.await(); + + Work.KeyGroup keyGroup1 = Work.KeyGroup.create(1, 1); + + // Create executable tasks + ExecutableWork work1 = createWorkWithCompIdAndKeyGroup("compA", keyGroup1, ignored -> {}); + ExecutableWork work2 = createWorkWithCompIdAndKeyGroup("compA", keyGroup1, ignored -> {}); + + // Mark work1 as failed + work1.work().setFailed(); + + // Enqueue tasks + testExecutor.execute(work1, 100); + testExecutor.execute(work2, 150); + + // Total outstanding elements must be 3 (blocker + work1 + work2) + assertEquals(3, testExecutor.elementsOutstanding()); + + // Steal work from keyGroup1. + // The first work in queue is work1, which is failed. + // It should be dropped, its handle closed, and work2 should be returned. + try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); + assertNotNull(stolen); + assertEquals(work2, stolen); + // blocker (1) + work2 (1) = 2. work1 (1) should have been released. + assertEquals(2, testExecutor.elementsOutstanding()); + } + // work2 should also be released now because stealHandle is closed. + // blocker (1) = 1. + assertEquals(1, testExecutor.elementsOutstanding()); + + // Unblock the blocker and shut down + blockerStop.countDown(); + testExecutor.shutdown(); + } } From 3670a036aa83a3c18e18f3b9435214f3cbb0ad13 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 11 Jun 2026 13:42:25 +0000 Subject: [PATCH 2/8] address comment --- .../worker/util/BoundedQueueExecutorTest.java | 51 ++++++++++--------- 1 file changed, 27 insertions(+), 24 deletions(-) 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 9106133cec24..c39b7f3a1d4d 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 @@ -594,30 +594,33 @@ public void testPollWorkDropsFailedWork() throws Exception { // Mark work1 as failed work1.work().setFailed(); - - // Enqueue tasks - testExecutor.execute(work1, 100); - testExecutor.execute(work2, 150); - - // Total outstanding elements must be 3 (blocker + work1 + work2) - assertEquals(3, testExecutor.elementsOutstanding()); - - // Steal work from keyGroup1. - // The first work in queue is work1, which is failed. - // It should be dropped, its handle closed, and work2 should be returned. - try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); - assertNotNull(stolen); - assertEquals(work2, stolen); - // blocker (1) + work2 (1) = 2. work1 (1) should have been released. - assertEquals(2, testExecutor.elementsOutstanding()); + try { + + // Enqueue tasks + testExecutor.execute(work1, 100); + testExecutor.execute(work2, 150); + + // Total outstanding elements must be 3 (blocker + work1 + work2) + assertEquals(3, testExecutor.elementsOutstanding()); + + // Steal work from keyGroup1. + // The first work in queue is work1, which is failed. + // It should be dropped, its handle closed, and work2 should be returned. + try (BoundedQueueExecutorWorkHandleImpl stealHandle = + testExecutor.createBudgetHandle(0, 0L)) { + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); + assertNotNull(stolen); + assertEquals(work2, stolen); + // blocker (1) + work2 (1) = 2. work1 (1) should have been released. + assertEquals(2, testExecutor.elementsOutstanding()); + } + // work2 should also be released now because stealHandle is closed. + // blocker (1) = 1. + assertEquals(1, testExecutor.elementsOutstanding()); + } finally { + // Unblock the blocker and shut down + blockerStop.countDown(); + testExecutor.shutdown(); } - // work2 should also be released now because stealHandle is closed. - // blocker (1) = 1. - assertEquals(1, testExecutor.elementsOutstanding()); - - // Unblock the blocker and shut down - blockerStop.countDown(); - testExecutor.shutdown(); } } From 2ddc10e8263f1cac65d998c59c744e644a5ee1fa Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Wed, 5 Aug 2026 23:18:27 +0000 Subject: [PATCH 3/8] Revert "address comment" This reverts commit 3670a036aa83a3c18e18f3b9435214f3cbb0ad13. --- .../worker/util/BoundedQueueExecutorTest.java | 51 +++++++++---------- 1 file changed, 24 insertions(+), 27 deletions(-) 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 12ada611c25c..2b437dd7f85e 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 @@ -616,33 +616,30 @@ public void testPollWorkDropsFailedWork() throws Exception { // Mark work1 as failed work1.work().setFailed(); - try { - - // Enqueue tasks - testExecutor.execute(work1, 100); - testExecutor.execute(work2, 150); - - // Total outstanding elements must be 3 (blocker + work1 + work2) - assertEquals(3, testExecutor.elementsOutstanding()); - - // Steal work from keyGroup1. - // The first work in queue is work1, which is failed. - // It should be dropped, its handle closed, and work2 should be returned. - try (BoundedQueueExecutorWorkHandleImpl stealHandle = - testExecutor.createBudgetHandle(0, 0L)) { - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); - assertNotNull(stolen); - assertEquals(work2, stolen); - // blocker (1) + work2 (1) = 2. work1 (1) should have been released. - assertEquals(2, testExecutor.elementsOutstanding()); - } - // work2 should also be released now because stealHandle is closed. - // blocker (1) = 1. - assertEquals(1, testExecutor.elementsOutstanding()); - } finally { - // Unblock the blocker and shut down - blockerStop.countDown(); - testExecutor.shutdown(); + + // Enqueue tasks + testExecutor.execute(work1, 100); + testExecutor.execute(work2, 150); + + // Total outstanding elements must be 3 (blocker + work1 + work2) + assertEquals(3, testExecutor.elementsOutstanding()); + + // Steal work from keyGroup1. + // The first work in queue is work1, which is failed. + // It should be dropped, its handle closed, and work2 should be returned. + try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); + assertNotNull(stolen); + assertEquals(work2, stolen); + // blocker (1) + work2 (1) = 2. work1 (1) should have been released. + assertEquals(2, testExecutor.elementsOutstanding()); } + // work2 should also be released now because stealHandle is closed. + // blocker (1) = 1. + assertEquals(1, testExecutor.elementsOutstanding()); + + // Unblock the blocker and shut down + blockerStop.countDown(); + testExecutor.shutdown(); } } From fc1530284dab136a73f933cfc7f762b6db4743b1 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Wed, 5 Aug 2026 23:19:02 +0000 Subject: [PATCH 4/8] Revert "Drop failed work in BoundedQueueExecutor::pollWork" This reverts commit 2c1278636a603fb5e06a5f5bf6f179b0fbf41d7c. --- .../worker/util/BoundedQueueExecutor.java | 16 ++--- .../worker/util/BoundedQueueExecutorTest.java | 67 ------------------- 2 files changed, 5 insertions(+), 78 deletions(-) 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 6924be11f3d6..9eb9a37b1b76 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,18 +395,12 @@ BoundedQueueExecutorWorkHandleImpl createBudgetHandle(Work work, long bytes) { if (keyGroupWorkQueue == null) { return null; } - while (true) { - @Nullable QueuedWork queuedWork = keyGroupWorkQueue.pollWork(computationId, keyGroup); - if (queuedWork == null) { - return null; - } - if (queuedWork.getWork().work().isFailed()) { - queuedWork.getHandle().close(); - } else { - internalHandle.merge(queuedWork.getHandle()); - return queuedWork.getWork(); - } + @Nullable QueuedWork queuedWork = keyGroupWorkQueue.pollWork(computationId, keyGroup); + if (queuedWork == null) { + return null; } + internalHandle.merge(queuedWork.getHandle()); + return queuedWork.getWork(); } private void decrementCounters(int elements, long bytes) { 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 2b437dd7f85e..0e75fa01f4f0 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 @@ -575,71 +575,4 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { blockerStop.countDown(); testExecutor.shutdown(); } - - @Test - public void testPollWorkDropsFailedWork() throws Exception { - BoundedQueueExecutor testExecutor = - new BoundedQueueExecutor( - /* initialMaximumPoolSize= */ 1, - /* keepAliveTime= */ 60, - /* unit= */ TimeUnit.SECONDS, - /* maximumElementsOutstanding= */ 100, - /* maximumBytesOutstanding= */ 10000000, - new ThreadFactoryBuilder().setNameFormat("testStealing-%d").setDaemon(true).build(), - useFairMonitor, - /*useKeyGroupWorkQueue=*/ true); - - // Create blocker task to occupy the worker thread - CountDownLatch blockerStart = new CountDownLatch(1); - CountDownLatch blockerStop = new CountDownLatch(1); - ExecutableWork blockerWork = - createWorkWithCompIdAndKeyGroup( - "blockerComp", - DEFAULT_KEY_GROUP, - ignored -> { - blockerStart.countDown(); - try { - blockerStop.await(); - } catch (InterruptedException e) { - throw new RuntimeException(e); - } - }); - - testExecutor.execute(blockerWork, 0); - blockerStart.await(); - - Work.KeyGroup keyGroup1 = Work.KeyGroup.create(1, 1); - - // Create executable tasks - ExecutableWork work1 = createWorkWithCompIdAndKeyGroup("compA", keyGroup1, ignored -> {}); - ExecutableWork work2 = createWorkWithCompIdAndKeyGroup("compA", keyGroup1, ignored -> {}); - - // Mark work1 as failed - work1.work().setFailed(); - - // Enqueue tasks - testExecutor.execute(work1, 100); - testExecutor.execute(work2, 150); - - // Total outstanding elements must be 3 (blocker + work1 + work2) - assertEquals(3, testExecutor.elementsOutstanding()); - - // Steal work from keyGroup1. - // The first work in queue is work1, which is failed. - // It should be dropped, its handle closed, and work2 should be returned. - try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); - assertNotNull(stolen); - assertEquals(work2, stolen); - // blocker (1) + work2 (1) = 2. work1 (1) should have been released. - assertEquals(2, testExecutor.elementsOutstanding()); - } - // work2 should also be released now because stealHandle is closed. - // blocker (1) = 1. - assertEquals(1, testExecutor.elementsOutstanding()); - - // Unblock the blocker and shut down - blockerStop.countDown(); - testExecutor.shutdown(); - } } From 7938f7f5aaca3ed85d11e10f6211c769733fff60 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Wed, 5 Aug 2026 23:43:20 +0000 Subject: [PATCH 5/8] [Dataflow Streaming] Remove finalizeCommits from processWork --- .../windmill/work/processing/StreamingWorkScheduler.java | 4 ---- 1 file changed, 4 deletions(-) diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java index 9e8265e509af..05a9ad82f182 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java @@ -232,10 +232,6 @@ private void processWork( KeyTransitionListener keyTransitionListener = createKeyTransitionListener(); keyTransitionListener.onKeyTransition(null, work); - // Before any processing starts, call any pending OnCommit callbacks. Nothing that requires - // cleanup should be done before this, since we might exit early here. - commitFinalizer.finalizeCommits(workItem.getSourceState().getFinalizeIdsList()); - if (workItem.getSourceState().getOnlyFinalize()) { handleOnlyFinalize(computationState, work, workItem); return; From fce54b05891828c8a321abce59d7784da4feaf92 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 6 Aug 2026 01:26:53 +0000 Subject: [PATCH 6/8] Plumb ComputationState to ProcessingContext --- .../worker/StreamingDataflowWorker.java | 34 +++++++++--------- .../worker/streaming/ActiveWorkState.java | 5 ++- .../worker/streaming/ComputationState.java | 2 +- .../dataflow/worker/streaming/Work.java | 26 ++++++++++---- .../FanOutStreamingEngineWorkerHarness.java | 23 +++++++----- .../harness/SingleSourceWorkerHarness.java | 4 +-- .../harness/WindmillStreamSender.java | 14 +++++--- .../client/grpc/GrpcDirectGetWorkStream.java | 35 ++++++++++++++----- .../grpc/GrpcWindmillStreamFactory.java | 8 +++-- .../windmill/work/WorkItemScheduler.java | 3 ++ .../worker/StreamingDataflowWorkerTest.java | 13 +++++-- .../StreamingModeExecutionContextTest.java | 12 ++++++- .../WindmillReaderIteratorBaseTest.java | 9 ++++- .../worker/WindowingWindmillReaderTest.java | 12 ++++++- .../worker/WorkerCustomSourcesTest.java | 14 ++++++-- .../worker/streaming/ActiveWorkStateTest.java | 11 +++++- .../streaming/ComputationStateCacheTest.java | 8 ++++- .../streaming/ComputationStateTest.java | 2 +- .../dataflow/worker/streaming/WorkTest.java | 9 ++++- ...anOutStreamingEngineWorkerHarnessTest.java | 14 +++++--- .../harness/WindmillStreamSenderTest.java | 23 ++++++++---- .../worker/util/BoundedQueueExecutorTest.java | 13 ++++++- .../worker/util/KeyGroupWorkQueueTest.java | 10 +++++- .../StreamingApplianceWorkCommitterTest.java | 12 +++++-- .../StreamingEngineWorkCommitterTest.java | 12 +++++-- .../grpc/GrpcDirectGetWorkStreamTest.java | 20 +++++++---- .../failures/WorkFailureProcessorTest.java | 10 +++++- .../work/refresh/ActiveWorkRefresherTest.java | 12 ++++++- 28 files changed, 281 insertions(+), 89 deletions(-) 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/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..ac7e257f2584 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( + 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(metadata.computationId()), + 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..9b7b27965a3c 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 @@ -382,7 +382,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()), @@ -391,6 +394,12 @@ computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.clas }); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private byte[] intervalWindowBytes(IntervalWindow window) throws Exception { return CoderUtils.encodeToByteArray( DEFAULT_WINDOW_COLLECTION_CODER, Collections.singletonList(window)); @@ -3966,7 +3975,7 @@ public void testLatencyAttributionProtobufsPopulated() { 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/StreamingModeExecutionContextTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java index c5efcea4e47c..d49117ecf21d 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 @@ -61,6 +61,7 @@ import org.apache.beam.runners.dataflow.worker.profiler.ScopedProfiler.NoopProfileScope; import org.apache.beam.runners.dataflow.worker.profiler.ScopedProfiler.ProfileScope; 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; @@ -169,13 +170,22 @@ public void setUp() { executionContext = createExecutionContext(options, globalConfigHandle); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createMockWork(Windmill.WorkItem workItem, Watermarks watermarks) { return Work.create( workItem, 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..e8d735728527 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 @@ -30,6 +30,7 @@ import java.util.Arrays; import java.util.List; import java.util.concurrent.ThreadLocalRandom; +import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.sdk.coders.CoderException; @@ -253,6 +254,12 @@ private void testForMessageBundleCounts(boolean skipErrors, int... messageBundle } } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createMockWork(Windmill.WorkItem workItem) { return Work.create( workItem, @@ -261,7 +268,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..0bcf54301f99 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 @@ -29,6 +29,7 @@ import java.io.IOException; import java.util.List; import org.apache.beam.runners.core.KeyedWorkItem; +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.util.common.worker.NativeReader; @@ -85,13 +86,22 @@ public void setUp() { coder, mockContext, ValueProvider.StaticValueProvider.of(false)); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createMockWork(Windmill.WorkItem workItem) { return Work.create( 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..340a1e06d016 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 @@ -89,6 +89,7 @@ import org.apache.beam.runners.dataflow.worker.counters.CounterSet; import org.apache.beam.runners.dataflow.worker.counters.NameContext; import org.apache.beam.runners.dataflow.worker.profiler.ScopedProfiler.NoopProfileScope; +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.streaming.config.FixedGlobalConfigHandle; @@ -201,13 +202,22 @@ public void testSplitAndReadBundlesBack() throws Exception { } } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createMockWork(Windmill.WorkItem workItem, Watermarks watermarks) { return Work.create( workItem, 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 +1056,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..6ddf1be99565 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 @@ -92,9 +92,18 @@ private static ExecutableWork expiredWork(Windmill.WorkItem workItem) { (work, handle) -> {}); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + 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..57ee5db9a4d2 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 @@ -58,6 +58,12 @@ public class ComputationStateCacheTest { private final ComputationConfig.Fetcher configFetcher = mock(ComputationConfig.Fetcher.class); private ComputationStateCache computationStateCache; + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static ExecutableWork createWork(ShardedKey shardedKey, long workToken, long cacheToken) { WorkItem workItem = WorkItem.newBuilder() @@ -72,7 +78,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/WorkTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/WorkTest.java index 61e52ddd61bd..3b96b395f15d 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 @@ -20,6 +20,7 @@ import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; @@ -39,6 +40,12 @@ @RunWith(JUnit4.class) public class WorkTest { + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createTestWork() { Windmill.WorkItem workItem = Windmill.WorkItem.newBuilder() @@ -51,7 +58,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..38d330e94568 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 @@ -25,6 +25,7 @@ import static org.junit.Assert.assertThat; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import java.util.Arrays; import java.util.Collection; @@ -34,6 +35,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 +92,12 @@ private static ExecutableWork createWorkWithCompIdAndKeyGroup( computationId, keyGroup, (work, handle) -> executeWorkFn.accept(work)); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static ExecutableWork createWorkWithHandle( String computationId, Work.KeyGroup keyGroup, @@ -112,7 +120,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()), 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..9699ad493124 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 @@ -24,6 +24,7 @@ import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import java.lang.Thread.State; import java.util.ArrayList; @@ -37,6 +38,7 @@ import java.util.concurrent.Future; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; +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; @@ -82,6 +84,12 @@ public void setUp() { private static final Work.KeyGroup TEST_KEY_GROUP = Work.KeyGroup.create(1, 2); + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private QueuedWork createQueuedWork(String computationId, long workBytes) { return createQueuedWork(computationId, TEST_KEY_GROUP, workBytes); } @@ -109,7 +117,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..10b5f4e094de 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 @@ -20,6 +20,7 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertNotNull; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.MapTask; import com.google.common.truth.Correspondence; @@ -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; @@ -56,6 +58,12 @@ public class StreamingApplianceWorkCommitterTest { private FakeWindmillServer fakeWindmillServer; private StreamingApplianceWorkCommitter workCommitter; + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createMockWork(long workToken) { WorkItem workItem = WorkItem.newBuilder() @@ -69,7 +77,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 +94,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..e01f9aa30a7b 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 @@ -23,6 +23,7 @@ import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.MapTask; import java.io.IOException; @@ -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; @@ -105,6 +107,12 @@ private static void waitForExpectedSetSize(Set s, int expectedSize) { assertThat(s).hasSize(expectedSize); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static Work createMockWork(long workToken) { WorkItem workItem = WorkItem.newBuilder() @@ -118,7 +126,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 +143,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..a5a6cd876133 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 @@ -20,6 +20,7 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertThrows; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import java.util.HashSet; import java.util.List; @@ -29,6 +30,7 @@ import java.util.concurrent.TimeUnit; import java.util.function.Consumer; import java.util.function.Supplier; +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; @@ -85,6 +87,12 @@ private static FailureTracker streamingApplianceFailureReporter(boolean isWorkFa ignored -> Windmill.ReportStatsResponse.newBuilder().setFailed(isWorkFailed).build()); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private static ExecutableWork createWork(Supplier clock, Consumer processWorkFn) { WorkItem workItem = WorkItem.newBuilder().setKey(ByteString.EMPTY).setWorkToken(1L).build(); return ExecutableWork.create( @@ -93,7 +101,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..66ada9862432 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 @@ -26,6 +26,7 @@ import static org.mockito.Mockito.spy; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.MapTask; import com.google.common.truth.Correspondence; @@ -121,6 +122,12 @@ private ExecutableWork createOldWork(int workIds, Consumer processWork) { return createOldWork(shardedKey, workIds, processWork); } + private static ComputationState createMockComputationState(String computationId) { + ComputationState computationState = mock(ComputationState.class); + when(computationState.getComputationId()).thenReturn(computationId); + return computationState; + } + private ExecutableWork createOldWork( ShardedKey shardedKey, int workIds, Consumer processWork) { WorkItem workItem = @@ -136,7 +143,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()), From 840ecdad63ad947dd00c3a690e4db2f710a7d639 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 6 Aug 2026 01:52:25 +0000 Subject: [PATCH 7/8] Drop failed workitems during pollwork --- .../worker/util/BoundedQueueExecutor.java | 19 ++- .../client/grpc/GrpcDirectGetWorkStream.java | 16 +- .../worker/StreamingDataflowWorkerTest.java | 152 +++++++++++++++++- .../worker/WorkerCustomSourcesTest.java | 1 + .../worker/streaming/ActiveWorkStateTest.java | 1 + .../worker/util/BoundedQueueExecutorTest.java | 109 +++++++++++++ 6 files changed, 284 insertions(+), 14 deletions(-) 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 ac7e257f2584..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 @@ -288,14 +288,14 @@ private void consumeAssembledWorkItem(AssembledWorkItem assembledWorkItem) { 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()); + 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()); } 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 9b7b27965a3c..8ffaaa6338a3 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 @@ -1533,13 +1533,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 = @@ -4846,6 +4977,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/WorkerCustomSourcesTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java index 340a1e06d016..8fd2a1411df0 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 @@ -50,6 +50,7 @@ import static org.junit.Assert.fail; import static org.junit.internal.matchers.ThrowableMessageMatcher.hasMessage; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.ApproximateReportedProgress; import com.google.api.services.dataflow.model.DataflowPackage; 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 6ddf1be99565..593d83a16b6b 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 @@ -26,6 +26,7 @@ import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; import java.util.Arrays; import java.util.Collections; 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 38d330e94568..8be24b94b2f2 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 @@ -25,6 +25,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 static org.mockito.Mockito.when; import java.util.Arrays; @@ -98,6 +99,39 @@ private static ComputationState createMockComputationState(String computationId) return computationState; } + 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, @@ -586,4 +620,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(); + } } From bbda0333ce4e7e51859103cd80652ea4e0c0c6b8 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 6 Aug 2026 04:53:56 +0000 Subject: [PATCH 8/8] Improve tests --- .../worker/StreamingDataflowWorkerTest.java | 7 +--- .../StreamingModeExecutionContextTest.java | 8 +---- .../WindmillReaderIteratorBaseTest.java | 8 +---- .../worker/WindowingWindmillReaderTest.java | 8 +---- .../worker/WorkerCustomSourcesTest.java | 9 +---- .../worker/streaming/ActiveWorkStateTest.java | 8 +---- .../streaming/ComputationStateCacheTest.java | 7 +--- .../streaming/ComputationStateTestUtils.java | 33 +++++++++++++++++++ .../dataflow/worker/streaming/WorkTest.java | 8 +---- .../worker/util/BoundedQueueExecutorTest.java | 8 +---- .../worker/util/KeyGroupWorkQueueTest.java | 9 +---- .../StreamingApplianceWorkCommitterTest.java | 8 +---- .../StreamingEngineWorkCommitterTest.java | 8 +---- .../failures/WorkFailureProcessorTest.java | 9 +---- .../work/refresh/ActiveWorkRefresherTest.java | 8 +---- 15 files changed, 47 insertions(+), 99 deletions(-) create mode 100644 runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTestUtils.java 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 8ffaaa6338a3..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; @@ -394,12 +395,6 @@ private static ExecutableWork createMockWork( }); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private byte[] intervalWindowBytes(IntervalWindow window) throws Exception { return CoderUtils.encodeToByteArray( DEFAULT_WINDOW_COLLECTION_CODER, Collections.singletonList(window)); 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 d49117ecf21d..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; @@ -61,7 +62,6 @@ import org.apache.beam.runners.dataflow.worker.profiler.ScopedProfiler.NoopProfileScope; import org.apache.beam.runners.dataflow.worker.profiler.ScopedProfiler.ProfileScope; 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; @@ -170,12 +170,6 @@ public void setUp() { executionContext = createExecutionContext(options, globalConfigHandle); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createMockWork(Windmill.WorkItem workItem, Watermarks watermarks) { return Work.create( workItem, 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 e8d735728527..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; @@ -30,7 +31,6 @@ import java.util.Arrays; import java.util.List; import java.util.concurrent.ThreadLocalRandom; -import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.sdk.coders.CoderException; @@ -254,12 +254,6 @@ private void testForMessageBundleCounts(boolean skipErrors, int... messageBundle } } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createMockWork(Windmill.WorkItem workItem) { return Work.create( workItem, 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 0bcf54301f99..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; @@ -29,7 +30,6 @@ import java.io.IOException; import java.util.List; import org.apache.beam.runners.core.KeyedWorkItem; -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.util.common.worker.NativeReader; @@ -86,12 +86,6 @@ public void setUp() { coder, mockContext, ValueProvider.StaticValueProvider.of(false)); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createMockWork(Windmill.WorkItem workItem) { return Work.create( workItem, 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 8fd2a1411df0..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; @@ -50,7 +51,6 @@ import static org.junit.Assert.fail; import static org.junit.internal.matchers.ThrowableMessageMatcher.hasMessage; import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.ApproximateReportedProgress; import com.google.api.services.dataflow.model.DataflowPackage; @@ -90,7 +90,6 @@ import org.apache.beam.runners.dataflow.worker.counters.CounterSet; import org.apache.beam.runners.dataflow.worker.counters.NameContext; import org.apache.beam.runners.dataflow.worker.profiler.ScopedProfiler.NoopProfileScope; -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.streaming.config.FixedGlobalConfigHandle; @@ -203,12 +202,6 @@ public void testSplitAndReadBundlesBack() throws Exception { } } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createMockWork(Windmill.WorkItem workItem, Watermarks watermarks) { return Work.create( workItem, 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 593d83a16b6b..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; @@ -26,7 +27,6 @@ import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.when; import java.util.Arrays; import java.util.Collections; @@ -93,12 +93,6 @@ private static ExecutableWork expiredWork(Windmill.WorkItem workItem) { (work, handle) -> {}); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work.ProcessingContext createWorkProcessingContext() { return Work.createProcessingContext( createMockComputationState("computationId"), 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 57ee5db9a4d2..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; @@ -58,12 +59,6 @@ public class ComputationStateCacheTest { private final ComputationConfig.Fetcher configFetcher = mock(ComputationConfig.Fetcher.class); private ComputationStateCache computationStateCache; - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static ExecutableWork createWork(ShardedKey shardedKey, long workToken, long cacheToken) { WorkItem workItem = WorkItem.newBuilder() 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 3b96b395f15d..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,10 +17,10 @@ */ 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; -import static org.mockito.Mockito.when; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; @@ -40,12 +40,6 @@ @RunWith(JUnit4.class) public class WorkTest { - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createTestWork() { Windmill.WorkItem workItem = Windmill.WorkItem.newBuilder() 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 8be24b94b2f2..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; @@ -26,7 +27,6 @@ import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.when; import java.util.Arrays; import java.util.Collection; @@ -93,12 +93,6 @@ private static ExecutableWork createWorkWithCompIdAndKeyGroup( computationId, keyGroup, (work, handle) -> executeWorkFn.accept(work)); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static ExecutableWork createWorkWithComputationStateAndKeyGroup( ComputationState computationState, Work.KeyGroup keyGroup, 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 9699ad493124..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; @@ -24,7 +25,6 @@ import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; import java.lang.Thread.State; import java.util.ArrayList; @@ -38,7 +38,6 @@ import java.util.concurrent.Future; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; -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; @@ -84,12 +83,6 @@ public void setUp() { private static final Work.KeyGroup TEST_KEY_GROUP = Work.KeyGroup.create(1, 2); - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private QueuedWork createQueuedWork(String computationId, long workBytes) { return createQueuedWork(computationId, TEST_KEY_GROUP, workBytes); } 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 10b5f4e094de..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,9 +18,9 @@ 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; -import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.MapTask; import com.google.common.truth.Correspondence; @@ -58,12 +58,6 @@ public class StreamingApplianceWorkCommitterTest { private FakeWindmillServer fakeWindmillServer; private StreamingApplianceWorkCommitter workCommitter; - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createMockWork(long workToken) { WorkItem workItem = WorkItem.newBuilder() 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 e01f9aa30a7b..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,12 +18,12 @@ 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; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.MapTask; import java.io.IOException; @@ -107,12 +107,6 @@ private static void waitForExpectedSetSize(Set s, int expectedSize) { assertThat(s).hasSize(expectedSize); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static Work createMockWork(long workToken) { WorkItem workItem = WorkItem.newBuilder() 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 a5a6cd876133..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,9 +18,9 @@ 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; -import static org.mockito.Mockito.when; import java.util.HashSet; import java.util.List; @@ -30,7 +30,6 @@ import java.util.concurrent.TimeUnit; import java.util.function.Consumer; import java.util.function.Supplier; -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; @@ -87,12 +86,6 @@ private static FailureTracker streamingApplianceFailureReporter(boolean isWorkFa ignored -> Windmill.ReportStatsResponse.newBuilder().setFailed(isWorkFailed).build()); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private static ExecutableWork createWork(Supplier clock, Consumer processWorkFn) { WorkItem workItem = WorkItem.newBuilder().setKey(ByteString.EMPTY).setWorkToken(1L).build(); return ExecutableWork.create( 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 66ada9862432..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; @@ -26,7 +27,6 @@ import static org.mockito.Mockito.spy; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.when; import com.google.api.services.dataflow.model.MapTask; import com.google.common.truth.Correspondence; @@ -122,12 +122,6 @@ private ExecutableWork createOldWork(int workIds, Consumer processWork) { return createOldWork(shardedKey, workIds, processWork); } - private static ComputationState createMockComputationState(String computationId) { - ComputationState computationState = mock(ComputationState.class); - when(computationState.getComputationId()).thenReturn(computationId); - return computationState; - } - private ExecutableWork createOldWork( ShardedKey shardedKey, int workIds, Consumer processWork) { WorkItem workItem =