From 2c1278636a603fb5e06a5f5bf6f179b0fbf41d7c Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 11 Jun 2026 13:27:09 +0000 Subject: [PATCH 01/19] 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 02/19] 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 03/19] 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 04/19] 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 05/19] [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 06/19] 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 07/19] 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 08/19] 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 = From 88b7b18d0a026988e3050fd34ab282eac428d8fd Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Tue, 11 Aug 2026 08:00:46 +0000 Subject: [PATCH 09/19] address comments --- .../worker/StreamingDataflowWorker.java | 34 +++---- .../worker/StreamingModeExecutionContext.java | 17 ++-- .../worker/streaming/ActiveWorkState.java | 5 +- .../worker/streaming/ComputationState.java | 2 +- .../streaming/ComputationWorkExecutor.java | 6 +- .../worker/streaming/FailedWorkHandler.java} | 17 +--- .../dataflow/worker/streaming/Work.java | 26 ++---- .../FanOutStreamingEngineWorkerHarness.java | 23 ++--- .../harness/SingleSourceWorkerHarness.java | 4 +- .../harness/WindmillStreamSender.java | 14 +-- .../worker/util/BoundedQueueExecutor.java | 11 ++- .../client/grpc/GrpcDirectGetWorkStream.java | 47 +++------- .../grpc/GrpcWindmillStreamFactory.java | 8 +- .../windmill/work/WorkItemScheduler.java | 3 - .../processing/StreamingWorkScheduler.java | 21 +++-- .../failures/WorkFailureProcessor.java | 6 +- .../worker/StreamingDataflowWorkerTest.java | 8 +- .../StreamingModeExecutionContextTest.java | 88 +++++++++++++++---- .../WindmillReaderIteratorBaseTest.java | 3 +- .../worker/WindowingWindmillReaderTest.java | 6 +- .../worker/WorkerCustomSourcesTest.java | 17 ++-- .../worker/streaming/ActiveWorkStateTest.java | 6 +- .../streaming/ComputationStateCacheTest.java | 3 +- .../streaming/ComputationStateTest.java | 2 +- .../dataflow/worker/streaming/WorkTest.java | 3 +- ...anOutStreamingEngineWorkerHarnessTest.java | 14 +-- .../harness/WindmillStreamSenderTest.java | 23 ++--- .../worker/util/BoundedQueueExecutorTest.java | 46 ++++------ .../worker/util/KeyGroupWorkQueueTest.java | 3 +- .../StreamingApplianceWorkCommitterTest.java | 6 +- .../StreamingEngineWorkCommitterTest.java | 6 +- .../grpc/GrpcDirectGetWorkStreamTest.java | 20 ++--- .../failures/WorkFailureProcessorTest.java | 3 +- .../work/refresh/ActiveWorkRefresherTest.java | 6 +- 34 files changed, 229 insertions(+), 278 deletions(-) rename runners/google-cloud-dataflow-java/worker/src/{test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTestUtils.java => main/java/org/apache/beam/runners/dataflow/worker/streaming/FailedWorkHandler.java} (62%) 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 64c7543b6ecb..2339430464c7 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,25 +405,28 @@ private StreamingWorkerHarnessFactoryOutput createFanOutStreamingEngineWorkerHar .setBytes(MAX_GET_WORK_FETCH_BYTES) .build(), windmillStreamFactory, - (computationState, - workItem, + (workItem, serializedWorkItemSize, watermarks, processingContext, drainMode, appliedFinalizeIds, - getWorkStreamLatencies) -> { - memoryMonitor.waitForResources("GetWork"); - streamingWorkScheduler.queueAppliedFinalizeIds(appliedFinalizeIds); - streamingWorkScheduler.scheduleWork( - computationState, - workItem, - serializedWorkItemSize, - watermarks, - processingContext, - drainMode, - getWorkStreamLatencies); - }, + getWorkStreamLatencies) -> + checkNotNull(computationStateCache) + .get(processingContext.computationId()) + .ifPresent( + computationState -> { + memoryMonitor.waitForResources("GetWork"); + streamingWorkScheduler.queueAppliedFinalizeIds(appliedFinalizeIds); + streamingWorkScheduler.scheduleWork( + computationState, + workItem, + serializedWorkItemSize, + watermarks, + processingContext, + drainMode, + getWorkStreamLatencies); + }), ChannelCachingRemoteStubFactory.create(options.getGcpCredential(), channelCache), GetWorkBudgetDistributors.distributeEvenly(), checkNotNull(dispatcherClient), @@ -438,8 +441,7 @@ private StreamingWorkerHarnessFactoryOutput createFanOutStreamingEngineWorkerHar .setCommitWorkStreamFactory( () -> CloseableStream.create(commitWorkStream, () -> {})) .build(), - getDataMetricTracker, - checkNotNull(this.computationStateCache)::get); + getDataMetricTracker); ChannelzServlet channelzServlet = createChannelzServlet( options, fanOutStreamingEngineWorkerHarness::currentWindmillEndpoints); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java index d577b8614078..cad52013415c 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java @@ -53,6 +53,7 @@ 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.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; import org.apache.beam.runners.dataflow.worker.streaming.KeyCommitTooLargeException; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; @@ -172,10 +173,7 @@ public class StreamingModeExecutionContext private @Nullable WorkExecutor workExecutor; private boolean finishKeyCalled = false; - @SuppressWarnings("UnusedVariable") private @Nullable BoundedQueueExecutor workQueueExecutor; - - @SuppressWarnings("UnusedVariable") private @Nullable BoundedQueueExecutorWorkHandle budgetHandle; private final HotKeyLogger hotKeyLogger; @@ -192,8 +190,8 @@ public interface KeyTransitionListener { void onKeyTransition(@Nullable Work oldWork, Work newWork); } - @SuppressWarnings("UnusedVariable") private @Nullable KeyTransitionListener keyTransitionListener; + private @Nullable FailedWorkHandler onFailedWorkHandler; private List executedWorks = Collections.emptyList(); private List outputBuilders = Collections.emptyList(); @@ -335,6 +333,7 @@ public void reset() { this.workQueueExecutor = null; this.budgetHandle = null; this.keyTransitionListener = null; + this.onFailedWorkHandler = null; this.work = null; this.key = null; this.outputBuilder = null; @@ -350,7 +349,8 @@ public void start( BoundedQueueExecutor workQueueExecutor, BoundedQueueExecutorWorkHandle budgetHandle, @Nullable Coder keyCoder, - KeyTransitionListener keyTransitionListener) + KeyTransitionListener keyTransitionListener, + @Nullable FailedWorkHandler onFailedWorkHandler) throws CoderException { reset(); this.executedWorks = new ArrayList<>(); @@ -361,6 +361,7 @@ public void start( this.workQueueExecutor = workQueueExecutor; this.budgetHandle = budgetHandle; this.keyTransitionListener = keyTransitionListener; + this.onFailedWorkHandler = onFailedWorkHandler; this.workItemsPolled = 1; this.bundleStartTimeNanos = System.nanoTime(); @@ -779,7 +780,11 @@ public boolean advance() throws CoderException { @Nullable ExecutableWork additionalWork = - executor.pollWork(computationId, activeWork.getKeyGroup(), handle); + executor.pollWork( + computationId, + activeWork.getKeyGroup(), + handle, + checkStateNotNull(onFailedWorkHandler)); if (additionalWork != null) { flushStateInternal(); Work newWork = additionalWork.work(); 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 519b2b2948a1..de4082581293 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,17 +28,18 @@ 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; @@ -77,7 +78,7 @@ public final class ActiveWorkState { private ActiveWorkState( Map> activeWork, - WindmillStateCache.ForComputation computationStateCache) { + 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 a03091824104..5e850d4312ea 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,6 +22,7 @@ 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; @@ -29,7 +30,6 @@ 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/ComputationWorkExecutor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationWorkExecutor.java index 9391b842f038..dabf72ba4eae 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationWorkExecutor.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationWorkExecutor.java @@ -65,7 +65,8 @@ public final StreamingModeExecutionContext executeWork( Work work, BoundedQueueExecutor workQueueExecutor, BoundedQueueExecutorWorkHandle budgetHandle, - KeyTransitionListener keyTransitionListener) + KeyTransitionListener keyTransitionListener, + FailedWorkHandler onFailedWorkHandler) throws Exception { context() .start( @@ -74,7 +75,8 @@ public final StreamingModeExecutionContext executeWork( workQueueExecutor, budgetHandle, keyCoder().orElse(null), - keyTransitionListener); + keyTransitionListener, + onFailedWorkHandler); workExecutor().execute(); return context(); } 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/main/java/org/apache/beam/runners/dataflow/worker/streaming/FailedWorkHandler.java similarity index 62% rename from runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTestUtils.java rename to runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/FailedWorkHandler.java index bfe3e0c87a21..683ecf5d6600 100644 --- 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/main/java/org/apache/beam/runners/dataflow/worker/streaming/FailedWorkHandler.java @@ -17,17 +17,8 @@ */ 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; - } +/** Handler for failed {@link Work}. */ +@FunctionalInterface +public interface FailedWorkHandler { + void onFailedWork(Work work); } 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 5759be7cecf6..4541a1c313a2 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,26 +144,22 @@ public static Work create( } public static ProcessingContext createProcessingContext( - ComputationState computationState, + String computationId, GetDataClient getDataClient, Consumer workCommitter, HeartbeatSender heartbeatSender) { return ProcessingContext.create( - computationState, - getDataClient, - workCommitter, - heartbeatSender, - /* backendWorkerToken= */ ""); + computationId, getDataClient, workCommitter, heartbeatSender, /* backendWorkerToken= */ ""); } public static ProcessingContext createProcessingContext( - ComputationState computationState, + String computationId, GetDataClient getDataClient, Consumer workCommitter, HeartbeatSender heartbeatSender, String backendWorkerToken) { return ProcessingContext.create( - computationState, getDataClient, workCommitter, heartbeatSender, backendWorkerToken); + computationId, getDataClient, workCommitter, heartbeatSender, backendWorkerToken); } private static LatencyAttribution.Builder createLatencyAttributionWithActiveLatencyBreakdown( @@ -211,10 +207,6 @@ public long getSerializedWorkItemSize() { return serializedWorkItemSize; } - public ComputationState getComputationState() { - return processingContext.computationState(); - } - public String getComputationId() { return processingContext.computationId(); } @@ -465,21 +457,17 @@ public KeyGroup getKeyGroup() { public abstract static class ProcessingContext { private static ProcessingContext create( - ComputationState computationState, + String computationId, GetDataClient getDataClient, Consumer workCommitter, HeartbeatSender heartbeatSender, String backendWorkerToken) { return new AutoValue_Work_ProcessingContext( - computationState, getDataClient, heartbeatSender, workCommitter, backendWorkerToken); + computationId, getDataClient, heartbeatSender, workCommitter, backendWorkerToken); } /** Computation that the {@link Work} belongs to. */ - public abstract ComputationState computationState(); - - public String computationId() { - return computationState().getComputationId(); - } + public abstract String computationId(); /** 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 a81e7537d077..f3262c17b698 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,7 +40,6 @@ 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; @@ -97,7 +96,6 @@ 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(); @@ -133,8 +131,7 @@ private FanOutStreamingEngineWorkerHarness( GrpcDispatcherClient dispatcherClient, Function workCommitterFactory, ThrottlingGetDataMetricTracker getDataMetricTracker, - ExecutorService workerMetadataConsumer, - Function> computationStateFetcher) { + ExecutorService workerMetadataConsumer) { this.jobHeader = jobHeader; this.getDataMetricTracker = getDataMetricTracker; this.started = false; @@ -153,7 +150,6 @@ private FanOutStreamingEngineWorkerHarness( this.activeMetadataType = WindmillEndpoints.Type.UNKNOWN; this.pendingMetadataType = WindmillEndpoints.Type.UNKNOWN; this.workCommitterFactory = workCommitterFactory; - this.computationStateFetcher = computationStateFetcher; } /** @@ -170,8 +166,7 @@ public static FanOutStreamingEngineWorkerHarness create( GetWorkBudgetDistributor getWorkBudgetDistributor, GrpcDispatcherClient dispatcherClient, Function workCommitterFactory, - ThrottlingGetDataMetricTracker getDataMetricTracker, - Function> computationStateFetcher) { + ThrottlingGetDataMetricTracker getDataMetricTracker) { return new FanOutStreamingEngineWorkerHarness( jobHeader, totalGetWorkBudget, @@ -183,8 +178,9 @@ public static FanOutStreamingEngineWorkerHarness create( workCommitterFactory, getDataMetricTracker, Executors.newSingleThreadExecutor( - new ThreadFactoryBuilder().setNameFormat(WORKER_METADATA_CONSUMER_THREAD_NAME).build()), - computationStateFetcher); + new ThreadFactoryBuilder() + .setNameFormat(WORKER_METADATA_CONSUMER_THREAD_NAME) + .build())); } @VisibleForTesting @@ -197,8 +193,7 @@ static FanOutStreamingEngineWorkerHarness forTesting( GetWorkBudgetDistributor getWorkBudgetDistributor, GrpcDispatcherClient dispatcherClient, Function workCommitterFactory, - ThrottlingGetDataMetricTracker getDataMetricTracker, - Function> computationStateFetcher) { + ThrottlingGetDataMetricTracker getDataMetricTracker) { FanOutStreamingEngineWorkerHarness fanOutStreamingEngineWorkProvider = new FanOutStreamingEngineWorkerHarness( jobHeader, @@ -215,8 +210,7 @@ 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(), - computationStateFetcher); + MoreExecutors.newDirectExecutorService()); fanOutStreamingEngineWorkProvider.start(); return fanOutStreamingEngineWorkProvider; } @@ -454,8 +448,7 @@ private WindmillStreamSender createAndStartWindmillStreamSender(Endpoint endpoin getDataStream -> StreamGetDataClient.create( getDataStream, this::getGlobalDataStream, getDataMetricTracker), - workCommitterFactory, - computationStateFetcher); + workCommitterFactory); 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 00c949009206..f41223310385 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( - computationState, + computationId, getDataClient, workCommitter::commit, heartbeatSender), @@ -250,7 +250,7 @@ private void applianceDispatchLoop(Supplier getWorkFn) workItem.getSerializedSize(), watermarks.setOutputDataWatermark(workItem.getOutputDataWatermark()).build(), Work.createProcessingContext( - computationState, getDataClient, workCommitter::commit, heartbeatSender), + computationId, 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 5abe93f234a1..d150ee6bf1d1 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,7 +19,6 @@ 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; @@ -28,7 +27,6 @@ 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; @@ -77,8 +75,7 @@ private WindmillStreamSender( GrpcWindmillStreamFactory streamingEngineStreamFactory, WorkItemScheduler workItemScheduler, Function getDataClientFactory, - Function workCommitterFactory, - Function> computationStateFetcher) { + Function workCommitterFactory) { this.started = new AtomicBoolean(false); this.getWorkBudget = getWorkBudget; @@ -94,8 +91,7 @@ private WindmillStreamSender( FixedStreamHeartbeatSender.create(getDataStream), getDataClientFactory.apply(getDataStream), workCommitter, - workItemScheduler, - computationStateFetcher); + workItemScheduler); // 3 threads, 1 for each stream type (GetWork, GetData, CommitWork). this.streamStarter = Executors.newFixedThreadPool( @@ -109,8 +105,7 @@ static WindmillStreamSender create( GrpcWindmillStreamFactory streamingEngineStreamFactory, WorkItemScheduler workItemScheduler, Function getDataClientFactory, - Function workCommitterFactory, - Function> computationStateFetcher) { + Function workCommitterFactory) { return new WindmillStreamSender( connection, getWorkRequest, @@ -118,8 +113,7 @@ static WindmillStreamSender create( streamingEngineStreamFactory, workItemScheduler, getDataClientFactory, - workCommitterFactory, - computationStateFetcher); + workCommitterFactory); } private static GetWorkRequest withRequestBudget(GetWorkRequest request, GetWorkBudget budget) { diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java index e11bb587bf0d..2dd0f971168e 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 @@ -18,6 +18,7 @@ package org.apache.beam.runners.dataflow.worker.util; import static org.apache.beam.sdk.util.Preconditions.checkArgumentNotNull; +import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull; import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkArgument; import java.util.ArrayList; @@ -31,6 +32,7 @@ import javax.annotation.concurrent.GuardedBy; import org.apache.beam.runners.dataflow.worker.streaming.BoundedQueueExecutorWorkHandle; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; import org.apache.beam.runners.dataflow.worker.streaming.Work; 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; @@ -387,10 +389,14 @@ BoundedQueueExecutorWorkHandleImpl createBudgetHandle(Work work, long bytes) { } public @Nullable ExecutableWork pollWork( - String computationId, Work.KeyGroup keyGroup, BoundedQueueExecutorWorkHandle handle) { + String computationId, + Work.KeyGroup keyGroup, + BoundedQueueExecutorWorkHandle handle, + FailedWorkHandler onFailedWorkHandler) { checkArgument( computationId != null && keyGroup != null && !keyGroup.equals(Work.KeyGroup.DEFAULT)); checkArgument(handle instanceof BoundedQueueExecutorWorkHandleImpl); + checkStateNotNull(onFailedWorkHandler); BoundedQueueExecutorWorkHandleImpl internalHandle = (BoundedQueueExecutorWorkHandleImpl) handle; if (keyGroupWorkQueue == null) { return null; @@ -403,8 +409,7 @@ BoundedQueueExecutorWorkHandleImpl createBudgetHandle(Work work, long bytes) { Work work = queuedWork.getWork().work(); if (work.isFailed()) { queuedWork.getHandle().close(); - work.getComputationState() - .completeWorkAndScheduleNextWorkForKey(work.getShardedKey(), work.id()); + onFailedWorkHandler.onFailedWork(work); continue; } internalHandle.merge(queuedWork.getHandle()); 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 546a957a4861..de8ebf14b709 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,7 +21,6 @@ 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; @@ -30,7 +29,6 @@ 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; @@ -83,7 +81,6 @@ 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; @@ -105,8 +102,7 @@ private GrpcDirectGetWorkStream( WorkCommitter workCommitter, WorkItemScheduler workItemScheduler, Duration halfClosePhysicalStreamAfter, - ScheduledExecutorService executorService, - Function> computationStateFetcher) { + ScheduledExecutorService executorService) { super( LOG, startGetWorkRpcFn, @@ -122,7 +118,6 @@ private GrpcDirectGetWorkStream( this.heartbeatSender = heartbeatSender; this.workCommitter = workCommitter; this.getDataClient = getDataClient; - this.computationStateFetcher = computationStateFetcher; this.lastRequest = new AtomicReference<>(); this.budgetTracker = new GetWorkBudgetTracker( @@ -150,8 +145,7 @@ static GrpcDirectGetWorkStream create( WorkCommitter workCommitter, WorkItemScheduler workItemScheduler, Duration halfClosePhysicalStreamAfter, - ScheduledExecutorService executor, - Function> computationStateFetcher) { + ScheduledExecutorService executor) { return new GrpcDirectGetWorkStream( backendWorkerToken, startGetWorkRpcFn, @@ -166,8 +160,7 @@ static GrpcDirectGetWorkStream create( workCommitter, workItemScheduler, halfClosePhysicalStreamAfter, - executor, - computationStateFetcher); + executor); } private static Watermarks createWatermarks( @@ -280,37 +273,25 @@ protected void sendHealthCheck() throws WindmillStreamShutdownException { } private void consumeAssembledWorkItem(AssembledWorkItem assembledWorkItem) { + WorkItem workItem = assembledWorkItem.workItem(); GetWorkResponseChunkAssembler.ComputationMetadata metadata = assembledWorkItem.computationMetadata(); - Optional maybeComputationState = - computationStateFetcher.apply(metadata.computationId()); - if (maybeComputationState.isPresent()) { - ComputationState computationState = maybeComputationState.get(); - WorkItem workItem = assembledWorkItem.workItem(); - workItemScheduler.scheduleWork( - computationState, - workItem, - assembledWorkItem.bufferedSize(), - createWatermarks(workItem, metadata), - createProcessingContext(computationState), - metadata.drainMode(), - assembledWorkItem.appliedFinalizeIds(), - assembledWorkItem.latencyAttributions()); - } else { - LOG.warn("Received work for unknown computation: {}", metadata.computationId()); - } + workItemScheduler.scheduleWork( + workItem, + assembledWorkItem.bufferedSize(), + createWatermarks(workItem, metadata), + createProcessingContext(metadata.computationId()), + metadata.drainMode(), + assembledWorkItem.appliedFinalizeIds(), + assembledWorkItem.latencyAttributions()); budgetTracker.recordBudgetReceived(assembledWorkItem.bufferedSize()); GetWorkBudget extension = budgetTracker.computeBudgetExtension(); maybeSendRequestExtension(extension); } - private Work.ProcessingContext createProcessingContext(ComputationState computationState) { + private Work.ProcessingContext createProcessingContext(String computationId) { return Work.createProcessingContext( - computationState, - getDataClient, - workCommitter::commit, - heartbeatSender, - backendWorkerToken()); + computationId, 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 f4465413f43d..97ca3c4e83d7 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,7 +23,6 @@ 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; @@ -38,7 +37,6 @@ 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; @@ -289,8 +287,7 @@ public GetWorkStream createDirectGetWorkStream( HeartbeatSender heartbeatSender, GetDataClient getDataClient, WorkCommitter workCommitter, - WorkItemScheduler workItemScheduler, - Function> computationStateFetcher) { + WorkItemScheduler workItemScheduler) { return GrpcDirectGetWorkStream.create( connection.backendWorkerToken(), responseObserver -> @@ -306,8 +303,7 @@ public GetWorkStream createDirectGetWorkStream( workCommitter, workItemScheduler, directStreamingRpcPhysicalStreamHalfCloseAfter, - executorForDirectStreams(connection.backendWorkerToken(), "GetWork"), - computationStateFetcher); + executorForDirectStreams(connection.backendWorkerToken(), "GetWork")); } 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 be4c5562031f..a2dfa50a0d63 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,7 +18,6 @@ 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; @@ -33,7 +32,6 @@ 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. @@ -43,7 +41,6 @@ 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/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 299c67128caf..a3bd1403199e 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 @@ -45,6 +45,7 @@ import org.apache.beam.runners.dataflow.worker.streaming.ComputationState; import org.apache.beam.runners.dataflow.worker.streaming.ComputationWorkExecutor; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; import org.apache.beam.runners.dataflow.worker.streaming.StageInfo; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; @@ -319,8 +320,11 @@ private ExecuteWorkResult executeWork( try { StreamingModeExecutionContext context = computationWorkExecutor.context(); + FailedWorkHandler onFailedWorkHandler = getFailedWorkHandler(computationState); + // Blocks while executing work. - computationWorkExecutor.executeWork(work, workExecutor, handle, keyTransitionListener); + computationWorkExecutor.executeWork( + work, workExecutor, handle, keyTransitionListener, onFailedWorkHandler); List workBatch; List workItemCommits; @@ -462,13 +466,10 @@ private void handleProcessWorkFailure( ExecutableWork.create(w, (retry, h) -> processWork(computationState, retry, h))); } + FailedWorkHandler onFailedWorkHandler = getFailedWorkHandler(computationState); + workFailureProcessor.logAndProcessFailureBatch( - computationId, - executableWorks, - t, - invalidWork -> - computationState.completeWorkAndScheduleNextWorkForKey( - invalidWork.getShardedKey(), invalidWork.id())); + computationId, executableWorks, t, onFailedWorkHandler); } catch (OutOfMemoryError oom) { throw oom; } catch (Throwable t2) { @@ -477,6 +478,12 @@ private void handleProcessWorkFailure( } } + private static FailedWorkHandler getFailedWorkHandler(ComputationState computationState) { + return failedWork -> + computationState.completeWorkAndScheduleNextWorkForKey( + failedWork.getShardedKey(), failedWork.id()); + } + private void recordProcessingTime( StageInfo stageInfo, List workBatch, long processingStartTimeNanos) { long processingTimeMsecs = diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java index 8af1840faf92..250cd925ab34 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java @@ -19,12 +19,12 @@ import java.util.List; import java.util.concurrent.TimeUnit; -import java.util.function.Consumer; import java.util.function.Supplier; import javax.annotation.Nullable; import javax.annotation.concurrent.ThreadSafe; import org.apache.beam.runners.dataflow.worker.status.LastExceptionDataProvider; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; import org.apache.beam.sdk.annotations.Internal; @@ -102,7 +102,7 @@ public void logAndProcessFailureBatch( String computationId, List executableWorks, Throwable t, - Consumer onInvalidWork) + FailedWorkHandler onFailedWorkHandler) throws Throwable { List worksToRetryLocally = new java.util.ArrayList<>(); @@ -111,7 +111,7 @@ public void logAndProcessFailureBatch( case DO_NOT_RETRY: // Consider the item invalid. It will eventually be retried by Windmill if it still needs // to be processed. - onInvalidWork.accept(executableWork.work()); + onFailedWorkHandler.onFailedWork(executableWork.work()); break; case RETRY_LOCALLY: // Try again after some delay and at the end of the queue to avoid a tight loop. 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 055890d0f6af..4ba855684243 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,7 +20,6 @@ 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; @@ -383,10 +382,7 @@ private static ExecutableWork createMockWork( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - createMockComputationState(computationId), - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)), + computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()), @@ -4101,7 +4097,7 @@ public void testLatencyAttributionProtobufsPopulated() { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - createMockComputationState("computationId"), + "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 5766b8196516..fd58e91432a1 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,7 +19,6 @@ 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; @@ -176,10 +175,7 @@ private static Work createMockWork(Windmill.WorkItem workItem, Watermarks waterm workItem.getSerializedSize(), watermarks, Work.createProcessingContext( - createMockComputationState(COMPUTATION_ID), - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)), + COMPUTATION_ID, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()); @@ -202,10 +198,11 @@ private void start(StreamingModeExecutionContext context, Work work, Coder ke context.start( work, workExecutor, - /* workQueueExecutor= */ null, - /* budgetHandle= */ null, + /* workQueueExecutor= */ mock(BoundedQueueExecutor.class), + /* budgetHandle= */ mock(BoundedQueueExecutorWorkHandle.class), keyCoder, - /* keyTransitionListener= */ (k, c) -> {}); + /* keyTransitionListener= */ (k, c) -> {}, + /* onFailedWorkHandler= */ ignored -> {}); } catch (CoderException e) { throw new RuntimeException(e); } @@ -550,12 +547,18 @@ public void testAdvance_success() throws Exception { workItem2, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); ExecutableWork executableWork2 = ExecutableWork.create(work2, (w, h) -> {}); - when(mockExecutor.pollWork(eq(COMPUTATION_ID), eq(work1.getKeyGroup()), eq(mockHandle))) + when(mockExecutor.pollWork(eq(COMPUTATION_ID), eq(work1.getKeyGroup()), eq(mockHandle), any())) .thenReturn(executableWork2) .thenReturn(null); executionContext.start( - work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); assertTrue(executionContext.advance()); assertEquals("key2", executionContext.getSerializedKey().toStringUtf8()); @@ -579,11 +582,17 @@ public void testAdvance_noMoreWork() throws Exception { createMockWork( workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); - when(mockExecutor.pollWork(eq(COMPUTATION_ID), eq(work1.getKeyGroup()), eq(mockHandle))) + when(mockExecutor.pollWork(eq(COMPUTATION_ID), eq(work1.getKeyGroup()), eq(mockHandle), any())) .thenReturn(null); executionContext.start( - work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); assertFalse(executionContext.advance()); } @@ -613,7 +622,14 @@ public void testAdvance_respectsMaxBatchSize() throws Exception { createMockWork( workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); - context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + context.start( + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); assertFalse(context.advance()); verifyNoInteractions(mockExecutor); @@ -644,7 +660,14 @@ public void testAdvance_respectsMaxBatchTime() throws Exception { createMockWork( workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); - context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + context.start( + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); assertFalse(context.advance()); verifyNoInteractions(mockExecutor); @@ -668,7 +691,13 @@ public void testAdvance_workFailed() throws Exception { workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); executionContext.start( - work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); work1.setFailed(); @@ -691,7 +720,13 @@ public void testAdvance_defaultKeyGroup() throws Exception { workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); executionContext.start( - work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); assertFalse(executionContext.advance()); verifyNoInteractions(mockExecutor); @@ -719,7 +754,14 @@ public void testAdvance_experimentDisabled() throws Exception { createMockWork( workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); - context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + context.start( + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); assertFalse(context.advance()); verifyNoInteractions(mockExecutor); @@ -752,11 +794,19 @@ public void testAdvance_respectsMaxBatchSinkBytes() throws Exception { createMockWork( workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); - context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + context.start( + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> {}, + ignored -> {}); context.reportBytesSinked(50); assertFalse(context.advance()); - verify(mockExecutor).pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle); + verify(mockExecutor) + .pollWork(eq(COMPUTATION_ID), eq(work1.getKeyGroup()), eq(mockHandle), any()); reset(mockExecutor); 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 7c6301fd0411..be77da540889 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,7 +17,6 @@ */ 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; @@ -262,7 +261,7 @@ private static Work createMockWork(Windmill.WorkItem workItem) { .setInputDataWatermark(new org.joda.time.Instant(1000)) .build(), Work.createProcessingContext( - createMockComputationState("computationId"), + "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 fa95e64fc90e..3c778650eb3e 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,7 +17,6 @@ */ 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; @@ -92,10 +91,7 @@ private static Work createMockWork(Windmill.WorkItem workItem) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(new Instant(1000)).build(), Work.createProcessingContext( - createMockComputationState("computationId"), - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)), + "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 2a0096b1ee18..2a741da7fd88 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,7 +26,6 @@ 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; @@ -90,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.BoundedQueueExecutorWorkHandle; 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; @@ -98,6 +98,7 @@ import org.apache.beam.runners.dataflow.worker.streaming.harness.StreamingCounters; import org.apache.beam.runners.dataflow.worker.streaming.sideinput.SideInputStateFetcherFactory; import org.apache.beam.runners.dataflow.worker.testing.TestCountingSource; +import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; import org.apache.beam.runners.dataflow.worker.util.common.worker.NativeReader; import org.apache.beam.runners.dataflow.worker.util.common.worker.NativeReader.NativeReaderIterator; import org.apache.beam.runners.dataflow.worker.util.common.worker.WorkExecutor; @@ -208,10 +209,7 @@ private static Work createMockWork(Windmill.WorkItem workItem, Watermarks waterm workItem.getSerializedSize(), watermarks, Work.createProcessingContext( - createMockComputationState(COMPUTATION_ID), - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)), + COMPUTATION_ID, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()); @@ -222,10 +220,11 @@ private void startContext(StreamingModeExecutionContext context, Work work) { context.start( work, mock(WorkExecutor.class), - /* workQueueExecutor= */ null, - /* budgetHandle= */ null, + /* workQueueExecutor= */ mock(BoundedQueueExecutor.class), + /* budgetHandle= */ mock(BoundedQueueExecutorWorkHandle.class), /* keyCoder= */ null, - /* keyTransitionListener= */ mock(KeyTransitionListener.class)); + /* keyTransitionListener= */ mock(KeyTransitionListener.class), + /* onFailedWorkHandler= */ ignored -> {}); } catch (CoderException e) { throw new RuntimeException(e); } @@ -1050,7 +1049,7 @@ public void testFailedWorkItemsAbort() throws Exception { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(new Instant(0)).build(), Work.createProcessingContext( - createMockComputationState(COMPUTATION_ID), + 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 9e65bc57119c..aa0eae0d159f 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,7 +18,6 @@ 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; @@ -95,10 +94,7 @@ private static ExecutableWork expiredWork(Windmill.WorkItem workItem) { private static Work.ProcessingContext createWorkProcessingContext() { return Work.createProcessingContext( - createMockComputationState("computationId"), - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)); + "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 0ea47e4037f7..f57e20d4b5fb 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,7 +18,6 @@ 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; @@ -73,7 +72,7 @@ private static ExecutableWork createWork(ShardedKey shardedKey, long workToken, workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - createMockComputationState("computationId"), + "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 6184560670c7..22ddc8e4de5b 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( - computationState, new FakeGetDataClient(), ignored -> {}, mockHeartbeatSender), + "computationId", 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 ad7ed86e2495..61e52ddd61bd 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,7 +17,6 @@ */ 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; @@ -52,7 +51,7 @@ private static Work createTestWork() { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - createMockComputationState("comp"), + "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 03944de29d9d..cf6df7f0e478 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,13 +32,11 @@ 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; @@ -124,8 +122,7 @@ public class FanOutStreamingEngineWorkerHarnessTest { private FanOutStreamingEngineWorkerHarness fanOutStreamingEngineWorkProvider; private static WorkItemScheduler noOpProcessWorkItemFn() { - return (computationState, - workItem, + return (workItem, serializedWorkItemSize, watermarks, processingContext, @@ -198,8 +195,7 @@ private FanOutStreamingEngineWorkerHarness newFanOutStreamingEngineWorkerHarness getWorkBudgetDistributor, dispatcherClient, ignored -> mock(WorkCommitter.class), - new ThrottlingGetDataMetricTracker(mock(MemoryMonitor.class)), - ignored -> Optional.of(mock(ComputationState.class))); + new ThrottlingGetDataMetricTracker(mock(MemoryMonitor.class))); getWorkerMetadataReady.await(); return harness; } @@ -250,8 +246,7 @@ public void testStreamsStartCorrectly() throws InterruptedException { any(), any(), any(), - eq(noOpProcessWorkItemFn()), - any()); + eq(noOpProcessWorkItemFn())); verify(streamFactory, times(1)) .createDirectGetWorkStream( any(), @@ -259,8 +254,7 @@ public void testStreamsStartCorrectly() throws InterruptedException { any(), any(), any(), - eq(noOpProcessWorkItemFn()), - any()); + eq(noOpProcessWorkItemFn())); 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 ae8c02a91278..457f75593e23 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,8 +26,6 @@ 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; @@ -66,8 +64,7 @@ public class WindmillStreamSenderTest { .build()) .build()); private final WorkItemScheduler workItemScheduler = - (computationState, - workItem, + (workItem, serializedWorkItemSize, watermarks, processingContext, @@ -118,8 +115,7 @@ public void testStartStream_startsAllStreams() { any(), any(), any(), - eq(workItemScheduler), - any()); + eq(workItemScheduler)); verify(streamFactory).createDirectGetDataStream(eq(connection)); verify(streamFactory).createDirectCommitWorkStream(eq(connection)); @@ -150,8 +146,7 @@ public void testStartStream_onlyStartsStreamsOnce() { any(), any(), any(), - eq(workItemScheduler), - any()); + eq(workItemScheduler)); verify(streamFactory, times(1)).createDirectGetDataStream(eq(connection)); verify(streamFactory, times(1)).createDirectCommitWorkStream(eq(connection)); @@ -185,8 +180,7 @@ public void testStartStream_onlyStartsStreamsOnceConcurrent() throws Interrupted any(), any(), any(), - eq(workItemScheduler), - any()); + eq(workItemScheduler)); verify(streamFactory, times(1)).createDirectGetDataStream(eq(connection)); verify(streamFactory, times(1)).createDirectCommitWorkStream(eq(connection)); @@ -209,8 +203,7 @@ public void testCloseAllStreams_closesAllStreams() { any(), any(), any(), - eq(workItemScheduler), - any())) + eq(workItemScheduler))) .thenReturn(mockGetWorkStream); when(mockStreamFactory.createDirectGetDataStream(eq(connection))).thenReturn(mockGetDataStream); @@ -247,8 +240,7 @@ public void testCloseAllStreams_doesNotStartStreamsAfterClose() { any(), any(), any(), - eq(workItemScheduler), - any())) + eq(workItemScheduler))) .thenReturn(mockGetWorkStream); when(mockStreamFactory.createDirectGetDataStream(eq(connection))).thenReturn(mockGetDataStream); @@ -297,7 +289,6 @@ private WindmillStreamSender newWindmillStreamSender( streamFactory, workItemScheduler, ignored -> mock(GetDataClient.class), - ignored -> mock(WorkCommitter.class), - ignored -> Optional.of(mock(ComputationState.class))); + ignored -> mock(WorkCommitter.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 51b13d1218fa..732aa83ee5d5 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,7 +17,6 @@ */ 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; @@ -36,8 +35,8 @@ 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.FailedWorkHandler; 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.BoundedQueueExecutor.BoundedQueueExecutorWorkHandleImpl; @@ -93,11 +92,8 @@ private static ExecutableWork createWorkWithCompIdAndKeyGroup( computationId, keyGroup, (work, handle) -> executeWorkFn.accept(work)); } - private static ExecutableWork createWorkWithComputationStateAndKeyGroup( - ComputationState computationState, - Work.KeyGroup keyGroup, - long workToken, - Consumer executeWorkFn) { + private static ExecutableWork createWorkWithCompIdAndKeyGroupAndWorkToken( + String computationId, Work.KeyGroup keyGroup, long workToken, Consumer executeWorkFn) { WorkItem workItem = WorkItem.newBuilder() .setKey(ByteString.EMPTY) @@ -116,10 +112,7 @@ private static ExecutableWork createWorkWithComputationStateAndKeyGroup( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - computationState, - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)), + computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()), @@ -148,10 +141,7 @@ private static ExecutableWork createWorkWithHandle( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - createMockComputationState(computationId), - new FakeGetDataClient(), - ignored -> {}, - mock(HeartbeatSender.class)), + computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, Instant::now, ImmutableList.of()), @@ -550,7 +540,7 @@ public void testPollWork() throws Exception { assertEquals(3, testExecutor.elementsOutstanding()); // Steal work2 using pollWork with compA and keyGroup2 - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup2, stealHandle); + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup2, stealHandle, ignored -> {}); assertNotNull(stolen); assertEquals(work2, stolen); @@ -559,7 +549,7 @@ public void testPollWork() throws Exception { targetStart.await(); // Steal work1 using pollWork with compA and keyGroup1 - ExecutableWork stolen1 = testExecutor.pollWork("compA", keyGroup1, stealHandle); + ExecutableWork stolen1 = testExecutor.pollWork("compA", keyGroup1, stealHandle, ignored -> {}); assertNotNull(stolen1); assertEquals(work1, stolen1); @@ -608,7 +598,7 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { ExecutableWork work = createWorkWithCompIdAndKeyGroup("compA", keyGroup, ignored -> {}); testExecutor.execute(work, 100); - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup, stealHandle); + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup, stealHandle, ignored -> {}); assertNull(stolen); blockerStop.countDown(); @@ -616,8 +606,7 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { } @Test - public void testPollWork_skipsFailedWorkAndCallsCompleteWorkAndScheduleNextWorkForKey() - throws Exception { + public void testPollWork_skipsFailedWorkAndCallsOnFailedWorkHandler() throws Exception { BoundedQueueExecutor testExecutor = new BoundedQueueExecutor( 1, @@ -653,12 +642,12 @@ public void testPollWork_skipsFailedWorkAndCallsCompleteWorkAndScheduleNextWorkF assertNotNull(stealHandle); Work.KeyGroup keyGroup = Work.KeyGroup.create(1, 1); - ComputationState mockCompState = createMockComputationState("compA"); + FailedWorkHandler onFailedWorkHandler = mock(FailedWorkHandler.class); ExecutableWork work1 = - createWorkWithComputationStateAndKeyGroup(mockCompState, keyGroup, 101, ignored -> {}); + createWorkWithCompIdAndKeyGroupAndWorkToken("compA", keyGroup, 101, ignored -> {}); ExecutableWork work2 = - createWorkWithComputationStateAndKeyGroup(mockCompState, keyGroup, 102, ignored -> {}); + createWorkWithCompIdAndKeyGroupAndWorkToken("compA", keyGroup, 102, ignored -> {}); // Enqueue both tasks (they will wait in the queue because the thread is blocked). testExecutor.execute(work1, 100); @@ -671,20 +660,19 @@ public void testPollWork_skipsFailedWorkAndCallsCompleteWorkAndScheduleNextWorkF 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); + // onFailedWorkHandler callback, and return work2. + ExecutableWork stolen = + testExecutor.pollWork("compA", keyGroup, stealHandle, onFailedWorkHandler); assertNotNull(stolen); assertEquals(work2, stolen); - verify(mockCompState) - .completeWorkAndScheduleNextWorkForKey(work1.work().getShardedKey(), work1.work().id()); + verify(onFailedWorkHandler).onFailedWork(work1.work()); // 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)); + assertNull(testExecutor.pollWork("compA", keyGroup, stealHandle, onFailedWorkHandler)); blockerStop.countDown(); testExecutor.shutdown(); diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueueTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueueTest.java index 6815df61cd04..77fcb0597586 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,7 +17,6 @@ */ 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; @@ -110,7 +109,7 @@ private QueuedWork createQueuedWork( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.now()).build(), Work.createProcessingContext( - createMockComputationState(computationId), + 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 c34f7b07616d..b0ca89ac4c2b 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,7 +18,6 @@ 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; @@ -38,7 +37,6 @@ 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; @@ -71,7 +69,7 @@ private static Work createMockWork(long workToken) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - createMockComputationState("computationId"), + "computationId", new FakeGetDataClient(), ignored -> { throw new UnsupportedOperationException(); @@ -88,7 +86,7 @@ private static ComputationState createComputationState(String computationId) { new MapTask().setSystemName("system").setStageName("stage"), mock(BoundedQueueExecutor.class), ImmutableMap.of(), - mock(WindmillStateCache.ForComputation.class)); + null); } 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 3961e4c02886..4c2b8a9f44fb 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,7 +18,6 @@ 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; @@ -62,7 +61,6 @@ 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; @@ -120,7 +118,7 @@ private static Work createMockWork(long workToken) { workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - createMockComputationState("computationId"), + "computationId", new FakeGetDataClient(), ignored -> { throw new UnsupportedOperationException(); @@ -137,7 +135,7 @@ private static ComputationState createComputationState(String computationId) { new MapTask().setSystemName("system").setStageName("stage"), mock(BoundedQueueExecutor.class), ImmutableMap.of(), - mock(WindmillStateCache.ForComputation.class)); + null); } 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 c53cfa3dc326..71e1300d90cf 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,13 +29,11 @@ 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; @@ -68,8 +66,7 @@ public class GrpcDirectGetWorkStreamTest { private static final WorkItemScheduler NO_OP_WORK_ITEM_SCHEDULER = - (computationState, - workItem, + (workItem, serializedWorkItemSize, watermarks, processingContext, @@ -164,8 +161,7 @@ private GrpcDirectGetWorkStream createGetWorkStream( mock(HeartbeatSender.class), mock(GetDataClient.class), mock(WorkCommitter.class), - workItemScheduler, - ignored -> Optional.of(mock(ComputationState.class))); + workItemScheduler); getWorkStream.start(); return getWorkStream; } @@ -285,8 +281,7 @@ public void testConsumedWorkItem_computesAndSendsCorrectExtension() throws Inter createGetWorkStream( testStub, initialBudget, - (computationState, - work, + (work, serializedWorkItemSize, watermarks, processingContext, @@ -336,8 +331,7 @@ public void testConsumedWorkItem_doesNotSendExtensionIfOutstandingBudgetHigh() createGetWorkStream( testStub, initialBudget, - (computationState, - work, + (work, serializedWorkItemSize, watermarks, processingContext, @@ -376,8 +370,7 @@ public void testConsumedWorkItems() throws InterruptedException { createGetWorkStream( testStub, initialBudget, - (computationState, - work, + (work, serializedWorkItemSize, watermarks, processingContext, @@ -422,8 +415,7 @@ public void testConsumedWorkItems_itemsSplitAcrossResponses() throws Interrupted createGetWorkStream( testStub, initialBudget, - (computationState, - work, + (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 ac069bfcf178..89f3aa0c0d98 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,7 +18,6 @@ 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; @@ -94,7 +93,7 @@ private static ExecutableWork createWork(Supplier clock, Consumer workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - createMockComputationState("computationId"), + "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 04d54b61aeb3..caa25bf83090 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,7 +18,6 @@ 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; @@ -137,10 +136,7 @@ private ExecutableWork createOldWork( workItem.getSerializedSize(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), Work.createProcessingContext( - createMockComputationState("computationId"), - new FakeGetDataClient(), - ignored -> {}, - heartbeatSender), + "computationId", new FakeGetDataClient(), ignored -> {}, heartbeatSender), false, ActiveWorkRefresherTest::aLongTimeAgo, ImmutableList.of()), From a02260e5121e6c016544b3a3cc0283aa8fad5ab3 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 13 Aug 2026 21:37:35 +0000 Subject: [PATCH 10/19] address comments --- .../worker/StreamingModeExecutionContext.java | 4 +-- .../worker/StreamingDataflowWorkerTest.java | 2 ++ .../StreamingModeExecutionContextTest.java | 30 ++++++++++++------ .../worker/util/BoundedQueueExecutorTest.java | 31 ++++++++++++++++--- 4 files changed, 51 insertions(+), 16 deletions(-) diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java index ab67fe5481a2..365ebbdc1f9d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java @@ -350,7 +350,7 @@ public void start( BoundedQueueExecutorWorkHandle budgetHandle, @Nullable Coder keyCoder, KeyTransitionListener keyTransitionListener, - @Nullable FailedWorkHandler onFailedWorkHandler) + FailedWorkHandler onFailedWorkHandler) throws CoderException { reset(); this.executedWorks = new ArrayList<>(); @@ -361,7 +361,7 @@ public void start( this.workQueueExecutor = workQueueExecutor; this.budgetHandle = budgetHandle; this.keyTransitionListener = keyTransitionListener; - this.onFailedWorkHandler = onFailedWorkHandler; + this.onFailedWorkHandler = checkStateNotNull(onFailedWorkHandler); this.workItemsPolled = 1; this.bundleStartTimeNanos = System.nanoTime(); 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 4ba855684243..6635bc11e79b 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 @@ -1628,6 +1628,8 @@ public void testMultiKeyCommit_queuedWorkItemFailsAndSubsequentWorkItemPickedUp( .setWorkToken(2) .setShardingKey(2) .setFailed(true); + + // Fake server processes heartbeat responses are processed synchronously in sendFailedHeartbeats server.sendFailedHeartbeats(Collections.singletonList(failedHeartbeat.build())); // Unblock key1 to allow bundle to poll key2 (token 2 -> failed, skipped) and key2 (token 3). 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 5dcc8f4861cc..93874f034ec2 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 @@ -62,6 +62,7 @@ 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.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; 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.FakeGlobalConfigHandle; @@ -96,6 +97,7 @@ import org.hamcrest.Matchers; import org.joda.time.Duration; import org.joda.time.Instant; +import org.junit.Assert; import org.junit.Before; import org.junit.Rule; import org.junit.Test; @@ -109,6 +111,10 @@ @RunWith(JUnit4.class) public class StreamingModeExecutionContextTest { + private static final FailedWorkHandler FAILING_FAILED_WORK_HANDLER = + ignored -> { + Assert.fail(); + }; @Rule public transient Timeout globalTimeout = Timeout.seconds(600); @Mock private WorkExecutor workExecutor; @@ -202,7 +208,7 @@ private void start(StreamingModeExecutionContext context, Work work, Coder ke /* budgetHandle= */ mock(BoundedQueueExecutorWorkHandle.class), keyCoder, /* keyTransitionListener= */ (k, c) -> {}, - /* onFailedWorkHandler= */ ignored -> {}); + /* onFailedWorkHandler= */ FAILING_FAILED_WORK_HANDLER); } catch (CoderException e) { throw new RuntimeException(e); } @@ -554,7 +560,13 @@ public void testAdvance_success() throws Exception { StreamingModeExecutionContext.KeyTransitionListener mockListener = mock(StreamingModeExecutionContext.KeyTransitionListener.class); executionContext.start( - work1, workExecutor, mockExecutor, mockHandle, null, mockListener, ignored -> {}); + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + mockListener, + FAILING_FAILED_WORK_HANDLER); assertTrue(executionContext.advance()); assertEquals("key2", executionContext.getSerializedKey().toStringUtf8()); @@ -589,7 +601,7 @@ public void testAdvance_noMoreWork() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); assertFalse(executionContext.advance()); } @@ -626,7 +638,7 @@ public void testAdvance_respectsMaxBatchSize() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); verifyNoInteractions(mockExecutor); @@ -664,7 +676,7 @@ public void testAdvance_respectsMaxBatchTime() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); verifyNoInteractions(mockExecutor); @@ -694,7 +706,7 @@ public void testAdvance_workFailed() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); work1.setFailed(); @@ -723,7 +735,7 @@ public void testAdvance_defaultKeyGroup() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); assertFalse(executionContext.advance()); verifyNoInteractions(mockExecutor); @@ -758,7 +770,7 @@ public void testAdvance_experimentDisabled() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); verifyNoInteractions(mockExecutor); @@ -798,7 +810,7 @@ public void testAdvance_respectsMaxBatchSinkBytes() throws Exception { mockHandle, null, (oldWork, newWork) -> {}, - ignored -> {}); + FAILING_FAILED_WORK_HANDLER); context.reportBytesSinked(50); assertFalse(context.advance()); 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 732aa83ee5d5..75230594fc12 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 @@ -48,6 +48,7 @@ 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.util.concurrent.ThreadFactoryBuilder; import org.joda.time.Instant; +import org.junit.Assert; import org.junit.Before; import org.junit.Rule; import org.junit.Test; @@ -69,6 +70,10 @@ public static Collection useFairMonitor() { return Arrays.asList(new Object[][] {{false}, {true}}); } + private static final FailedWorkHandler FAILING_FAILED_WORK_HANDLER = + ignored -> { + Assert.fail(); + }; private static final long MAXIMUM_BYTES_OUTSTANDING = 10000000; private static final int DEFAULT_MAX_THREADS = 2; private static final int DEFAULT_THREAD_EXPIRATION_SEC = 60; @@ -549,7 +554,8 @@ public void testPollWork() throws Exception { targetStart.await(); // Steal work1 using pollWork with compA and keyGroup1 - ExecutableWork stolen1 = testExecutor.pollWork("compA", keyGroup1, stealHandle, ignored -> {}); + ExecutableWork stolen1 = + testExecutor.pollWork("compA", keyGroup1, stealHandle, FAILING_FAILED_WORK_HANDLER); assertNotNull(stolen1); assertEquals(work1, stolen1); @@ -598,7 +604,8 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { ExecutableWork work = createWorkWithCompIdAndKeyGroup("compA", keyGroup, ignored -> {}); testExecutor.execute(work, 100); - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup, stealHandle, ignored -> {}); + ExecutableWork stolen = + testExecutor.pollWork("compA", keyGroup, stealHandle, FAILING_FAILED_WORK_HANDLER); assertNull(stolen); blockerStop.countDown(); @@ -648,16 +655,20 @@ public void testPollWork_skipsFailedWorkAndCallsOnFailedWorkHandler() throws Exc createWorkWithCompIdAndKeyGroupAndWorkToken("compA", keyGroup, 101, ignored -> {}); ExecutableWork work2 = createWorkWithCompIdAndKeyGroupAndWorkToken("compA", keyGroup, 102, ignored -> {}); + ExecutableWork work3 = + createWorkWithCompIdAndKeyGroupAndWorkToken("compA", keyGroup, 103, ignored -> {}); // Enqueue both tasks (they will wait in the queue because the thread is blocked). testExecutor.execute(work1, 100); testExecutor.execute(work2, 150); + testExecutor.execute(work3, 200); - assertEquals(3, testExecutor.elementsOutstanding()); - assertEquals(260, testExecutor.bytesOutstanding()); + assertEquals(4, testExecutor.elementsOutstanding()); + assertEquals(460, testExecutor.bytesOutstanding()); - // Mark work1 as failed while waiting in the queue. + // Mark work1, work3 as failed while waiting in the queue. work1.work().setFailed(); + work3.work().setFailed(); // pollWork should skip work1, close work1's handle, invoke // onFailedWorkHandler callback, and return work2. @@ -671,10 +682,20 @@ public void testPollWork_skipsFailedWorkAndCallsOnFailedWorkHandler() throws Exc // Verify stealHandle merged (blockerWork: 10 bytes, work2: 150 bytes). assertEquals(160, stealHandle.bytes()); + stolen = testExecutor.pollWork("compA", keyGroup, stealHandle, onFailedWorkHandler); + assertNull(stolen); + verify(onFailedWorkHandler).onFailedWork(work3.work()); + + // Still 160, nothing should be merged in. + assertEquals(160, stealHandle.bytes()); + // Polling again should return null since no more tasks exist for keyGroup. assertNull(testExecutor.pollWork("compA", keyGroup, stealHandle, onFailedWorkHandler)); blockerStop.countDown(); + stealHandle.close(); + assertEquals(0, testExecutor.elementsOutstanding()); + assertEquals(0, testExecutor.bytesOutstanding()); testExecutor.shutdown(); } } From b3d191a2c284cf947f0fe3284d95fd72b636afcc Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 13 Aug 2026 22:06:00 +0000 Subject: [PATCH 11/19] address comments --- .../StreamingModeExecutionContextTest.java | 22 ++++++++++++------- 1 file changed, 14 insertions(+), 8 deletions(-) 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 93874f034ec2..5aceb0ca9564 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 @@ -54,6 +54,7 @@ import org.apache.beam.runners.dataflow.options.DataflowWorkerHarnessOptions; import org.apache.beam.runners.dataflow.worker.DataflowExecutionContext.DataflowExecutionStateTracker; import org.apache.beam.runners.dataflow.worker.MetricsToCounterUpdateConverter.Kind; +import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext.KeyTransitionListener; import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext.StreamingModeExecutionState; import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext.StreamingModeExecutionStateRegistry; import org.apache.beam.runners.dataflow.worker.counters.CounterSet; @@ -115,6 +116,11 @@ public class StreamingModeExecutionContextTest { ignored -> { Assert.fail(); }; + + private static final KeyTransitionListener FAILING_KEY_TRANSISITON = + (oldWork, newWork) -> { + Assert.fail(); + }; @Rule public transient Timeout globalTimeout = Timeout.seconds(600); @Mock private WorkExecutor workExecutor; @@ -207,7 +213,7 @@ private void start(StreamingModeExecutionContext context, Work work, Coder ke /* workQueueExecutor= */ mock(BoundedQueueExecutor.class), /* budgetHandle= */ mock(BoundedQueueExecutorWorkHandle.class), keyCoder, - /* keyTransitionListener= */ (k, c) -> {}, + FAILING_KEY_TRANSISITON, /* onFailedWorkHandler= */ FAILING_FAILED_WORK_HANDLER); } catch (CoderException e) { throw new RuntimeException(e); @@ -600,7 +606,7 @@ public void testAdvance_noMoreWork() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); assertFalse(executionContext.advance()); @@ -637,7 +643,7 @@ public void testAdvance_respectsMaxBatchSize() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); @@ -675,7 +681,7 @@ public void testAdvance_respectsMaxBatchTime() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); @@ -705,7 +711,7 @@ public void testAdvance_workFailed() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); work1.setFailed(); @@ -734,7 +740,7 @@ public void testAdvance_defaultKeyGroup() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); assertFalse(executionContext.advance()); @@ -769,7 +775,7 @@ public void testAdvance_experimentDisabled() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); @@ -809,7 +815,7 @@ public void testAdvance_respectsMaxBatchSinkBytes() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> {}, + FAILING_KEY_TRANSISITON, FAILING_FAILED_WORK_HANDLER); context.reportBytesSinked(50); From 49db96acb5dd63789a2329790b851e0ec66e050e Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Thu, 13 Aug 2026 23:13:58 +0000 Subject: [PATCH 12/19] [Dataflow Streaming] Commit size validation for multi key commits --- .../worker/StreamingModeExecutionContext.java | 30 +- .../worker/streaming/ExecutableWork.java | 8 + .../MultiKeyCommitValidationException.java | 28 ++ .../dataflow/worker/streaming/Work.java | 20 + .../worker/util/KeyGroupWorkQueue.java | 12 +- .../processing/StreamingWorkScheduler.java | 10 +- .../failures/WorkFailureProcessor.java | 11 + .../worker/StreamingDataflowWorkerTest.java | 412 ++++++++++++++++++ .../StreamingModeExecutionContextTest.java | 71 ++- .../worker/util/KeyGroupWorkQueueTest.java | 21 + .../failures/WorkFailureProcessorTest.java | 21 + 11 files changed, 628 insertions(+), 16 deletions(-) create mode 100644 runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/MultiKeyCommitValidationException.java diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java index 68dbd61f15f6..27dcc2dfffcc 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java @@ -54,6 +54,7 @@ import org.apache.beam.runners.dataflow.worker.streaming.BoundedQueueExecutorWorkHandle; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; import org.apache.beam.runners.dataflow.worker.streaming.KeyCommitTooLargeException; +import org.apache.beam.runners.dataflow.worker.streaming.MultiKeyCommitValidationException; 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.StreamingGlobalConfig; @@ -717,6 +718,25 @@ private void validateCommitRequestSize() { return; } + // Look at budgetHandle instead of executedWorks because when intermediate work items are + // validated during advance(), the next work item (additionalWork) has already been polled + // and merged into budgetHandle before startForNewKey() adds it to executedWorks. This means + // any item hitting truncation in a multi-key bundle will be retried at least once. + // TODO: Can we request truncation without retrying if the first commit exceed the limits? + BoundedQueueExecutorWorkHandle handle = checkNotNull(budgetHandle); + checkState(!handle.getWorkBatch().isEmpty()); + List currentBatch = handle.getWorkBatch(); + if (currentBatch.size() > 1) { + LOG.warn( + "Windmill Commit limit exceeded on a multi key bundle. Retrying without batching. Batch size: {}", + currentBatch.size()); + for (Work w : currentBatch) { + w.setMultiKeyBatchingDisabled(true); + } + throw new MultiKeyCommitValidationException( + "Commit size validation failed for batch. Retrying individually."); + } + KeyCommitTooLargeException e = KeyCommitTooLargeException.causedBy( systemName, byteLimit, commitRequest, key, hotKeyLoggingEnabled); @@ -730,11 +750,6 @@ private void validateCommitRequestSize() { buildWorkItemTruncationRequestBuilder(currentWork, estimatedCommitSize); currentBuilder.clear(); currentBuilder.mergeFrom(truncationBuilder.build()); - - // TODO: throw and retry when truncation is not on a single key bundle. - checkState( - !multiKeyBundleOptions.multiKeyBundleEnabled(), - "Commit truncation not implemented for multikey bundles"); } private Windmill.WorkItemCommitRequest.Builder buildWorkItemTruncationRequestBuilder( @@ -773,7 +788,9 @@ public boolean advance() throws CoderException { throw new WorkItemCancelledException(activeWork.getWorkItem().getShardingKey()); } - if (activeWork.getKeyGroup().equals(Work.KeyGroup.DEFAULT) || shouldStopBatching()) { + if (activeWork.getKeyGroup().equals(Work.KeyGroup.DEFAULT) + || activeWork.isMultiKeyBatchingDisabled() + || shouldStopBatching()) { return false; } @@ -792,7 +809,6 @@ public boolean advance() throws CoderException { } private boolean shouldStopBatching() { - // TODO: stop batching if the previous work item requested truncation if (workItemsPolled >= multiKeyBundleOptions.maxKeyGroupBatchSize()) { return true; } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ExecutableWork.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ExecutableWork.java index 7748a554f0fc..4a992e872a4c 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ExecutableWork.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/ExecutableWork.java @@ -82,4 +82,12 @@ public String getComputationId() { public Work.KeyGroup getKeyGroup() { return work().getKeyGroup(); } + + /** + * Returns true if multi-key batching is disabled for this work item (e.g. after a prior batch + * commit size validation failure). + */ + public boolean isMultiKeyBatchingDisabled() { + return work().isMultiKeyBatchingDisabled(); + } } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/MultiKeyCommitValidationException.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/MultiKeyCommitValidationException.java new file mode 100644 index 000000000000..f147d380d073 --- /dev/null +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/MultiKeyCommitValidationException.java @@ -0,0 +1,28 @@ +/* + * 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; + +/** + * Thrown when a multi-key bundle exceeds commit size limits, triggering unbatching and local retry + * of individual work items. + */ +public final class MultiKeyCommitValidationException extends RuntimeException { + public MultiKeyCommitValidationException(String message) { + super(message); + } +} 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..2acee9410fa3 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 @@ -83,6 +83,10 @@ public final class Work implements RefreshableWork { private final long serializedWorkItemSize; private volatile TimedState currentState; private volatile boolean isFailed; + // If true, this work item will not be batched with other work items in a multi-key bundle. + // This is used to isolate work items that failed validation (e.g. commit size limit exceeded) + // so they can be retried individually and potentially truncated. + private volatile boolean disableMultiKeyBatching = false; private volatile String processingThreadName = ""; private final AtomicReference<@Nullable AtomicBoolean> onFailureListener = new AtomicReference<>(null); @@ -399,6 +403,22 @@ public boolean isFailed() { return isFailed; } + /** + * Sets whether multi-key batching should be disabled for this work item. When true, this work + * item will not be batched with other work items upon local retry. + */ + public void setMultiKeyBatchingDisabled(boolean disableMultiKeyBatching) { + this.disableMultiKeyBatching = disableMultiKeyBatching; + } + + /** + * Returns true if multi-key batching is disabled for this work item (e.g. after a prior batch + * commit size validation failure). + */ + public boolean isMultiKeyBatchingDisabled() { + return disableMultiKeyBatching; + } + boolean isStuckCommittingAt(Instant stuckCommitDeadline) { return currentState.state() == Work.State.COMMITTING && currentState.startTime().isBefore(stuckCommitDeadline); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java index d151157ec68f..07570030edab 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java @@ -20,6 +20,7 @@ import static org.apache.beam.sdk.util.Preconditions.checkArgumentNotNull; import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull; import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkArgument; +import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkState; import java.util.AbstractQueue; import java.util.Collection; @@ -67,9 +68,14 @@ static class Node { @Nullable Node prevKeyGroupNode; @Nullable Node nextKeyGroupNode; + private static boolean isMultiKeyBatchingDisabled(Runnable task) { + return (task instanceof QueuedWork) + && ((QueuedWork) task).getWork().isMultiKeyBatchingDisabled(); + } + Node(Runnable task) { this.task = task; - if (task instanceof QueuedWork) { + if (task instanceof QueuedWork && !isMultiKeyBatchingDisabled(task)) { this.computationId = ((QueuedWork) task).getWork().getComputationId(); this.keyGroup = ((QueuedWork) task).getWork().getKeyGroup(); } else { @@ -193,6 +199,10 @@ private void unlinkNode(Node node) { if (firstNode == keyGroupWorkList.tail) { return null; } + + // MultiKeyBatchingDisabled items should not be in keyGroupWorkList + checkState(!Node.isMultiKeyBatchingDisabled(firstNode.task)); + unlinkNode(firstNode); return (QueuedWork) firstNode.task; 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 299c67128caf..6130299a5ac3 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 @@ -261,7 +261,7 @@ private void processWork( } catch (Throwable t) { handleProcessWorkFailure(computationState, handle.getWorkBatch(), computationId, work, t); } finally { - List processedWorkBatch = workBatch != null ? workBatch : ImmutableList.of(work); + List processedWorkBatch = workBatch != null ? workBatch : handle.getWorkBatch(); // Update total processing time counters. Updating in finally clause ensures that // work items causing exceptions are also accounted in time spent. recordProcessingTime(stageInfo, processedWorkBatch, processingStartTimeNanos); @@ -413,10 +413,6 @@ private void commitMultiKeyWorkBatch( } for (int i = 0; i < workBatch.size(); i++) { Windmill.WorkItemCommitRequest commit = workItemCommits.get(i); - // TODO: Retry on commit truncations - checkState( - !commit.getExceedsMaxWorkItemCommitBytes(), - "Commit truncation with multikey bundles not implemented"); Work w = workBatch.get(i); multiKeyBuilder.addRequests( commit @@ -425,6 +421,8 @@ private void commitMultiKeyWorkBatch( .build()); } + Windmill.MultiKeyWorkItemCommitRequest multiKeyCommitRequest = multiKeyBuilder.build(); + // Transition states of all completed works in the batch to COMMIT_QUEUED and submit for (Work w : workBatch) { w.setState(Work.State.COMMIT_QUEUED); @@ -435,7 +433,7 @@ private void commitMultiKeyWorkBatch( .workCommitter() .accept( Commit.createMultiKey( - multiKeyBuilder.build(), computationState, ImmutableList.copyOf(workBatch))); + multiKeyCommitRequest, computationState, ImmutableList.copyOf(workBatch))); } private void commitSingleKeyWork( diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java index 8af1840faf92..fc255f4d491e 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java @@ -25,6 +25,7 @@ import javax.annotation.concurrent.ThreadSafe; import org.apache.beam.runners.dataflow.worker.status.LastExceptionDataProvider; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.MultiKeyCommitValidationException; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; import org.apache.beam.sdk.annotations.Internal; @@ -160,6 +161,16 @@ private RetryEvaluation evaluateRetry(String computationId, Work work, Throwable @Nullable final Throwable cause = t.getCause(); Throwable parsedException = (t instanceof UserCodeException && cause != null) ? cause : t; + if (parsedException instanceof MultiKeyCommitValidationException) { + LOG.info( + "Execution of work for computation '{}' on sharding key '{}' for work token '{}' exceeded commit size limits. " + + "Work will be retried locally in smaller batches.", + computationId, + work.getWorkItem().getShardingKey(), + work.getWorkItem().getWorkToken()); + return RetryEvaluation.RETRY_LOCALLY; + } + LastExceptionDataProvider.reportException(parsedException); LOG.debug("Failed work: {}", work); Duration elapsedTimeSinceStart = new Duration(work.getStartTime(), clock.get()); 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 2d35da51a79d..f37ee72502f0 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 @@ -154,6 +154,7 @@ import org.apache.beam.sdk.state.StateSpec; import org.apache.beam.sdk.state.StateSpecs; import org.apache.beam.sdk.state.ValueState; +import org.apache.beam.sdk.testing.ExpectedLogs; import org.apache.beam.sdk.transforms.DoFn; import org.apache.beam.sdk.transforms.DoFnSchemaInformation; import org.apache.beam.sdk.transforms.windowing.AfterPane; @@ -301,6 +302,11 @@ public Long get() { }; @Rule public transient Timeout globalTimeout = Timeout.seconds(600); + + @Rule + public ExpectedLogs expectedStreamingModeExecutionContextLogs = + ExpectedLogs.none(StreamingModeExecutionContext.class); + @Rule public BlockingFn blockingFn = new BlockingFn(); @Rule public TestRule restoreMDC = new RestoreDataflowLoggingMDC(); @Rule public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); @@ -4751,6 +4757,395 @@ public void testSkipInputElementsWithDecodingExceptions() throws Exception { "12345", commit.getOutputMessages(0).getBundles(0).getMessages(0).getData().toStringUtf8()); } + @Test + public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSucceeds() throws Exception { + if (!streamingEngine) { + return; + } + KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); + + List instructions = + Arrays.asList( + makeSourceInstruction(kvCoder), + makeDoFnInstruction(new FixedSizeCommitFn(500), 0, kvCoder), + makeSinkInstruction(kvCoder, 1)); + + StreamingDataflowWorker worker = + makeWorker( + defaultWorkerParams( + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--numberOfWorkerHarnessThreads=1") + .setLocalRetryTimeoutMs(100) + .setInstructions(instructions) + .setStreamingGlobalConfig( + StreamingGlobalConfig.builder() + .setOperationalLimits( + OperationalLimits.builder().setMaxWorkItemCommitBytes(1000).build()) + .build()) + .build()); + worker.start(); + + String batchInputText = + "work {" + + " computation_id: \"" + + DEFAULT_COMPUTATION_ID + + "\"" + + " input_data_watermark: 0" + + " work {" + + " key: \"key1\"" + + " sharding_key: 1" + + " work_token: 1" + + " cache_token: 1" + + " 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: 2" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data2\"" + + " }" + + " }" + + " }" + + "}"; + Windmill.GetWorkResponse batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + + server.whenGetWorkCalled().thenReturn(batchInput); + + Map result = server.waitForAndGetCommits(2); + + assertEquals(2, result.size()); + assertTrue(result.containsKey(1L)); + assertTrue(result.containsKey(2L)); + for (Windmill.WorkItemCommitRequest commitRequest : result.values()) { + assertFalse(commitRequest.getExceedsMaxWorkItemCommitBytes()); + } + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(2, multiKeyCommits.size()); + assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); + assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); + + expectedStreamingModeExecutionContextLogs.verifyWarn( + "Windmill Commit limit exceeded on a multi key bundle"); + + worker.stop(); + } + + @Test + public void testMultiKeyCommit_batchCommitSizeExceededUnBatchTruncates() throws Exception { + if (!streamingEngine) { + return; + } + KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); + + List instructions = + Arrays.asList( + makeSourceInstruction(kvCoder), + makeDoFnInstruction(new FixedSizeCommitFn(500), 0, kvCoder), + makeSinkInstruction(kvCoder, 1)); + + StreamingDataflowWorker worker = + makeWorker( + defaultWorkerParams( + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--numberOfWorkerHarnessThreads=1") + .setLocalRetryTimeoutMs(100) + .setInstructions(instructions) + .setStreamingGlobalConfig( + StreamingGlobalConfig.builder() + .setOperationalLimits( + // Both workitems exceed commit limits + OperationalLimits.builder().setMaxWorkItemCommitBytes(400).build()) + .build()) + .build()); + worker.start(); + + String batchInputText = + "work {" + + " computation_id: \"" + + DEFAULT_COMPUTATION_ID + + "\"" + + " input_data_watermark: 0" + + " work {" + + " key: \"key1\"" + + " sharding_key: 1" + + " work_token: 1" + + " cache_token: 1" + + " 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: 2" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data2\"" + + " }" + + " }" + + " }" + + "}"; + Windmill.GetWorkResponse batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + + server.whenGetWorkCalled().thenReturn(batchInput); + + Map result = server.waitForAndGetCommits(2); + + assertEquals(2, result.size()); + assertTrue(result.containsKey(1L)); + assertTrue(result.containsKey(2L)); + for (Windmill.WorkItemCommitRequest commitRequest : result.values()) { + assertTrue(commitRequest.getExceedsMaxWorkItemCommitBytes()); + } + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(2, multiKeyCommits.size()); + assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); + assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); + + expectedStreamingModeExecutionContextLogs.verifyWarn( + "Windmill Commit limit exceeded on a multi key bundle"); + + worker.stop(); + } + + @Test + public void testMultiKeyCommit_batchCommitSizeExceededUnBatchFirstItemTruncates() + throws Exception { + if (!streamingEngine) { + return; + } + KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); + + List instructions = + Arrays.asList( + makeSourceInstruction(kvCoder), + makeDoFnInstruction(new LargeCommitFn(), 0, kvCoder), + makeSinkInstruction(kvCoder, 1)); + + StreamingDataflowWorker worker = + makeWorker( + defaultWorkerParams( + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--numberOfWorkerHarnessThreads=1") + .setLocalRetryTimeoutMs(100) + .setInstructions(instructions) + .setStreamingGlobalConfig( + StreamingGlobalConfig.builder() + .setOperationalLimits( + OperationalLimits.builder().setMaxWorkItemCommitBytes(1000).build()) + .build()) + .build()); + worker.start(); + + String batchInputText = + "work {" + + " computation_id: \"" + + DEFAULT_COMPUTATION_ID + + "\"" + + " input_data_watermark: 0" + + " work {" + + " key: \"large_key\"" + + " sharding_key: 1" + + " work_token: 1" + + " cache_token: 1" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data1\"" + + " }" + + " }" + + " }" + + " work {" + + " key: \"small_key\"" + + " sharding_key: 2" + + " work_token: 2" + + " cache_token: 2" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data2\"" + + " }" + + " }" + + " }" + + "}"; + Windmill.GetWorkResponse batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + + server.whenGetWorkCalled().thenReturn(batchInput); + + Map result = server.waitForAndGetCommits(2); + + assertEquals(2, result.size()); + assertTrue(result.containsKey(1L)); + assertTrue(result.containsKey(2L)); + assertTrue(result.get(1L).getExceedsMaxWorkItemCommitBytes()); + assertFalse(result.get(2L).getExceedsMaxWorkItemCommitBytes()); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(2, multiKeyCommits.size()); + assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); + assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); + + expectedStreamingModeExecutionContextLogs.verifyWarn( + "Windmill Commit limit exceeded on a multi key bundle"); + + worker.stop(); + } + + @Test + public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSecondItemTruncates() + throws Exception { + if (!streamingEngine) { + return; + } + KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); + + List instructions = + Arrays.asList( + makeSourceInstruction(kvCoder), + makeDoFnInstruction(new LargeCommitFn(), 0, kvCoder), + makeSinkInstruction(kvCoder, 1)); + + StreamingDataflowWorker worker = + makeWorker( + defaultWorkerParams( + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--numberOfWorkerHarnessThreads=1") + .setLocalRetryTimeoutMs(100) + .setInstructions(instructions) + .setStreamingGlobalConfig( + StreamingGlobalConfig.builder() + .setOperationalLimits( + OperationalLimits.builder().setMaxWorkItemCommitBytes(1000).build()) + .build()) + .build()); + worker.start(); + + String batchInputText = + "work {" + + " computation_id: \"" + + DEFAULT_COMPUTATION_ID + + "\"" + + " input_data_watermark: 0" + + " work {" + + " key: \"small_key\"" + + " sharding_key: 1" + + " work_token: 1" + + " cache_token: 1" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data1\"" + + " }" + + " }" + + " }" + + " work {" + + " key: \"large_key\"" + + " sharding_key: 2" + + " work_token: 2" + + " cache_token: 2" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data2\"" + + " }" + + " }" + + " }" + + "}"; + Windmill.GetWorkResponse batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + + server.whenGetWorkCalled().thenReturn(batchInput); + + Map result = server.waitForAndGetCommits(2); + + assertEquals(2, result.size()); + assertTrue(result.containsKey(1L)); + assertTrue(result.containsKey(2L)); + assertFalse(result.get(1L).getExceedsMaxWorkItemCommitBytes()); + assertTrue(result.get(2L).getExceedsMaxWorkItemCommitBytes()); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(2, multiKeyCommits.size()); + assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); + assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); + + expectedStreamingModeExecutionContextLogs.verifyWarn( + "Windmill Commit limit exceeded on a multi key bundle"); + + worker.stop(); + } + static class BlockingFn extends DoFn implements TestRule { public static AtomicReference blocker = @@ -4841,6 +5236,23 @@ public void processElement(ProcessContext c) { } } + static class FixedSizeCommitFn extends DoFn, KV> { + private final int size; + + FixedSizeCommitFn(int size) { + this.size = size; + } + + @ProcessElement + public void processElement(ProcessContext c) { + StringBuilder s = new StringBuilder(); + for (int i = 0; i < size; ++i) { + s.append("a"); + } + c.output(KV.of(c.element().getKey(), s.toString())); + } + } + static class ExceptionCatchingFn extends DoFn, KV> { @ProcessElement diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java index eb6bb51e4207..a7aee0215a5f 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 @@ -76,7 +76,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache; import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodingV1; import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodingV2; -import org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.FailureTracker; +import org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.StreamingEngineFailureTracker; import org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender; import org.apache.beam.sdk.Pipeline; import org.apache.beam.sdk.coders.Coder; @@ -152,7 +152,7 @@ private StreamingModeExecutionContext createExecutionContext( /*stepName=*/ "stepName", /*systemName=*/ "systemName", StreamingCounters.create(), - mock(FailureTracker.class), + StreamingEngineFailureTracker.create(10, 10), "sourceBytesProcessCounterName", MultiKeyBundleOptions.fromOptions(options), SideInputStateFetcherFactory.fromOptions(options)); @@ -695,6 +695,32 @@ public void testAdvance_defaultKeyGroup() throws Exception { verifyNoInteractions(mockExecutor); } + @Test + public void testAdvance_batchingDisabled() throws Exception { + BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class); + BoundedQueueExecutorWorkHandle mockHandle = mock(BoundedQueueExecutorWorkHandle.class); + + Windmill.Uint128Proto keyGroup = + Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build(); + Windmill.WorkItem workItem1 = + Windmill.WorkItem.newBuilder() + .setKey(ByteString.copyFromUtf8("key1")) + .setWorkToken(1L) + .setKeyGroup(keyGroup) + .build(); + Work work1 = + createMockWork( + workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); + + work1.setMultiKeyBatchingDisabled(true); + + executionContext.start( + work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(executionContext.advance()); + verifyNoInteractions(mockExecutor); + } + @Test public void testAdvance_experimentDisabled() throws Exception { DataflowWorkerHarnessOptions optionsDisabled = @@ -833,4 +859,45 @@ public void testInternalsPoisonedAfterFlushState() throws Exception { assertThat(e.getMessage(), Matchers.containsString("poisoned")); } } + + @Test + public void testAdvance_stopsWhenQueuedWorkBatchingDisabled() throws Exception { + DataflowWorkerHarnessOptions optionsMultiKey = + PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + optionsMultiKey + .as(ExperimentalOptions.class) + .setExperiments(Arrays.asList("unstable_enable_multi_key_bundle")); + StreamingModeExecutionContext context = + createExecutionContext(optionsMultiKey, globalConfigHandle); + + BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class); + BoundedQueueExecutorWorkHandle mockHandle = mock(BoundedQueueExecutorWorkHandle.class); + Windmill.Uint128Proto keyGroup = + Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build(); + + Work work1 = + createMockWork( + Windmill.WorkItem.newBuilder() + .setKey(ByteString.copyFromUtf8("key1")) + .setWorkToken(1L) + .setKeyGroup(keyGroup) + .build(), + Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); + assertFalse(work1.isMultiKeyBatchingDisabled()); + + when(mockExecutor.pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle)).thenReturn(null); + + AtomicBoolean transitionListenerCalled = new AtomicBoolean(false); + context.start( + work1, + workExecutor, + mockExecutor, + mockHandle, + null, + (oldWork, newWork) -> transitionListenerCalled.set(true)); + + assertFalse(context.advance()); + assertFalse(transitionListenerCalled.get()); + verify(mockExecutor).pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle); + } } 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..c7be44525502 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 @@ -489,6 +489,27 @@ public void testPollWorkWithKeyGroup() { assertTrue(queue.isEmpty()); } + @Test + public void testOffer_multiKeyBatchingDisabled_notInsertedInKeyGroupQueue() { + KeyGroupWorkQueue queue = new KeyGroupWorkQueue(fairQueue); + QueuedWork workDisabled = createQueuedWork("compA", 100); + workDisabled.getWork().work().setMultiKeyBatchingDisabled(true); + QueuedWork workEnabled = createQueuedWork("compA", 200); + + queue.offer(workDisabled); + queue.offer(workEnabled); + assertEquals(2, queue.size()); + + QueuedWork polledWork = queue.pollWork("compA", TEST_KEY_GROUP); + assertNotNull(polledWork); + assertEquals(workEnabled, polledWork); + assertEquals(1, queue.size()); + + assertNull(queue.pollWork("compA", TEST_KEY_GROUP)); + assertEquals(workDisabled, queue.poll()); + assertTrue(queue.isEmpty()); + } + private void waitForThreadState(Thread t, State state) throws InterruptedException { long timeoutMs = 30000; long start = System.currentTimeMillis(); 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..46c1e1b5c4e0 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 @@ -30,6 +30,7 @@ import java.util.function.Consumer; import java.util.function.Supplier; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.MultiKeyCommitValidationException; 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.BoundedQueueExecutor; @@ -239,4 +240,24 @@ public void logAndProcessFailureBatch_mixRetryAndAbort() throws Throwable { assertThat(executedWork2).isEmpty(); assertThat(invalidWork).containsExactly(work2.work()); } + + @Test + public void logAndProcessFailureBatch_retriesOnMultiKeyCommitValidationException() + throws Throwable { + CountDownLatch runWork = new CountDownLatch(1); + ExecutableWork work = createWork(ignored -> runWork.countDown()); + FailureTracker failureTracker = streamingEngineFailureReporter(); + WorkFailureProcessor workFailureProcessor = createWorkFailureProcessor(failureTracker); + Set invalidWork = new HashSet<>(); + + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, + List.of(work), + new MultiKeyCommitValidationException("test"), + invalidWork::add); + + runWork.await(); + assertThat(invalidWork).isEmpty(); + assertThat(failureTracker.drainPendingFailuresToReport()).isEmpty(); + } } From 96c3b0c3b39aeffddac87a0459c946fc52597548 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Fri, 14 Aug 2026 23:38:22 +0000 Subject: [PATCH 13/19] address comments --- .../dataflow/worker/StreamingModeExecutionContext.java | 9 +++++---- .../work/processing/failures/WorkFailureProcessor.java | 2 +- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java index ae34528b064f..7c3c71d6000d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java @@ -719,10 +719,11 @@ private void validateCommitRequestSize() { return; } - // Look at budgetHandle instead of executedWorks because when intermediate work items are - // validated during advance(), the next work item (additionalWork) has already been polled - // and merged into budgetHandle before startForNewKey() adds it to executedWorks. This means - // any item hitting truncation in a multi-key bundle will be retried at least once. + // If this is a multi-key work item, then we need to retry all of the individual work items + // without merging so that we can identify large commits to truncate. We determine the work + // items that were part of the bundle by looking at the budgethandle instead of executedWorks + // because validateCommitRequestSize is called when transitioning and the handle has been + // updated but executedWorks has not. // TODO: Can we request truncation without retrying if the first commit exceed the limits? BoundedQueueExecutorWorkHandle handle = checkNotNull(budgetHandle); checkState(!handle.getWorkBatch().isEmpty()); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java index 84fd4a011d6a..b635bde7e08a 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessor.java @@ -24,8 +24,8 @@ import javax.annotation.concurrent.ThreadSafe; import org.apache.beam.runners.dataflow.worker.status.LastExceptionDataProvider; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; -import org.apache.beam.runners.dataflow.worker.streaming.MultiKeyCommitValidationException; import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; +import org.apache.beam.runners.dataflow.worker.streaming.MultiKeyCommitValidationException; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; import org.apache.beam.sdk.annotations.Internal; From cc27c02abdaac28444d5f6ac960b6947c56aa22c Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Sat, 15 Aug 2026 00:12:37 +0000 Subject: [PATCH 14/19] fix test --- .../work/processing/failures/WorkFailureProcessorTest.java | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) 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 b7fe03c7fe67..f1cc33c963f1 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 @@ -30,6 +30,7 @@ import java.util.function.Consumer; import java.util.function.Supplier; import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; +import org.apache.beam.runners.dataflow.worker.streaming.FailedWorkHandler; import org.apache.beam.runners.dataflow.worker.streaming.MultiKeyCommitValidationException; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; @@ -284,10 +285,11 @@ public void logAndProcessFailureBatch_retriesOnMultiKeyCommitValidationException Set invalidWork = new HashSet<>(); workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, DEFAULT_COMPUTATION_ID, List.of(work), new MultiKeyCommitValidationException("test"), - invalidWork::add); + (FailedWorkHandler) invalidWork::add); runWork.await(); assertThat(invalidWork).isEmpty(); From dd8267a112f4a36c79c4cfb16b7c446594797995 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Sat, 15 Aug 2026 07:48:03 +0000 Subject: [PATCH 15/19] fix merge --- .../worker/StreamingModeExecutionContextTest.java | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) 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 8d780010e795..880ba0ca49d1 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 @@ -926,7 +926,9 @@ public void testAdvance_stopsWhenQueuedWorkBatchingDisabled() throws Exception { Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); assertFalse(work1.isMultiKeyBatchingDisabled()); - when(mockExecutor.pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle)).thenReturn(null); + when(mockExecutor.pollWork( + COMPUTATION_ID, work1.getKeyGroup(), mockHandle, FAILING_FAILED_WORK_HANDLER)) + .thenReturn(null); AtomicBoolean transitionListenerCalled = new AtomicBoolean(false); context.start( @@ -935,10 +937,12 @@ public void testAdvance_stopsWhenQueuedWorkBatchingDisabled() throws Exception { mockExecutor, mockHandle, null, - (oldWork, newWork) -> transitionListenerCalled.set(true)); + (oldWork, newWork) -> transitionListenerCalled.set(true), + FAILING_FAILED_WORK_HANDLER); assertFalse(context.advance()); assertFalse(transitionListenerCalled.get()); - verify(mockExecutor).pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle); + verify(mockExecutor) + .pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle, FAILING_FAILED_WORK_HANDLER); } } From 0505956e96c9213c6e073cad605b8c2526083849 Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Mon, 17 Aug 2026 20:58:09 +0000 Subject: [PATCH 16/19] address comments --- gradle.properties | 3 + .../worker/StreamingModeExecutionContext.java | 7 +- .../BoundedQueueExecutorWorkHandle.java | 8 +- .../worker/util/BoundedQueueExecutor.java | 11 +- .../worker/util/KeyGroupWorkQueue.java | 6 +- .../processing/StreamingWorkScheduler.java | 21 +-- .../worker/StreamingDataflowWorkerTest.java | 177 +++++++++++++++--- 7 files changed, 174 insertions(+), 59 deletions(-) diff --git a/gradle.properties b/gradle.properties index 9503a0933dd3..4bd47f5f3653 100644 --- a/gradle.properties +++ b/gradle.properties @@ -47,3 +47,6 @@ python_versions=3.10,3.11,3.12,3.13,3.14 # Maven Central fallback mirror URL mavenCentralMirrorUrl=https://maven-central.storage-download.googleapis.com/maven2/ + +# Enabled parallel sync for Gradle 9.4+ +org.gradle.tooling.parallel=true diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java index 7c3c71d6000d..8706ff05b6ac 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java @@ -726,8 +726,8 @@ private void validateCommitRequestSize() { // updated but executedWorks has not. // TODO: Can we request truncation without retrying if the first commit exceed the limits? BoundedQueueExecutorWorkHandle handle = checkNotNull(budgetHandle); - checkState(!handle.getWorkBatch().isEmpty()); List currentBatch = handle.getWorkBatch(); + checkState(!currentBatch.isEmpty()); if (currentBatch.size() > 1) { LOG.warn( "Windmill Commit limit exceeded on a multi key bundle. Retrying without batching. Batch size: {}", @@ -879,11 +879,6 @@ public List getWorkItemCommits() { return commits; } - // Returns list of Work that was executed in the bundle - public List getExecutedWorks() { - return executedWorks; - } - // Returns finalization callbacks recorded during the bundle execution public Map> getFinalizationCallbacks() { return finalizationCallbacks; diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/BoundedQueueExecutorWorkHandle.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/BoundedQueueExecutorWorkHandle.java index 20661aae0a04..d7a61562bc58 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/BoundedQueueExecutorWorkHandle.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/BoundedQueueExecutorWorkHandle.java @@ -17,13 +17,15 @@ */ package org.apache.beam.runners.dataflow.worker.streaming; -import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList; +import java.util.List; /** * A handle to use when requesting pulling more work from @BoundedQueueExecutor * via @BoundedQueueExecutor.pollWork */ public interface BoundedQueueExecutorWorkHandle { - // Returns all work that are tracked by the handle - ImmutableList getWorkBatch(); + // Returns all work that are tracked by the handle. + // Returned list cannot be modified. Copying the list is fine. + // Don't keep reference to the returned list after the processing exits the harness threads. + List getWorkBatch(); } 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 2dd0f971168e..046d8cae9f9d 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 @@ -22,6 +22,7 @@ import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkArgument; import java.util.ArrayList; +import java.util.Collections; import java.util.List; import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.LinkedBlockingQueue; @@ -36,7 +37,6 @@ import org.apache.beam.runners.dataflow.worker.streaming.Work; 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.util.concurrent.Monitor; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.Monitor.Guard; import org.checkerframework.checker.nullness.qual.Nullable; @@ -306,8 +306,13 @@ public synchronized boolean isClosed() { } @Override - public synchronized ImmutableList getWorkBatch() { - return ImmutableList.copyOf(workBatch); + /* + * Returns an unmodifiable view over the underlying list. + * It is unsafe to use the returned list with concurrent calls to mutating methods + * like merge/close + */ + public synchronized List getWorkBatch() { + return Collections.unmodifiableList(workBatch); } @VisibleForTesting diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java index 07570030edab..dd409616ab9d 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/KeyGroupWorkQueue.java @@ -69,13 +69,13 @@ static class Node { @Nullable Node nextKeyGroupNode; private static boolean isMultiKeyBatchingDisabled(Runnable task) { - return (task instanceof QueuedWork) - && ((QueuedWork) task).getWork().isMultiKeyBatchingDisabled(); + return !(task instanceof QueuedWork) + || ((QueuedWork) task).getWork().isMultiKeyBatchingDisabled(); } Node(Runnable task) { this.task = task; - if (task instanceof QueuedWork && !isMultiKeyBatchingDisabled(task)) { + if (!isMultiKeyBatchingDisabled(task)) { this.computationId = ((QueuedWork) task).getWork().getComputationId(); this.keyGroup = ((QueuedWork) task).getWork().getKeyGroup(); } else { 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 93aa06ae929d..958cd62f5eb3 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 @@ -242,7 +242,6 @@ private void processWork( long processingStartTimeNanos = System.nanoTime(); StageInfo stageInfo = getStageInfo(computationState); - @Nullable List workBatch = null; try { if (work.isFailed()) { throw new WorkItemCancelledException(workItem.getShardingKey()); @@ -251,7 +250,7 @@ private void processWork( // Execute the user code for the Work batch. ExecuteWorkResult executeWorkResult = executeWork(work, stageInfo, computationState, handle, keyTransitionListener); - workBatch = executeWorkResult.workBatch(); + List workBatch = handle.getWorkBatch(); List workItemCommits = executeWorkResult.workItemCommits(); commitFinalizer.cacheCommitFinalizers(executeWorkResult.finalizationCallbacks()); @@ -264,7 +263,7 @@ private void processWork( handleProcessWorkFailure( computationState, handle.getWorkBatch(), computationId, systemName, work, t); } finally { - List processedWorkBatch = workBatch != null ? workBatch : handle.getWorkBatch(); + List processedWorkBatch = handle.getWorkBatch(); // Update total processing time counters. Updating in finally clause ensures that // work items causing exceptions are also accounted in time spent. recordProcessingTime(stageInfo, processedWorkBatch, processingStartTimeNanos); @@ -328,7 +327,6 @@ private ExecuteWorkResult executeWork( computationWorkExecutor.executeWork( work, workExecutor, handle, keyTransitionListener, onFailedWorkHandler); - List workBatch; List workItemCommits; Map> finalizationCallbacks; long stateBytesRead; @@ -338,9 +336,6 @@ private ExecuteWorkResult executeWork( } context.flushState(); - // Retrieve executed works, work item commits, and accumulated callbacks from execution - // context - workBatch = context.getExecutedWorks(); workItemCommits = context.getWorkItemCommits(); finalizationCallbacks = context.getFinalizationCallbacks(); stateBytesRead = context.getStateBytesRead(); @@ -351,8 +346,7 @@ private ExecuteWorkResult executeWork( computationState.releaseComputationWorkExecutor(computationWorkExecutor); computationWorkExecutor = null; - return ExecuteWorkResult.create( - workBatch, workItemCommits, finalizationCallbacks, stateBytesRead); + return ExecuteWorkResult.create(workItemCommits, finalizationCallbacks, stateBytesRead); } catch (Throwable t) { if (computationWorkExecutor != null) { // If processing failed due to a thrown exception, close the executionState. Do not @@ -427,8 +421,6 @@ private void commitMultiKeyWorkBatch( .build()); } - Windmill.MultiKeyWorkItemCommitRequest multiKeyCommitRequest = multiKeyBuilder.build(); - // Transition states of all completed works in the batch to COMMIT_QUEUED and submit for (Work w : workBatch) { w.setState(Work.State.COMMIT_QUEUED); @@ -439,7 +431,7 @@ private void commitMultiKeyWorkBatch( .workCommitter() .accept( Commit.createMultiKey( - multiKeyCommitRequest, computationState, ImmutableList.copyOf(workBatch))); + multiKeyBuilder.build(), computationState, ImmutableList.copyOf(workBatch))); } private void commitSingleKeyWork( @@ -521,16 +513,13 @@ private KeyTransitionListener createKeyTransitionListener() { @AutoValue abstract static class ExecuteWorkResult { static ExecuteWorkResult create( - List workBatch, List workItemCommits, Map> finalizationCallbacks, long stateBytesRead) { return new AutoValue_StreamingWorkScheduler_ExecuteWorkResult( - workBatch, workItemCommits, finalizationCallbacks, stateBytesRead); + workItemCommits, finalizationCallbacks, stateBytesRead); } - abstract List workBatch(); - abstract List workItemCommits(); // Map> 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 be76527cdc10..d8063ae66d44 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 @@ -151,6 +151,7 @@ import org.apache.beam.sdk.coders.VarIntCoder; import org.apache.beam.sdk.extensions.gcp.util.Transport; import org.apache.beam.sdk.options.PipelineOptionsFactory; +import org.apache.beam.sdk.state.BagState; import org.apache.beam.sdk.state.StateSpec; import org.apache.beam.sdk.state.StateSpecs; import org.apache.beam.sdk.state.ValueState; @@ -352,6 +353,8 @@ private Iterable buildCounters() { @Before public void setUp() { + FixedSizeBagCommitFn.SEEN_ELEMENTS.set(0); + LargeBagCommitFn.SEEN_ELEMENTS.set(0); server.clearCommitsReceived(); streamingCounters = StreamingCounters.create(); } @@ -4890,6 +4893,8 @@ public void testSkipInputElementsWithDecodingExceptions() throws Exception { "12345", commit.getOutputMessages(0).getBundles(0).getMessages(0).getData().toStringUtf8()); } + // TODO: Add similar tests with productions after changing WindmillSink to flush in finishKey. + @Test public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSucceeds() throws Exception { if (!streamingEngine) { @@ -4900,13 +4905,13 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSucceeds() throws E List instructions = Arrays.asList( makeSourceInstruction(kvCoder), - makeDoFnInstruction(new FixedSizeCommitFn(500), 0, kvCoder), + makeDoFnInstruction(new FixedSizeBagCommitFn(500), 0, kvCoder), makeSinkInstruction(kvCoder, 1)); StreamingDataflowWorker worker = makeWorker( defaultWorkerParams( - "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000", "--numberOfWorkerHarnessThreads=1") .setLocalRetryTimeoutMs(100) .setInstructions(instructions) @@ -4956,6 +4961,22 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSucceeds() throws E + " }" + " }" + " }" + + " work {" + + " key: \"key3\"" + + " sharding_key: 3" + + " work_token: 3" + + " cache_token: 3" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data3\"" + + " }" + + " }" + + " }" + "}"; Windmill.GetWorkResponse batchInput = buildInput( @@ -4966,21 +4987,25 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSucceeds() throws E server.whenGetWorkCalled().thenReturn(batchInput); - Map result = server.waitForAndGetCommits(2); + Map result = server.waitForAndGetCommits(3); - assertEquals(2, result.size()); + assertEquals(3, result.size()); assertTrue(result.containsKey(1L)); assertTrue(result.containsKey(2L)); + assertTrue(result.containsKey(3L)); for (Windmill.WorkItemCommitRequest commitRequest : result.values()) { assertFalse(commitRequest.getExceedsMaxWorkItemCommitBytes()); } List multiKeyCommits = server.getMultiKeyCommitsReceived(); - assertEquals(2, multiKeyCommits.size()); + assertEquals(3, multiKeyCommits.size()); assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); - + assertEquals(1, multiKeyCommits.get(2).getRequestsCount()); + // 2 in initial batch (item 1 succeeds with 500 bytes, item 2 fails after accumulating 1000 + // bytes) + 3 unbatched retries + assertEquals(5, FixedSizeBagCommitFn.SEEN_ELEMENTS.get()); expectedStreamingModeExecutionContextLogs.verifyWarn( "Windmill Commit limit exceeded on a multi key bundle"); @@ -4997,20 +5022,20 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchTruncates() throws List instructions = Arrays.asList( makeSourceInstruction(kvCoder), - makeDoFnInstruction(new FixedSizeCommitFn(500), 0, kvCoder), + makeDoFnInstruction(new FixedSizeBagCommitFn(500), 0, kvCoder), makeSinkInstruction(kvCoder, 1)); StreamingDataflowWorker worker = makeWorker( defaultWorkerParams( - "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000", "--numberOfWorkerHarnessThreads=1") .setLocalRetryTimeoutMs(100) .setInstructions(instructions) .setStreamingGlobalConfig( StreamingGlobalConfig.builder() .setOperationalLimits( - // Both workitems exceed commit limits + // All workitems exceed commit limits OperationalLimits.builder().setMaxWorkItemCommitBytes(400).build()) .build()) .build()); @@ -5054,6 +5079,22 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchTruncates() throws + " }" + " }" + " }" + + " work {" + + " key: \"key3\"" + + " sharding_key: 3" + + " work_token: 3" + + " cache_token: 3" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data3\"" + + " }" + + " }" + + " }" + "}"; Windmill.GetWorkResponse batchInput = buildInput( @@ -5064,21 +5105,24 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchTruncates() throws server.whenGetWorkCalled().thenReturn(batchInput); - Map result = server.waitForAndGetCommits(2); + Map result = server.waitForAndGetCommits(3); - assertEquals(2, result.size()); + assertEquals(3, result.size()); assertTrue(result.containsKey(1L)); assertTrue(result.containsKey(2L)); + assertTrue(result.containsKey(3L)); for (Windmill.WorkItemCommitRequest commitRequest : result.values()) { assertTrue(commitRequest.getExceedsMaxWorkItemCommitBytes()); } List multiKeyCommits = server.getMultiKeyCommitsReceived(); - assertEquals(2, multiKeyCommits.size()); + assertEquals(3, multiKeyCommits.size()); assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); - + assertEquals(1, multiKeyCommits.get(2).getRequestsCount()); + // 1 in initial batch (fails after first item's bag write exceeds limit) + 3 unbatched retries + assertEquals(4, FixedSizeBagCommitFn.SEEN_ELEMENTS.get()); expectedStreamingModeExecutionContextLogs.verifyWarn( "Windmill Commit limit exceeded on a multi key bundle"); @@ -5096,13 +5140,13 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchFirstItemTruncates( List instructions = Arrays.asList( makeSourceInstruction(kvCoder), - makeDoFnInstruction(new LargeCommitFn(), 0, kvCoder), + makeDoFnInstruction(new LargeBagCommitFn(), 0, kvCoder), makeSinkInstruction(kvCoder, 1)); StreamingDataflowWorker worker = makeWorker( defaultWorkerParams( - "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000", "--numberOfWorkerHarnessThreads=1") .setLocalRetryTimeoutMs(100) .setInstructions(instructions) @@ -5152,6 +5196,22 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchFirstItemTruncates( + " }" + " }" + " }" + + " work {" + + " key: \"small_key\"" + + " sharding_key: 3" + + " work_token: 3" + + " cache_token: 3" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data3\"" + + " }" + + " }" + + " }" + "}"; Windmill.GetWorkResponse batchInput = buildInput( @@ -5162,19 +5222,24 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchFirstItemTruncates( server.whenGetWorkCalled().thenReturn(batchInput); - Map result = server.waitForAndGetCommits(2); + Map result = server.waitForAndGetCommits(3); - assertEquals(2, result.size()); + assertEquals(3, result.size()); assertTrue(result.containsKey(1L)); assertTrue(result.containsKey(2L)); + assertTrue(result.containsKey(3L)); assertTrue(result.get(1L).getExceedsMaxWorkItemCommitBytes()); assertFalse(result.get(2L).getExceedsMaxWorkItemCommitBytes()); + assertFalse(result.get(3L).getExceedsMaxWorkItemCommitBytes()); List multiKeyCommits = server.getMultiKeyCommitsReceived(); - assertEquals(2, multiKeyCommits.size()); + assertEquals(3, multiKeyCommits.size()); assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); + assertEquals(1, multiKeyCommits.get(2).getRequestsCount()); + // 1 in initial batch (fails after first item's bag write exceeds limit) + 3 unbatched retries + assertEquals(4, LargeBagCommitFn.SEEN_ELEMENTS.get()); expectedStreamingModeExecutionContextLogs.verifyWarn( "Windmill Commit limit exceeded on a multi key bundle"); @@ -5193,13 +5258,13 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSecondItemTruncates List instructions = Arrays.asList( makeSourceInstruction(kvCoder), - makeDoFnInstruction(new LargeCommitFn(), 0, kvCoder), + makeDoFnInstruction(new LargeBagCommitFn(), 0, kvCoder), makeSinkInstruction(kvCoder, 1)); StreamingDataflowWorker worker = makeWorker( defaultWorkerParams( - "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000", "--numberOfWorkerHarnessThreads=1") .setLocalRetryTimeoutMs(100) .setInstructions(instructions) @@ -5249,6 +5314,22 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSecondItemTruncates + " }" + " }" + " }" + + " work {" + + " key: \"small_key\"" + + " sharding_key: 3" + + " work_token: 3" + + " cache_token: 3" + + " key_group { high: 0 low: 1 }" + + " message_bundles {" + + " source_computation_id: \"" + + DEFAULT_SOURCE_COMPUTATION_ID + + "\"" + + " messages {" + + " timestamp: 0" + + " data: \"data3\"" + + " }" + + " }" + + " }" + "}"; Windmill.GetWorkResponse batchInput = buildInput( @@ -5259,19 +5340,25 @@ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchSecondItemTruncates server.whenGetWorkCalled().thenReturn(batchInput); - Map result = server.waitForAndGetCommits(2); + Map result = server.waitForAndGetCommits(3); - assertEquals(2, result.size()); + assertEquals(3, result.size()); assertTrue(result.containsKey(1L)); assertTrue(result.containsKey(2L)); + assertTrue(result.containsKey(3L)); assertFalse(result.get(1L).getExceedsMaxWorkItemCommitBytes()); assertTrue(result.get(2L).getExceedsMaxWorkItemCommitBytes()); + assertFalse(result.get(3L).getExceedsMaxWorkItemCommitBytes()); List multiKeyCommits = server.getMultiKeyCommitsReceived(); - assertEquals(2, multiKeyCommits.size()); + assertEquals(3, multiKeyCommits.size()); assertEquals(1, multiKeyCommits.get(0).getRequestsCount()); assertEquals(1, multiKeyCommits.get(1).getRequestsCount()); + assertEquals(1, multiKeyCommits.get(2).getRequestsCount()); + // 2 in initial batch (item 1 succeeds, fails after item 2's bag write exceeds limit) + 3 + // unbatched retries + assertEquals(5, LargeBagCommitFn.SEEN_ELEMENTS.get()); expectedStreamingModeExecutionContextLogs.verifyWarn( "Windmill Commit limit exceeded on a multi key bundle"); @@ -5373,7 +5460,6 @@ public static void reset() { } static class LargeCommitFn extends DoFn, KV> { - @ProcessElement public void processElement(ProcessContext c) { if (c.element().getKey().equals("large_key")) { @@ -5388,20 +5474,55 @@ public void processElement(ProcessContext c) { } } - static class FixedSizeCommitFn extends DoFn, KV> { + static class LargeBagCommitFn extends DoFn, KV> { + @StateId("bag") + private final StateSpec> bagSpec = StateSpecs.bag(StringUtf8Coder.of()); + + public static AtomicInteger SEEN_ELEMENTS = new AtomicInteger(); + + @ProcessElement + public void processElement(ProcessContext c, @StateId("bag") BagState bag) { + SEEN_ELEMENTS.incrementAndGet(); + if (c.element().getKey().equals("large_key")) { + StringBuilder s = new StringBuilder(); + for (int i = 0; i < 100; ++i) { + s.append("large_commit"); + } + bag.add(s.toString()); + } else { + bag.add(c.element().getValue()); + } + } + } + + static class FixedSizeBagCommitFn extends DoFn, KV> { + @StateId("bag") + private final StateSpec> bagSpec = StateSpecs.bag(StringUtf8Coder.of()); + private final int size; + public static AtomicInteger SEEN_ELEMENTS = new AtomicInteger(); + private List bundleElements = new ArrayList<>(); - FixedSizeCommitFn(int size) { + FixedSizeBagCommitFn(int size) { this.size = size; } + @StartBundle + public void startBundle() { + bundleElements = new ArrayList<>(); + } + @ProcessElement - public void processElement(ProcessContext c) { + public void processElement(ProcessContext c, @StateId("bag") BagState bag) { + SEEN_ELEMENTS.incrementAndGet(); StringBuilder s = new StringBuilder(); for (int i = 0; i < size; ++i) { s.append("a"); } - c.output(KV.of(c.element().getKey(), s.toString())); + bundleElements.add(s.toString()); + for (String elem : bundleElements) { + bag.add(elem); + } } } From d4ae79663708d385c6170f8422d6fde152a0094f Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Tue, 18 Aug 2026 17:46:02 +0000 Subject: [PATCH 17/19] address comments --- .../worker/StreamingModeExecutionContextTest.java | 11 +++-------- 1 file changed, 3 insertions(+), 8 deletions(-) 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 880ba0ca49d1..53dd96620a55 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 @@ -902,7 +902,7 @@ public void testInternalsPoisonedAfterFlushState() throws Exception { } @Test - public void testAdvance_stopsWhenQueuedWorkBatchingDisabled() throws Exception { + public void testAdvance_stopsWhenCurrentWorkBatchingDisabled() throws Exception { DataflowWorkerHarnessOptions optionsMultiKey = PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); optionsMultiKey @@ -924,11 +924,7 @@ public void testAdvance_stopsWhenQueuedWorkBatchingDisabled() throws Exception { .setKeyGroup(keyGroup) .build(), Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); - assertFalse(work1.isMultiKeyBatchingDisabled()); - - when(mockExecutor.pollWork( - COMPUTATION_ID, work1.getKeyGroup(), mockHandle, FAILING_FAILED_WORK_HANDLER)) - .thenReturn(null); + work1.setMultiKeyBatchingDisabled(true); AtomicBoolean transitionListenerCalled = new AtomicBoolean(false); context.start( @@ -942,7 +938,6 @@ public void testAdvance_stopsWhenQueuedWorkBatchingDisabled() throws Exception { assertFalse(context.advance()); assertFalse(transitionListenerCalled.get()); - verify(mockExecutor) - .pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle, FAILING_FAILED_WORK_HANDLER); + verifyNoInteractions(mockExecutor); } } From 61a4e0fc4a7a0c4e1020ba2b13dfe4150aa6d45e Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Tue, 18 Aug 2026 17:57:38 +0000 Subject: [PATCH 18/19] remove unrelated diff --- gradle.properties | 3 --- 1 file changed, 3 deletions(-) diff --git a/gradle.properties b/gradle.properties index 4bd47f5f3653..9503a0933dd3 100644 --- a/gradle.properties +++ b/gradle.properties @@ -47,6 +47,3 @@ python_versions=3.10,3.11,3.12,3.13,3.14 # Maven Central fallback mirror URL mavenCentralMirrorUrl=https://maven-central.storage-download.googleapis.com/maven2/ - -# Enabled parallel sync for Gradle 9.4+ -org.gradle.tooling.parallel=true From 48b6e82aa9672da1e55fe297645917d815ed608a Mon Sep 17 00:00:00 2001 From: Arun Pandian Date: Tue, 18 Aug 2026 22:24:31 +0000 Subject: [PATCH 19/19] address comments --- .../worker/StreamingModeExecutionContext.java | 21 +++++++------------ 1 file changed, 8 insertions(+), 13 deletions(-) diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java index 8706ff05b6ac..7e9c3eca13f9 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java @@ -194,7 +194,6 @@ public interface KeyTransitionListener { private @Nullable KeyTransitionListener keyTransitionListener; private @Nullable FailedWorkHandler onFailedWorkHandler; - private List executedWorks = Collections.emptyList(); private List outputBuilders = Collections.emptyList(); // Map> @@ -320,7 +319,6 @@ public byte[] getCurrentRecordOffset() { public void reset() { // these lists and maps are returned to callers after processing // don't clear and reuse, instead reset the reference. - this.executedWorks = Collections.emptyList(); this.outputBuilders = Collections.emptyList(); this.finalizationCallbacks = Collections.emptyMap(); // Work from prior bundles might have a reference to the old workBatchFailed. @@ -354,7 +352,6 @@ public void start( FailedWorkHandler onFailedWorkHandler) throws CoderException { reset(); - this.executedWorks = new ArrayList<>(); this.outputBuilders = new ArrayList<>(); this.finalizationCallbacks = new HashMap<>(); this.keyCoder = keyCoder; @@ -579,11 +576,13 @@ public void setActiveReader(UnboundedReader reader) { /** Invalidate the state and reader caches for this computation and key. */ public void invalidateCache() { - for (Work w : executedWorks) { - WindmillComputationKey compKey = - WindmillComputationKey.create(computationId, w.getShardedKey()); - readerCache.invalidateReader(compKey); - stateCache.invalidate(w.getShardedKey()); + if (budgetHandle != null) { + for (Work w : budgetHandle.getWorkBatch()) { + WindmillComputationKey compKey = + WindmillComputationKey.create(computationId, w.getShardedKey()); + readerCache.invalidateReader(compKey); + stateCache.invalidate(w.getShardedKey()); + } } if (activeReader != null) { try { @@ -720,10 +719,7 @@ private void validateCommitRequestSize() { } // If this is a multi-key work item, then we need to retry all of the individual work items - // without merging so that we can identify large commits to truncate. We determine the work - // items that were part of the bundle by looking at the budgethandle instead of executedWorks - // because validateCommitRequestSize is called when transitioning and the handle has been - // updated but executedWorks has not. + // without merging so that we can identify large commits to truncate. // TODO: Can we request truncation without retrying if the first commit exceed the limits? BoundedQueueExecutorWorkHandle handle = checkNotNull(budgetHandle); List currentBatch = handle.getWorkBatch(); @@ -838,7 +834,6 @@ private void startForNewKey(Work newWork) throws CoderException { this.outputBuilder = createOutputBuilder(newWork); this.outputBuilders.add(this.outputBuilder); newWork.setOnFailureListener(this.workBatchFailed); - this.executedWorks.add(newWork); logHotKeyIfDetected(newWork, this.key);