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 6894ac20ef97..31fa4d9eba98 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 @@ -80,6 +80,7 @@ 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.state.WindmillTimerData; +import org.apache.beam.runners.dataflow.worker.windmill.work.processing.StreamingWorkScheduler.MultiKeyCommitValidationException; import org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.FailureTracker; import org.apache.beam.sdk.annotations.Internal; import org.apache.beam.sdk.coders.Coder; @@ -713,6 +714,17 @@ private void validateCommitRequestSize() { return; } + if (executedWorks.size() > 1) { + LOG.warn( + "Windmill Commit limit exceeded on a multi key bundle. Retrying without batching. Batch size: {}", + executedWorks.size()); + for (Work w : executedWorks) { + w.setDisableMultiKeyBatching(true); + } + throw new MultiKeyCommitValidationException( + "Commit size validation failed for batch. Retrying individually."); + } + KeyCommitTooLargeException e = KeyCommitTooLargeException.causedBy( systemName, byteLimit, commitRequest, key, hotKeyLoggingEnabled); @@ -769,7 +781,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; } @@ -789,7 +803,10 @@ public boolean advance() throws CoderException { } private boolean shouldStopBatching() { - // TODO: stop batching if the previous work item requested truncation + // stop batching if the previous work item requested truncation + if (getOutputBuilder().getExceedsMaxWorkItemCommitBytes()) { + return true; + } if (workItemsPolled >= multiKeyBundleOptions.maxKeyGroupBatchSize()) { return true; } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java index 9cb4a1c0a5be..4bb3be29cfb0 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java @@ -15,7 +15,23 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -package org.apache.beam.runners.dataflow.worker; +package org.apache.beam.runners.dataflow.worker; /* + * 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. + */ import org.checkerframework.checker.nullness.qual.Nullable; 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..f7c494391212 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,8 @@ public String getComputationId() { public Work.KeyGroup getKeyGroup() { return work().getKeyGroup(); } + + 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/Work.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/Work.java index 4541a1c313a2..ab7801822941 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,14 @@ public boolean isFailed() { return isFailed; } + public void setDisableMultiKeyBatching(boolean disableMultiKeyBatching) { + this.disableMultiKeyBatching = disableMultiKeyBatching; + } + + 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 9e8265e509af..022c105f7c2e 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 @@ -417,10 +417,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 @@ -429,6 +425,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); @@ -439,7 +437,7 @@ private void commitMultiKeyWorkBatch( .workCommitter() .accept( Commit.createMultiKey( - multiKeyBuilder.build(), computationState, ImmutableList.copyOf(workBatch))); + multiKeyCommitRequest, computationState, ImmutableList.copyOf(workBatch))); } private void commitSingleKeyWork( @@ -534,4 +532,10 @@ static ExecuteWorkResult create( abstract long stateBytesRead(); } + + public static 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/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..170d412fefe3 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 @@ -27,6 +27,7 @@ import org.apache.beam.runners.dataflow.worker.streaming.ExecutableWork; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; +import org.apache.beam.runners.dataflow.worker.windmill.work.processing.StreamingWorkScheduler.MultiKeyCommitValidationException; import org.apache.beam.sdk.annotations.Internal; import org.apache.beam.sdk.util.UserCodeException; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.annotations.VisibleForTesting; @@ -160,6 +161,17 @@ 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 '{}' failed batch validation. " + + "Work will be retried locally.", + computationId, + work.getWorkItem().getShardingKey()); + return RetryEvaluation.RETRY_LOCALLY; + } + @Nullable final Throwable cause = t.getCause(); + Throwable parsedException = (t instanceof UserCodeException && cause != null) ? cause : t; + 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 9ed705550bc6..a29e9e3c042f 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 @@ -142,6 +142,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.client.grpc.stubs.WindmillChannels; import org.apache.beam.runners.dataflow.worker.windmill.testing.FakeWindmillStubFactory; import org.apache.beam.runners.dataflow.worker.windmill.testing.FakeWindmillStubFactoryFactory; +import org.apache.beam.runners.dataflow.worker.windmill.work.processing.StreamingWorkScheduler; import org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender; import org.apache.beam.sdk.coders.Coder; import org.apache.beam.sdk.coders.Coder.Context; @@ -156,6 +157,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; @@ -303,6 +305,10 @@ public Long get() { }; @Rule public transient Timeout globalTimeout = Timeout.seconds(600); + + @Rule + public ExpectedLogs expectedWorkSchedulerLogs = ExpectedLogs.none(StreamingWorkScheduler.class); + @Rule public BlockingFn blockingFn = new BlockingFn(); @Rule public TestRule restoreMDC = new RestoreDataflowLoggingMDC(); @Rule public final GrpcCleanupRule grpcCleanup = new GrpcCleanupRule(); @@ -4853,6 +4859,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 c5efcea4e47c..0b526330cc21 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)); @@ -693,6 +693,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.setDisableMultiKeyBatching(true); + + executionContext.start( + work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(executionContext.advance()); + verifyNoInteractions(mockExecutor); + } + @Test public void testAdvance_experimentDisabled() throws Exception { DataflowWorkerHarnessOptions optionsDisabled = @@ -831,4 +857,79 @@ public void testInternalsPoisonedAfterFlushState() throws Exception { assertThat(e.getMessage(), Matchers.containsString("poisoned")); } } + + @Test + public void testAdvance_stopsBatchingWhenCommitTruncated() 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()); + + context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + context.getOutputBuilder().setExceedsMaxWorkItemCommitBytes(true); + + assertFalse(context.advance()); + org.mockito.Mockito.verifyNoInteractions(mockExecutor); + } + + @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()); + + org.mockito.Mockito.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()); + org.mockito.Mockito.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..3100b92c6dcf 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().setDisableMultiKeyBatching(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..b1279e4bf9b3 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 @@ -239,4 +239,25 @@ 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 org.apache.beam.runners.dataflow.worker.windmill.work.processing.StreamingWorkScheduler + .MultiKeyCommitValidationException("test"), + invalidWork::add); + + runWork.await(); + assertThat(invalidWork).isEmpty(); + assertThat(failureTracker.drainPendingFailuresToReport()).isEmpty(); + } }