This is an automated email from the ASF dual-hosted git repository.
scwhittle pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/master by this push:
new 6356a3c3d28 [Dataflow Streaming] Commit size validation for multi key
commits (#39473)
6356a3c3d28 is described below
commit 6356a3c3d28e700c4557dce13ffb228ad793a393
Author: Arun Pandian <[email protected]>
AuthorDate: Wed Aug 19 01:18:06 2026 -0700
[Dataflow Streaming] Commit size validation for multi key commits (#39473)
---
.../worker/StreamingModeExecutionContext.java | 49 +-
.../streaming/BoundedQueueExecutorWorkHandle.java | 8 +-
.../dataflow/worker/streaming/ExecutableWork.java | 8 +
...java => MultiKeyCommitValidationException.java} | 13 +-
.../runners/dataflow/worker/streaming/Work.java | 20 +
.../dataflow/worker/util/BoundedQueueExecutor.java | 11 +-
.../dataflow/worker/util/KeyGroupWorkQueue.java | 12 +-
.../work/processing/StreamingWorkScheduler.java | 21 +-
.../processing/failures/WorkFailureProcessor.java | 11 +
.../worker/StreamingDataflowWorkerTest.java | 535 ++++++++++++++++++++-
.../worker/StreamingModeExecutionContextTest.java | 44 +-
.../worker/util/KeyGroupWorkQueueTest.java | 21 +
.../failures/WorkFailureProcessorTest.java | 23 +
13 files changed, 721 insertions(+), 55 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 365ebbdc1f9..7e9c3eca13f 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
@@ -55,6 +55,7 @@ import
org.apache.beam.runners.dataflow.worker.streaming.BoundedQueueExecutorWor
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.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;
@@ -193,7 +194,6 @@ public class StreamingModeExecutionContext
private @Nullable KeyTransitionListener keyTransitionListener;
private @Nullable FailedWorkHandler onFailedWorkHandler;
- private List<Work> executedWorks = Collections.emptyList();
private List<Windmill.WorkItemCommitRequest.Builder> outputBuilders =
Collections.emptyList();
// Map<finalizerId, Pair<callbackExpiration, callback>>
@@ -319,7 +319,6 @@ public class StreamingModeExecutionContext
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.
@@ -353,7 +352,6 @@ public class StreamingModeExecutionContext
FailedWorkHandler onFailedWorkHandler)
throws CoderException {
reset();
- this.executedWorks = new ArrayList<>();
this.outputBuilders = new ArrayList<>();
this.finalizationCallbacks = new HashMap<>();
this.keyCoder = keyCoder;
@@ -578,11 +576,13 @@ public class StreamingModeExecutionContext
/** 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 {
@@ -718,6 +718,23 @@ public class StreamingModeExecutionContext
return;
}
+ // 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.
+ // TODO: Can we request truncation without retrying if the first commit
exceed the limits?
+ BoundedQueueExecutorWorkHandle handle = checkNotNull(budgetHandle);
+ List<Work> 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: {}",
+ 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);
@@ -731,11 +748,6 @@ public class StreamingModeExecutionContext
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(
@@ -774,7 +786,9 @@ public class StreamingModeExecutionContext
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;
}
@@ -797,7 +811,6 @@ public class StreamingModeExecutionContext
}
private boolean shouldStopBatching() {
- // TODO: stop batching if the previous work item requested truncation
if (workItemsPolled >= multiKeyBundleOptions.maxKeyGroupBatchSize()) {
return true;
}
@@ -821,7 +834,6 @@ public class StreamingModeExecutionContext
this.outputBuilder = createOutputBuilder(newWork);
this.outputBuilders.add(this.outputBuilder);
newWork.setOnFailureListener(this.workBatchFailed);
- this.executedWorks.add(newWork);
logHotKeyIfDetected(newWork, this.key);
@@ -862,11 +874,6 @@ public class StreamingModeExecutionContext
return commits;
}
- // Returns list of Work that was executed in the bundle
- public List<Work> getExecutedWorks() {
- return executedWorks;
- }
-
// Returns finalization callbacks recorded during the bundle execution
public Map<Long, Pair<Instant, Runnable>> 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 20661aae0a0..d7a61562bc5 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<Work> 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<Work> getWorkBatch();
}
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 7748a554f0f..4a992e872a4 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 final class ExecutableWork {
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/BoundedQueueExecutorWorkHandle.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/MultiKeyCommitValidationException.java
similarity index 71%
copy from
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/BoundedQueueExecutorWorkHandle.java
copy to
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/MultiKeyCommitValidationException.java
index 20661aae0a0..f147d380d07 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/MultiKeyCommitValidationException.java
@@ -17,13 +17,12 @@
*/
package org.apache.beam.runners.dataflow.worker.streaming;
-import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList;
-
/**
- * A handle to use when requesting pulling more work from @BoundedQueueExecutor
- * via @BoundedQueueExecutor.pollWork
+ * Thrown when a multi-key bundle exceeds commit size limits, triggering
unbatching and local retry
+ * of individual work items.
*/
-public interface BoundedQueueExecutorWorkHandle {
- // Returns all work that are tracked by the handle
- ImmutableList<Work> getWorkBatch();
+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 4541a1c313a..2acee9410fa 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 final class Work implements RefreshableWork {
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/BoundedQueueExecutor.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java
index 2dd0f971168..046d8cae9f9 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.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;
+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.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;
-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 class BoundedQueueExecutor {
}
@Override
- public synchronized ImmutableList<Work> 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<Work> 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 d151157ec68..dd409616ab9 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 @@ 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 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 @@ class KeyGroupWorkQueue extends AbstractQueue<Runnable>
implements BlockingQueue
@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 (!isMultiKeyBatchingDisabled(task)) {
this.computationId = ((QueuedWork) task).getWork().getComputationId();
this.keyGroup = ((QueuedWork) task).getWork().getKeyGroup();
} else {
@@ -193,6 +199,10 @@ class KeyGroupWorkQueue extends AbstractQueue<Runnable>
implements BlockingQueue
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 63cfad5a9a6..958cd62f5eb 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 @@ public class StreamingWorkScheduler {
long processingStartTimeNanos = System.nanoTime();
StageInfo stageInfo = getStageInfo(computationState);
- @Nullable List<Work> workBatch = null;
try {
if (work.isFailed()) {
throw new WorkItemCancelledException(workItem.getShardingKey());
@@ -251,7 +250,7 @@ public class StreamingWorkScheduler {
// Execute the user code for the Work batch.
ExecuteWorkResult executeWorkResult =
executeWork(work, stageInfo, computationState, handle,
keyTransitionListener);
- workBatch = executeWorkResult.workBatch();
+ List<Work> workBatch = handle.getWorkBatch();
List<Windmill.WorkItemCommitRequest> workItemCommits =
executeWorkResult.workItemCommits();
commitFinalizer.cacheCommitFinalizers(executeWorkResult.finalizationCallbacks());
@@ -264,7 +263,7 @@ public class StreamingWorkScheduler {
handleProcessWorkFailure(
computationState, handle.getWorkBatch(), computationId, systemName,
work, t);
} finally {
- List<Work> processedWorkBatch = workBatch != null ? workBatch :
ImmutableList.of(work);
+ List<Work> 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 @@ public class StreamingWorkScheduler {
computationWorkExecutor.executeWork(
work, workExecutor, handle, keyTransitionListener,
onFailedWorkHandler);
- List<Work> workBatch;
List<Windmill.WorkItemCommitRequest> workItemCommits;
Map<Long, Pair<Instant, Runnable>> finalizationCallbacks;
long stateBytesRead;
@@ -338,9 +336,6 @@ public class StreamingWorkScheduler {
}
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 @@ public class StreamingWorkScheduler {
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
@@ -419,10 +413,6 @@ public class StreamingWorkScheduler {
}
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
@@ -523,16 +513,13 @@ public class StreamingWorkScheduler {
@AutoValue
abstract static class ExecuteWorkResult {
static ExecuteWorkResult create(
- List<Work> workBatch,
List<Windmill.WorkItemCommitRequest> workItemCommits,
Map<Long, Pair<Instant, Runnable>> finalizationCallbacks,
long stateBytesRead) {
return new AutoValue_StreamingWorkScheduler_ExecuteWorkResult(
- workBatch, workItemCommits, finalizationCallbacks, stateBytesRead);
+ workItemCommits, finalizationCallbacks, stateBytesRead);
}
- abstract List<Work> workBatch();
-
abstract List<Windmill.WorkItemCommitRequest> workItemCommits();
// Map<finalizerId, Pair<callbackExpiration, callback>>
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 d23c870178e..b635bde7e08 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.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;
@@ -162,6 +163,16 @@ public final class WorkFailureProcessor {
@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 c48b30ecf64..d8063ae66d4 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,9 +151,11 @@ import org.apache.beam.sdk.coders.TextualIntegerCoder;
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;
+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 +303,11 @@ public class StreamingDataflowWorkerTest {
};
@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();
@@ -346,6 +353,8 @@ public class StreamingDataflowWorkerTest {
@Before
public void setUp() {
+ FixedSizeBagCommitFn.SEEN_ELEMENTS.set(0);
+ LargeBagCommitFn.SEEN_ELEMENTS.set(0);
server.clearCommitsReceived();
streamingCounters = StreamingCounters.create();
}
@@ -4884,6 +4893,479 @@ public class StreamingDataflowWorkerTest {
"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) {
+ return;
+ }
+ KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(),
StringUtf8Coder.of());
+
+ List<ParallelInstruction> instructions =
+ Arrays.asList(
+ makeSourceInstruction(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=50000",
+ "--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\""
+ + " }"
+ + " }"
+ + " }"
+ + " 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(
+ batchInputText,
+ CoderUtils.encodeToByteArray(
+ CollectionCoder.of(IntervalWindow.getCoder()),
+ Collections.singletonList(DEFAULT_WINDOW)));
+
+ server.whenGetWorkCalled().thenReturn(batchInput);
+
+ Map<Long, Windmill.WorkItemCommitRequest> result =
server.waitForAndGetCommits(3);
+
+ 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<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+ server.getMultiKeyCommitsReceived();
+ 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");
+
+ worker.stop();
+ }
+
+ @Test
+ public void testMultiKeyCommit_batchCommitSizeExceededUnBatchTruncates()
throws Exception {
+ if (!streamingEngine) {
+ return;
+ }
+ KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(),
StringUtf8Coder.of());
+
+ List<ParallelInstruction> instructions =
+ Arrays.asList(
+ makeSourceInstruction(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=50000",
+ "--numberOfWorkerHarnessThreads=1")
+ .setLocalRetryTimeoutMs(100)
+ .setInstructions(instructions)
+ .setStreamingGlobalConfig(
+ StreamingGlobalConfig.builder()
+ .setOperationalLimits(
+ // All 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\""
+ + " }"
+ + " }"
+ + " }"
+ + " 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(
+ batchInputText,
+ CoderUtils.encodeToByteArray(
+ CollectionCoder.of(IntervalWindow.getCoder()),
+ Collections.singletonList(DEFAULT_WINDOW)));
+
+ server.whenGetWorkCalled().thenReturn(batchInput);
+
+ Map<Long, Windmill.WorkItemCommitRequest> result =
server.waitForAndGetCommits(3);
+
+ 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<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+ server.getMultiKeyCommitsReceived();
+ 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");
+
+ worker.stop();
+ }
+
+ @Test
+ public void
testMultiKeyCommit_batchCommitSizeExceededUnBatchFirstItemTruncates()
+ throws Exception {
+ if (!streamingEngine) {
+ return;
+ }
+ KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(),
StringUtf8Coder.of());
+
+ List<ParallelInstruction> instructions =
+ Arrays.asList(
+ makeSourceInstruction(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=50000",
+ "--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\""
+ + " }"
+ + " }"
+ + " }"
+ + " 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(
+ batchInputText,
+ CoderUtils.encodeToByteArray(
+ CollectionCoder.of(IntervalWindow.getCoder()),
+ Collections.singletonList(DEFAULT_WINDOW)));
+
+ server.whenGetWorkCalled().thenReturn(batchInput);
+
+ Map<Long, Windmill.WorkItemCommitRequest> result =
server.waitForAndGetCommits(3);
+
+ 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<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+ server.getMultiKeyCommitsReceived();
+ 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");
+
+ worker.stop();
+ }
+
+ @Test
+ public void
testMultiKeyCommit_batchCommitSizeExceededUnBatchSecondItemTruncates()
+ throws Exception {
+ if (!streamingEngine) {
+ return;
+ }
+ KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(),
StringUtf8Coder.of());
+
+ List<ParallelInstruction> instructions =
+ Arrays.asList(
+ makeSourceInstruction(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=50000",
+ "--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\""
+ + " }"
+ + " }"
+ + " }"
+ + " 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(
+ batchInputText,
+ CoderUtils.encodeToByteArray(
+ CollectionCoder.of(IntervalWindow.getCoder()),
+ Collections.singletonList(DEFAULT_WINDOW)));
+
+ server.whenGetWorkCalled().thenReturn(batchInput);
+
+ Map<Long, Windmill.WorkItemCommitRequest> result =
server.waitForAndGetCommits(3);
+
+ 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<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+ server.getMultiKeyCommitsReceived();
+ 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");
+
+ worker.stop();
+ }
+
static class BlockingFn extends DoFn<String, String> implements TestRule {
public static AtomicReference<CountDownLatch> blocker =
@@ -4978,7 +5460,6 @@ public class StreamingDataflowWorkerTest {
}
static class LargeCommitFn extends DoFn<KV<String, String>, KV<String,
String>> {
-
@ProcessElement
public void processElement(ProcessContext c) {
if (c.element().getKey().equals("large_key")) {
@@ -4993,6 +5474,58 @@ public class StreamingDataflowWorkerTest {
}
}
+ static class LargeBagCommitFn extends DoFn<KV<String, String>, KV<String,
String>> {
+ @StateId("bag")
+ private final StateSpec<BagState<String>> bagSpec =
StateSpecs.bag(StringUtf8Coder.of());
+
+ public static AtomicInteger SEEN_ELEMENTS = new AtomicInteger();
+
+ @ProcessElement
+ public void processElement(ProcessContext c, @StateId("bag")
BagState<String> 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<String, String>,
KV<String, String>> {
+ @StateId("bag")
+ private final StateSpec<BagState<String>> bagSpec =
StateSpecs.bag(StringUtf8Coder.of());
+
+ private final int size;
+ public static AtomicInteger SEEN_ELEMENTS = new AtomicInteger();
+ private List<String> bundleElements = new ArrayList<>();
+
+ FixedSizeBagCommitFn(int size) {
+ this.size = size;
+ }
+
+ @StartBundle
+ public void startBundle() {
+ bundleElements = new ArrayList<>();
+ }
+
+ @ProcessElement
+ public void processElement(ProcessContext c, @StateId("bag")
BagState<String> bag) {
+ SEEN_ELEMENTS.incrementAndGet();
+ StringBuilder s = new StringBuilder();
+ for (int i = 0; i < size; ++i) {
+ s.append("a");
+ }
+ bundleElements.add(s.toString());
+ for (String elem : bundleElements) {
+ bag.add(elem);
+ }
+ }
+ }
+
static class ExceptionCatchingFn extends DoFn<KV<String, String>, KV<String,
String>> {
@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 5aceb0ca956..53dd96620a5 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
@@ -78,7 +78,7 @@ import
org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDa
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;
@@ -164,7 +164,7 @@ public class StreamingModeExecutionContextTest {
/*stepName=*/ "stepName",
/*systemName=*/ "systemName",
StreamingCounters.create(),
- mock(FailureTracker.class),
+ StreamingEngineFailureTracker.create(10, 10),
"sourceBytesProcessCounterName",
MultiKeyBundleOptions.fromOptions(options),
SideInputStateFetcherFactory.fromOptions(options));
@@ -900,4 +900,44 @@ public class StreamingModeExecutionContextTest {
assertThat(e.getMessage(), Matchers.containsString("poisoned"));
}
}
+
+ @Test
+ public void testAdvance_stopsWhenCurrentWorkBatchingDisabled() 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());
+ work1.setMultiKeyBatchingDisabled(true);
+
+ AtomicBoolean transitionListenerCalled = new AtomicBoolean(false);
+ context.start(
+ work1,
+ workExecutor,
+ mockExecutor,
+ mockHandle,
+ null,
+ (oldWork, newWork) -> transitionListenerCalled.set(true),
+ FAILING_FAILED_WORK_HANDLER);
+
+ assertFalse(context.advance());
+ assertFalse(transitionListenerCalled.get());
+ verifyNoInteractions(mockExecutor);
+ }
}
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 77fcb059758..c7be4452550 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 class KeyGroupWorkQueueTest {
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 741cc35376f..f1cc33c963f 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,8 @@ import java.util.concurrent.TimeUnit;
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;
import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor;
@@ -272,4 +274,25 @@ public class WorkFailureProcessorTest {
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<Work> invalidWork = new HashSet<>();
+
+ workFailureProcessor.logAndProcessFailureBatch(
+ DEFAULT_COMPUTATION_ID,
+ DEFAULT_COMPUTATION_ID,
+ List.of(work),
+ new MultiKeyCommitValidationException("test"),
+ (FailedWorkHandler) invalidWork::add);
+
+ runWork.await();
+ assertThat(invalidWork).isEmpty();
+ assertThat(failureTracker.drainPendingFailuresToReport()).isEmpty();
+ }
}