diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/DataflowExecutionContext.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/DataflowExecutionContext.java index 6ff05b4b4452..888e954c1c9f 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/DataflowExecutionContext.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/DataflowExecutionContext.java @@ -150,6 +150,10 @@ boolean isSinkFullHintSet() { // the state size might grow unbounded. } + protected final long getBytesSinked() { + return bytesSinked; + } + /** * Sets a flag to indicate that a sink has enough data written to it. This hint is read by * upstream producers to stop producing if they can. Mainly used in streaming. diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/MultiKeyBundleOptions.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/MultiKeyBundleOptions.java new file mode 100644 index 000000000000..8a264dc57ef6 --- /dev/null +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/MultiKeyBundleOptions.java @@ -0,0 +1,138 @@ +/* + * 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; + +import com.google.auto.value.AutoValue; +import java.util.concurrent.TimeUnit; +import org.apache.beam.sdk.annotations.Internal; +import org.apache.beam.sdk.options.ExperimentalOptions; +import org.apache.beam.sdk.options.PipelineOptions; +import org.checkerframework.checker.nullness.qual.Nullable; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +@AutoValue +@Internal +public abstract class MultiKeyBundleOptions { + // TODO: consider moving this to be part of PipelineOptions after the feature is stable + + private static final Logger LOG = LoggerFactory.getLogger(MultiKeyBundleOptions.class); + + // Don't use. Experiment guarding multi key bundles. The feature is work in progress and + // incomplete. + public static final String UNSTABLE_ENABLE_MULTI_KEY_BUNDLE = "unstable_enable_multi_key_bundle"; + + private static final String WINDMILL_MAX_KEY_GROUP_BATCH_SIZE = + "windmill_max_key_group_batch_size"; + private static final String WINDMILL_MAX_KEY_GROUP_BATCH_TIME_MS = + "windmill_max_key_group_batch_time_ms"; + private static final String WINDMILL_MAX_KEY_GROUP_BATCH_SINK_BYTES = + "windmill_max_key_group_batch_sink_bytes"; + + public abstract int maxKeyGroupBatchSize(); + + public abstract long maxKeyGroupBatchTimeNanos(); + + public abstract boolean multiKeyBundleEnabled(); + + public abstract long maxKeyGroupBatchSinkBytes(); + + public static Builder builder() { + return new AutoValue_MultiKeyBundleOptions.Builder(); + } + + public static MultiKeyBundleOptions fromOptions(PipelineOptions options) { + int maxKeyGroupBatchSize = + tryParseInt( + ExperimentalOptions.getExperimentValue(options, WINDMILL_MAX_KEY_GROUP_BATCH_SIZE), + 100, + WINDMILL_MAX_KEY_GROUP_BATCH_SIZE); + + long batchTimeMs = + tryParseLong( + ExperimentalOptions.getExperimentValue(options, WINDMILL_MAX_KEY_GROUP_BATCH_TIME_MS), + 100, + WINDMILL_MAX_KEY_GROUP_BATCH_TIME_MS); + + boolean multiKeyBundleEnabled = + ExperimentalOptions.hasExperiment(options, UNSTABLE_ENABLE_MULTI_KEY_BUNDLE); + + long maxKeyGroupBatchSinkBytes = + tryParseLong( + ExperimentalOptions.getExperimentValue( + options, WINDMILL_MAX_KEY_GROUP_BATCH_SINK_BYTES), + StreamingDataflowWorker.MAX_SINK_BYTES, + WINDMILL_MAX_KEY_GROUP_BATCH_SINK_BYTES); + + return builder() + .setMaxKeyGroupBatchSize(maxKeyGroupBatchSize) + .setMaxKeyGroupBatchTimeNanos(TimeUnit.MILLISECONDS.toNanos(batchTimeMs)) + .setMultiKeyBundleEnabled(multiKeyBundleEnabled) + .setMaxKeyGroupBatchSinkBytes(maxKeyGroupBatchSinkBytes) + .build(); + } + + private static int tryParseInt(@Nullable String value, int defaultValue, String experimentName) { + if (value == null) { + return defaultValue; + } + try { + return Integer.parseInt(value); + } catch (NumberFormatException e) { + LOG.warn( + "Failed to parse experiment {} value '{}' as integer, falling back to default: {}", + experimentName, + value, + defaultValue, + e); + return defaultValue; + } + } + + private static long tryParseLong( + @Nullable String value, long defaultValue, String experimentName) { + if (value == null) { + return defaultValue; + } + try { + return Long.parseLong(value); + } catch (NumberFormatException e) { + LOG.warn( + "Failed to parse experiment {} value '{}' as long, falling back to default: {}", + experimentName, + value, + defaultValue, + e); + return defaultValue; + } + } + + @AutoValue.Builder + public abstract static class Builder { + + public abstract Builder setMaxKeyGroupBatchSize(int size); + + public abstract Builder setMaxKeyGroupBatchTimeNanos(long nanos); + + public abstract Builder setMultiKeyBundleEnabled(boolean enabled); + + public abstract Builder setMaxKeyGroupBatchSinkBytes(long bytes); + + public abstract MultiKeyBundleOptions build(); + } +} 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 180dda153bb6..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 @@ -179,9 +179,6 @@ public final class StreamingDataflowWorker { // Experiment make the monitor within BoundedQueueExecutor fair public static final String BOUNDED_QUEUE_EXECUTOR_USE_FAIR_MONITOR_EXPERIMENT = "windmill_bounded_queue_executor_use_fair_monitor"; - // Don't use. Experiment guarding multi key bundles. The feature is work in progress and - // incomplete. - private static final String UNSTABLE_ENABLE_MULTI_KEY_BUNDLE = "unstable_enable_multi_key_bundle"; private final WindmillStateCache stateCache; private AtomicReference statusPages = new AtomicReference<>(); @@ -211,7 +208,7 @@ public final class StreamingDataflowWorker { private StreamingDataflowWorker( WindmillServerStub windmillServer, long clientId, - ComputationConfig.Fetcher configFetcher, + Fetcher configFetcher, ComputationStateCache computationStateCache, WindmillStateCache windmillStateCache, BoundedQueueExecutor workUnitExecutor, @@ -228,7 +225,8 @@ private StreamingDataflowWorker( GrpcWindmillStreamFactory windmillStreamFactory, ScheduledExecutorService activeWorkRefreshExecutorFn, ConcurrentMap stageInfoMap, - @Nullable GrpcDispatcherClient dispatcherClient) { + @Nullable GrpcDispatcherClient dispatcherClient, + MultiKeyBundleOptions multiKeyBundleOptions) { // Register standard file systems. FileSystems.setDefaultPipelineOptions(options); this.configFetcher = configFetcher; @@ -257,6 +255,7 @@ private StreamingDataflowWorker( this.streamingWorkScheduler = StreamingWorkScheduler.create( options, + multiKeyBundleOptions, clock, readerCache, mapTaskExecutorFactory, @@ -627,7 +626,8 @@ public static StreamingDataflowWorker fromOptions(DataflowWorkerHarnessOptions o ConcurrentMap stageInfo = new ConcurrentHashMap<>(); StreamingCounters streamingCounters = StreamingCounters.create(); WorkUnitClient dataflowServiceClient = new DataflowWorkUnitClient(options, LOG); - BoundedQueueExecutor workExecutor = createWorkUnitExecutor(options); + MultiKeyBundleOptions multiKeyBundleOptions = MultiKeyBundleOptions.fromOptions(options); + BoundedQueueExecutor workExecutor = createWorkUnitExecutor(options, multiKeyBundleOptions); ScheduledExecutorService commitFinalizerCleanupExecutor = Executors.newScheduledThreadPool( 1, @@ -726,7 +726,8 @@ public static StreamingDataflowWorker fromOptions(DataflowWorkerHarnessOptions o Executors.newSingleThreadScheduledExecutor( new ThreadFactoryBuilder().setNameFormat("RefreshWork").build()), stageInfo, - configFetcherComputationStateCacheAndWindmillClient.windmillDispatcherClient()); + configFetcherComputationStateCacheAndWindmillClient.windmillDispatcherClient(), + multiKeyBundleOptions); } /** @@ -876,7 +877,8 @@ static StreamingDataflowWorker forTesting( StreamingCounters streamingCounters, WindmillStubFactoryFactory stubFactory) { ConcurrentMap stageInfo = new ConcurrentHashMap<>(); - BoundedQueueExecutor workExecutor = createWorkUnitExecutor(options); + MultiKeyBundleOptions multiKeyBundleOptions = MultiKeyBundleOptions.fromOptions(options); + BoundedQueueExecutor workExecutor = createWorkUnitExecutor(options, multiKeyBundleOptions); ScheduledExecutorService commitFinalizerCleanupExecutor = Executors.newScheduledThreadPool( 1, @@ -990,7 +992,8 @@ static StreamingDataflowWorker forTesting( : windmillStreamFactory.build(), executorSupplier.apply("RefreshWork"), stageInfo, - grpcDispatcherClient); + grpcDispatcherClient, + multiKeyBundleOptions); } private static GrpcWindmillStreamFactory.Builder createGrpcwindmillStreamFactoryBuilder( @@ -1020,11 +1023,11 @@ private static JobHeader createJobHeader(DataflowWorkerHarnessOptions options, l .build(); } - private static BoundedQueueExecutor createWorkUnitExecutor(DataflowWorkerHarnessOptions options) { + private static BoundedQueueExecutor createWorkUnitExecutor( + DataflowWorkerHarnessOptions options, MultiKeyBundleOptions multiKeyBundleOptions) { boolean useFairMonitor = DataflowRunner.hasExperiment(options, BOUNDED_QUEUE_EXECUTOR_USE_FAIR_MONITOR_EXPERIMENT); - boolean useKeyGroupWorkQueue = - DataflowRunner.hasExperiment(options, UNSTABLE_ENABLE_MULTI_KEY_BUNDLE); + boolean useKeyGroupWorkQueue = multiKeyBundleOptions.multiKeyBundleEnabled(); return new BoundedQueueExecutor( chooseMaxThreads(options), THREAD_EXPIRATION_TIME_SEC, @@ -1206,9 +1209,14 @@ private void onCompleteCommit(CompleteCommit completeCommit) { computationStateCache .getIfPresent(completeCommit.computationId()) .ifPresent( - state -> + state -> { + if (completeCommit.retryableFailure()) { + state.reexecuteActiveWork(completeCommit.shardedKey(), completeCommit.workId()); + } else { state.completeWorkAndScheduleNextWorkForKey( - completeCommit.shardedKey(), completeCommit.workId())); + completeCommit.shardedKey(), completeCommit.workId()); + } + }); } @AutoValue 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 db930d53f76d..f4696752bab1 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 @@ -52,6 +52,7 @@ import org.apache.beam.runners.dataflow.worker.counters.NameContext; 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.KeyCommitTooLargeException; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; @@ -79,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; @@ -187,7 +189,7 @@ public class StreamingModeExecutionContext // Key switch listener to delegate MDC logging context and thread name updates public interface KeyTransitionListener { - void onKeyTransition(Work oldWork, Work newWork); + void onKeyTransition(@Nullable Work oldWork, Work newWork); } @SuppressWarnings("UnusedVariable") @@ -203,6 +205,10 @@ public interface KeyTransitionListener { private long stateBytesRead = 0; private final String sourceBytesProcessCounterName; + private final MultiKeyBundleOptions multiKeyBundleOptions; + private int workItemsPolled = 0; + private long bundleStartTimeNanos = 0; + public StreamingModeExecutionContext( CounterFactory counterFactory, String computationId, @@ -222,6 +228,7 @@ public StreamingModeExecutionContext( StreamingCounters streamingCounters, FailureTracker failureTracker, String sourceBytesProcessCounterName, + MultiKeyBundleOptions multiKeyBundleOptions, SideInputStateFetcherFactory sideInputStateFetcherFactory) { super( counterFactory, @@ -244,7 +251,10 @@ public StreamingModeExecutionContext( this.streamingCounters = checkNotNull(streamingCounters); this.failureTracker = checkNotNull(failureTracker); this.sourceBytesProcessCounterName = checkNotNull(sourceBytesProcessCounterName); - this.sideInputStateFetcherFactory = sideInputStateFetcherFactory; + this.sideInputStateFetcherFactory = checkNotNull(sideInputStateFetcherFactory); + + this.multiKeyBundleOptions = checkNotNull(multiKeyBundleOptions); + StreamingGlobalConfig config = globalConfigHandle.getConfig(); this.operationalLimits = config.operationalLimits(); this.windmillTagEncoding = @@ -332,7 +342,6 @@ public void reset() { public void start( Work work, - WindmillStateReader stateReader, WorkExecutor workExecutor, BoundedQueueExecutor workQueueExecutor, BoundedQueueExecutorWorkHandle budgetHandle, @@ -349,11 +358,14 @@ public void start( this.budgetHandle = budgetHandle; this.keyTransitionListener = keyTransitionListener; + this.workItemsPolled = 1; + this.bundleStartTimeNanos = System.nanoTime(); + StreamingGlobalConfig config = globalConfigHandle.getConfig(); // Snapshot the limits for entire bundle processing. this.operationalLimits = config.operationalLimits(); - startForNewKey(work, stateReader); + startForNewKey(work); } private @Nullable Object decodeKey(Work work) throws CoderException { @@ -701,6 +713,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); @@ -739,12 +762,56 @@ private final long computeSourceBytesProcessed(String sourceBytesCounterName) { .orElse(0L); } - public boolean advance() { - // TODO: get more work from workQueueExecutor and merge into the bundle here + public boolean advance() throws CoderException { + if (!multiKeyBundleOptions.multiKeyBundleEnabled()) { + return false; + } + + Work activeWork = checkStateNotNull(work); + BoundedQueueExecutor executor = checkStateNotNull(workQueueExecutor); + BoundedQueueExecutorWorkHandle handle = checkStateNotNull(budgetHandle); + + if (workIsFailed()) { + throw new WorkItemCancelledException(activeWork.getWorkItem().getShardingKey()); + } + + if (activeWork.getKeyGroup().equals(Work.KeyGroup.DEFAULT) + || activeWork.isMultiKeyBatchingDisabled() + || shouldStopBatching()) { + return false; + } + + @Nullable + ExecutableWork additionalWork = + executor.pollWork(computationId, activeWork.getKeyGroup(), handle); + if (additionalWork != null) { + flushStateInternal(); + Work newWork = additionalWork.work(); + ++workItemsPolled; + checkStateNotNull(keyTransitionListener).onKeyTransition(activeWork, newWork); + startForNewKey(newWork); + return true; + } + return false; } - private void startForNewKey(Work newWork, WindmillStateReader reader) throws CoderException { + private boolean shouldStopBatching() { + // stop batching if the previous work item requested truncation + if (getOutputBuilder().getExceedsMaxWorkItemCommitBytes()) { + return true; + } + if (workItemsPolled >= multiKeyBundleOptions.maxKeyGroupBatchSize()) { + return true; + } + long elapsedNanos = System.nanoTime() - bundleStartTimeNanos; + if (elapsedNanos >= multiKeyBundleOptions.maxKeyGroupBatchTimeNanos()) { + return true; + } + return getBytesSinked() >= multiKeyBundleOptions.maxKeyGroupBatchSinkBytes(); + } + + private void startForNewKey(Work newWork) throws CoderException { newWork.setState(Work.State.PROCESSING); if (keyTransitionListener != null && this.work != null && this.work != newWork) { keyTransitionListener.onKeyTransition(this.work, newWork); @@ -779,8 +846,8 @@ private void startForNewKey(Work newWork, WindmillStateReader reader) throws Cod WindmillStateCache.ForKey cacheForKey = stateCache.forKey( getComputationKey(), newWork.getWorkItem().getCacheToken(), getWorkToken()); - this.activeStateReader = reader; - startStepContexts(reader, processingTime, cacheForKey, newWork.watermarks()); + this.activeStateReader = newWork.createWindmillStateReader(this::workIsFailed); + startStepContexts(this.activeStateReader, processingTime, cacheForKey, newWork.watermarks()); } } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/KeyTokenInvalidException.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java similarity index 50% rename from runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/KeyTokenInvalidException.java rename to runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java index 29b16b71883f..9cb4a1c0a5be 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/KeyTokenInvalidException.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkCancellingException.java @@ -17,21 +17,31 @@ */ package org.apache.beam.runners.dataflow.worker; -import javax.annotation.Nullable; +import org.checkerframework.checker.nullness.qual.Nullable; -/** Indicates that the key token was invalid when data was attempted to be fetched. */ -public class KeyTokenInvalidException extends RuntimeException { - public KeyTokenInvalidException(String key) { - super("Unable to fetch data due to token mismatch for key " + key); +/** + * Indicates that the work is no longer valid and should be canceled. It is thrown as a signal for + * upper layers to mark the work as failed. This is different from WorkItemCancelledException, which + * is thrown after marking the work as failed. + */ +public class WorkCancellingException extends RuntimeException { + + public WorkCancellingException(long sharding_key) { + super("Work cancelling exception for key " + sharding_key); + } + + public WorkCancellingException(Throwable cause) { + super(cause); } - /** Returns whether an exception was caused by a {@link KeyTokenInvalidException}. */ - public static boolean isKeyTokenInvalidException(@Nullable Throwable t) { - while (t != null) { - if (t instanceof KeyTokenInvalidException) { + /** Returns whether an exception was caused by a {@link WorkCancellingException}. */ + public static boolean isWorkCancellingException(Throwable t) { + @Nullable Throwable throwable = t; + while (throwable != null) { + if (throwable instanceof WorkCancellingException) { return true; } - t = t.getCause(); + throwable = throwable.getCause(); } return false; } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkItemCancelledException.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkItemCancelledException.java index a12a5075c5ee..8053767491bd 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkItemCancelledException.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WorkItemCancelledException.java @@ -17,31 +17,14 @@ */ package org.apache.beam.runners.dataflow.worker; -/** Indicates that the work item was cancelled and should not be retried. */ -@SuppressWarnings({ - "nullness" // TODO(https://github.com/apache/beam/issues/20497) -}) +/** + * Indicates that the work item was canceled. When this is thrown, the work is already marked as + * failed. This is different from WorkItemCancellingException which is thrown before marking work as + * failed. + */ public class WorkItemCancelledException extends RuntimeException { + public WorkItemCancelledException(long sharding_key) { super("Work item cancelled for key " + sharding_key); } - - public WorkItemCancelledException(String message, Throwable cause) { - super(message, cause); - } - - public WorkItemCancelledException(Throwable cause) { - super(cause); - } - - /** Returns whether an exception was caused by a {@link WorkItemCancelledException}. */ - public static boolean isWorkItemCancelledException(Throwable t) { - while (t != null) { - if (t instanceof WorkItemCancelledException) { - return true; - } - t = t.getCause(); - } - return false; - } } 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 e430f6c8f638..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 @@ -88,6 +88,11 @@ static ActiveWorkState create(WindmillStateCache.ForComputation computationState return new ActiveWorkState(new HashMap<>(), computationStateCache); } + synchronized @Nullable ExecutableWork getActiveWork(ShardedKey shardedKey, WorkId workId) { + LinkedHashMap workQueue = activeWork.get(shardedKey.shardingKey()); + return workQueue == null ? null : workQueue.get(workId); + } + @VisibleForTesting static ActiveWorkState forTesting( Map> activeWork, 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 1ca534966947..20661aae0a04 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,8 +17,13 @@ */ 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 */ -public interface BoundedQueueExecutorWorkHandle {} +public interface BoundedQueueExecutorWorkHandle { + // Returns all work that are tracked by the handle + ImmutableList getWorkBatch(); +} 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 3886d4fbc01b..8020eda1b25d 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 @@ -131,6 +131,13 @@ public void completeWorkAndScheduleNextWorkForKey(ShardedKey shardedKey, WorkId .ifPresent(this::forceExecute); } + public void reexecuteActiveWork(ShardedKey shardedKey, WorkId workId) { + ExecutableWork activeWork = activeWorkState.getActiveWork(shardedKey, workId); + if (activeWork != null) { + forceExecute(activeWork); + } + } + public void invalidateStuckCommits(Instant stuckCommitDeadline) { activeWorkState.invalidateStuckCommits( stuckCommitDeadline, this::completeWorkAndScheduleNextWorkForKey); 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 b8ee42c8ef88..9391b842f038 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 @@ -25,7 +25,6 @@ import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext; import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext.KeyTransitionListener; import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; -import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateReader; import org.apache.beam.sdk.annotations.Internal; import org.apache.beam.sdk.coders.Coder; import org.slf4j.Logger; @@ -53,7 +52,7 @@ public static ComputationWorkExecutor.Builder builder() { public abstract DataflowWorkExecutor workExecutor(); - abstract StreamingModeExecutionContext context(); + public abstract StreamingModeExecutionContext context(); public abstract Optional> keyCoder(); @@ -64,7 +63,6 @@ public static ComputationWorkExecutor.Builder builder() { */ public final StreamingModeExecutionContext executeWork( Work work, - WindmillStateReader stateReader, BoundedQueueExecutor workQueueExecutor, BoundedQueueExecutorWorkHandle budgetHandle, KeyTransitionListener keyTransitionListener) @@ -72,7 +70,6 @@ public final StreamingModeExecutionContext executeWork( context() .start( work, - stateReader, workExecutor(), workQueueExecutor, budgetHandle, 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 252a16a38bc9..51dd8ee045c8 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 @@ -36,6 +36,7 @@ import org.apache.beam.repackaged.core.org.apache.commons.lang3.tuple.Pair; import org.apache.beam.runners.dataflow.worker.ActiveMessageMetadata; import org.apache.beam.runners.dataflow.worker.DataflowExecutionStateSampler; +import org.apache.beam.runners.dataflow.worker.WorkCancellingException; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GlobalData; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.GlobalDataRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.KeyedGetDataRequest; @@ -44,7 +45,6 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution.ActiveLatencyBreakdown; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution.ActiveLatencyBreakdown.ActiveElementMetadata; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution.ActiveLatencyBreakdown.Distribution; -import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution.State; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItem; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItemCommitRequest; import org.apache.beam.runners.dataflow.worker.windmill.client.commits.Commit; @@ -83,10 +83,15 @@ 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); private final boolean drainMode; + private ImmutableList getWorkStreamLatencies; private Work( WorkItem workItem, @@ -94,7 +99,8 @@ private Work( Watermarks watermarks, ProcessingContext processingContext, boolean drainMode, - Supplier clock) { + Supplier clock, + ImmutableList getWorkStreamLatencies) { this.shardedKey = ShardedKey.create(workItem.getKey(), workItem.getShardingKey()); this.workItem = workItem; this.serializedWorkItemSize = serializedWorkItemSize; @@ -118,6 +124,7 @@ private Work( + Long.toHexString(workItem.getWorkToken()); this.currentState = TimedState.initialState(startTime); this.isFailed = false; + this.getWorkStreamLatencies = getWorkStreamLatencies; } public static Work create( @@ -126,9 +133,16 @@ public static Work create( Watermarks watermarks, ProcessingContext processingContext, boolean drainMode, - Supplier clock) { + Supplier clock, + ImmutableList getWorkStreamLatencies) { return new Work( - workItem, serializedWorkItemSize, watermarks, processingContext, drainMode, clock); + workItem, + serializedWorkItemSize, + watermarks, + processingContext, + drainMode, + clock, + getWorkStreamLatencies); } public static ProcessingContext createProcessingContext( @@ -205,11 +219,31 @@ public ShardedKey getShardedKey() { } public Optional fetchKeyedState(KeyedGetDataRequest keyedGetDataRequest) { - return processingContext.fetchKeyedState(keyedGetDataRequest); + try { + Optional response = + processingContext.fetchKeyedState(keyedGetDataRequest); + if (response.isPresent() && response.get().getFailed()) { + // Work is not valid in backend anymore. + this.setFailed(); + } + return response; + } catch (RuntimeException e) { + if (WorkCancellingException.isWorkCancellingException(e)) { + this.setFailed(); + } + throw e; + } } public GlobalData fetchSideInput(GlobalDataRequest request) { - return processingContext.getDataClient().getSideInputData(request); + try { + return processingContext.getDataClient().getSideInputData(request); + } catch (RuntimeException e) { + if (WorkCancellingException.isWorkCancellingException(e)) { + this.setFailed(); + } + throw e; + } } public String backendWorkerToken() { @@ -293,8 +327,8 @@ public Consumer workCommitter() { return processingContext.workCommitter(); } - public WindmillStateReader createWindmillStateReader() { - return WindmillStateReader.forWork(this); + public WindmillStateReader createWindmillStateReader(Supplier workIsFailed) { + return WindmillStateReader.forWork(this, workIsFailed); } @Override @@ -302,11 +336,13 @@ public WorkId id() { return id; } - public void recordGetWorkStreamLatencies( - ImmutableList getWorkStreamLatencies) { - for (LatencyAttribution latency : getWorkStreamLatencies) { - totalDurationPerState.put( - latency.getState(), Duration.millis(latency.getTotalDurationMillis())); + public void recordGetWorkStreamLatencies() { + if (!getWorkStreamLatencies.isEmpty()) { + for (LatencyAttribution latency : getWorkStreamLatencies) { + totalDurationPerState.put( + latency.getState(), Duration.millis(latency.getTotalDurationMillis())); + } + this.getWorkStreamLatencies = ImmutableList.of(); } } @@ -365,6 +401,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/BoundedQueueExecutor.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/BoundedQueueExecutor.java index 8964246c1160..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 @@ -20,6 +20,8 @@ import static org.apache.beam.sdk.util.Preconditions.checkArgumentNotNull; import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkArgument; +import java.util.ArrayList; +import java.util.List; import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.LinkedBlockingQueue; import java.util.concurrent.ThreadFactory; @@ -30,9 +32,9 @@ 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.Work; -import org.apache.beam.runners.dataflow.worker.streaming.Work.KeyGroup; 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; @@ -260,7 +262,7 @@ final class BoundedQueueExecutorWorkHandleImpl implements BoundedQueueExecutorWorkHandle, AutoCloseable { @GuardedBy("this") - private int elements; + private final List workBatch; @GuardedBy("this") private long bytes; @@ -268,16 +270,17 @@ final class BoundedQueueExecutorWorkHandleImpl @GuardedBy("this") private boolean closed = false; - private BoundedQueueExecutorWorkHandleImpl(int elements, long bytes) { - checkArgument(elements >= 0 && bytes >= 0); - this.elements = elements; + private BoundedQueueExecutorWorkHandleImpl(Work work, long bytes) { + checkArgument(bytes >= 0); + this.workBatch = new ArrayList<>(); + this.workBatch.add(checkArgumentNotNull(work)); this.bytes = bytes; } /** * Merges the budget from another handle into this handle. * - *

This transfers the budget (elements and bytes) from the {@code other} handle to this + *

This transfers the budget (workBatch and bytes) from the {@code other} handle to this * handle, and marks the {@code other} handle as closed to prevent it from releasing the budget * again if it is closed. */ @@ -287,10 +290,10 @@ public void merge(BoundedQueueExecutorWorkHandleImpl other) { Preconditions.checkState(!closed, "Cannot merge into a closed handle"); synchronized (other) { Preconditions.checkState(!other.closed, "Cannot merge a closed handle"); - this.elements += other.elements; + this.workBatch.addAll(other.workBatch); this.bytes += other.bytes; other.closed = true; - other.elements = 0; + other.workBatch.clear(); other.bytes = 0; } } @@ -300,9 +303,9 @@ public synchronized boolean isClosed() { return closed; } - @VisibleForTesting - synchronized int elements() { - return elements; + @Override + public synchronized ImmutableList getWorkBatch() { + return ImmutableList.copyOf(workBatch); } @VisibleForTesting @@ -314,7 +317,7 @@ synchronized long bytes() { public synchronized void close() { if (closed) return; closed = true; - decrementCounters(this.elements, this.bytes); + decrementCounters(this.workBatch.size(), this.bytes); } } @@ -350,7 +353,7 @@ private void executeMonitorHeld(ExecutableWork work, long workBytes) { bytesOutstanding += workBytes; monitor.leave(); BoundedQueueExecutorWorkHandleImpl handle = - new BoundedQueueExecutorWorkHandleImpl(1, workBytes); + new BoundedQueueExecutorWorkHandleImpl(work.work(), workBytes); try { executor.execute(new QueuedWork(work, handle)); } catch (Throwable t) { @@ -379,14 +382,15 @@ private void executeMonitorHeld(Runnable work) { } @VisibleForTesting - BoundedQueueExecutorWorkHandleImpl createBudgetHandle(int elements, long bytes) { - return new BoundedQueueExecutorWorkHandleImpl(elements, bytes); + BoundedQueueExecutorWorkHandleImpl createBudgetHandle(Work work, long bytes) { + return new BoundedQueueExecutorWorkHandleImpl(work, bytes); } public @Nullable ExecutableWork pollWork( String computationId, Work.KeyGroup keyGroup, BoundedQueueExecutorWorkHandle handle) { + checkArgument( + computationId != null && keyGroup != null && !keyGroup.equals(Work.KeyGroup.DEFAULT)); checkArgument(handle instanceof BoundedQueueExecutorWorkHandleImpl); - checkArgument(computationId != null && keyGroup != null && !keyGroup.equals(KeyGroup.DEFAULT)); BoundedQueueExecutorWorkHandleImpl internalHandle = (BoundedQueueExecutorWorkHandleImpl) handle; if (keyGroupWorkQueue == null) { return null; 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/client/commits/CompleteCommit.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java index 6c0a5a98e2ab..7e2be8308954 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java @@ -38,8 +38,13 @@ public abstract class CompleteCommit { public static CompleteCommit create( - String computationId, ShardedKey shardedKey, WorkId workId, CommitStatus status) { - return new AutoValue_CompleteCommit(computationId, shardedKey, workId, status); + String computationId, + ShardedKey shardedKey, + WorkId workId, + CommitStatus status, + boolean retryableFailure) { + return new AutoValue_CompleteCommit( + computationId, shardedKey, workId, status, retryableFailure); } public abstract String computationId(); @@ -49,4 +54,10 @@ public static CompleteCommit create( public abstract WorkId workId(); public abstract CommitStatus status(); + + /** + * If retryableFailure true, the workitem will be retried locally. Used to retry partial work + * failures in multi key bundles. + */ + public abstract boolean retryableFailure(); } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java index d627490dfe98..ffb9b64595c5 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java @@ -159,7 +159,8 @@ private void completeWork( .setCacheToken(workRequest.getCacheToken()) .setWorkToken(workRequest.getWorkToken()) .build(), - Windmill.CommitStatus.OK)); + Windmill.CommitStatus.OK, + /* retryableFailure= */ false)); } } } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java index 83d0dfc6cda4..1a09a75ccb81 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java @@ -145,11 +145,41 @@ private void drainCommitQueue() { } private void failQueuedCommit(Commit commit) { + if (!isRunning.get()) { + // Shutting down, fail everything unconditionally to prevent infinite loops + for (Work w : commit.workBatch()) { + w.setFailed(); + onCommitComplete.accept( + CompleteCommit.create( + commit.computationId(), + w.getShardedKey(), + w.id(), + CommitStatus.ABORTED, + /* retryableFailure= */ false)); + } + return; + } + + // Still running, only fail actually failed work, and request re-execution for valid ones for (Work w : commit.workBatch()) { - w.setFailed(); - onCommitComplete.accept( - CompleteCommit.create( - commit.computationId(), w.getShardedKey(), w.id(), CommitStatus.ABORTED)); + if (w.isFailed()) { + onCommitComplete.accept( + CompleteCommit.create( + commit.computationId(), + w.getShardedKey(), + w.id(), + CommitStatus.ABORTED, + /* retryableFailure= */ false)); + } else { + LOG.debug("Requesting re-execution for valid work {} from failed commit", w.id()); + onCommitComplete.accept( + CompleteCommit.create( + commit.computationId(), + w.getShardedKey(), + w.id(), + CommitStatus.ABORTED, + /* retryableFailure= */ true)); + } } } @@ -227,7 +257,11 @@ private boolean tryAddToCommitBatch(Commit commit, CommitWorkStream.RequestBatch for (Work w : commit.workBatch()) { onCommitComplete.accept( CompleteCommit.create( - commit.computationId(), w.getShardedKey(), w.id(), commitStatus)); + commit.computationId(), + w.getShardedKey(), + w.id(), + commitStatus, + /* retryableFailure= */ false)); } activeCommitBytes.addAndGet(-commit.getSerializedByteSize()); }); @@ -240,7 +274,11 @@ private boolean tryAddToCommitBatch(Commit commit, CommitWorkStream.RequestBatch Work w = commit.workBatch().get(0); onCommitComplete.accept( CompleteCommit.create( - commit.computationId(), w.getShardedKey(), w.id(), commitStatus)); + commit.computationId(), + w.getShardedKey(), + w.id(), + commitStatus, + /* retryableFailure= */ false)); activeCommitBytes.addAndGet(-commit.getSerializedByteSize()); }); } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/getdata/StreamGetDataClient.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/getdata/StreamGetDataClient.java index ab12946ad18b..a9134c677520 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/getdata/StreamGetDataClient.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/getdata/StreamGetDataClient.java @@ -19,6 +19,7 @@ import java.io.PrintWriter; import java.util.function.Function; +import org.apache.beam.runners.dataflow.worker.WorkCancellingException; import org.apache.beam.runners.dataflow.worker.WorkItemCancelledException; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.GetDataStream; @@ -62,7 +63,7 @@ public Windmill.KeyedGetDataResponse getStateData( try (AutoCloseable ignored = getDataMetricTracker.trackStateDataFetchWithThrottling()) { return getDataStream.requestKeyedData(computationId, request); } catch (WindmillStreamShutdownException e) { - throw new WorkItemCancelledException(request.getShardingKey()); + throw new WorkCancellingException(request.getShardingKey()); } catch (Exception e) { throw new GetDataException( "Error occurred fetching state for computation=" @@ -87,7 +88,7 @@ public Windmill.GlobalData getSideInputData(Windmill.GlobalDataRequest request) try (AutoCloseable ignored = getDataMetricTracker.trackSideInputFetchWithThrottling()) { return sideInputGetDataStream.requestGlobalData(request); } catch (WindmillStreamShutdownException e) { - throw new WorkItemCancelledException(e); + throw new WorkCancellingException(e); } catch (Exception e) { throw new GetDataException( "Error occurred fetching side input for tag=" + request.getDataId(), e); diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReader.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReader.java index c609bed4eae0..4ac5d006e35b 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReader.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReader.java @@ -36,8 +36,8 @@ import java.util.function.Supplier; import java.util.stream.Collectors; import javax.annotation.Nullable; -import org.apache.beam.runners.dataflow.worker.KeyTokenInvalidException; import org.apache.beam.runners.dataflow.worker.WindmillTimeUtils; +import org.apache.beam.runners.dataflow.worker.WorkCancellingException; import org.apache.beam.runners.dataflow.worker.WorkItemCancelledException; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; @@ -153,7 +153,7 @@ static WindmillStateReader forTesting( fetchStateFromWindmillFn, key, shardingKey, workToken, () -> null, () -> Boolean.FALSE); } - public static WindmillStateReader forWork(Work work) { + public static WindmillStateReader forWork(Work work, Supplier workItemIsFailed) { return new WindmillStateReader( work::fetchKeyedState, work.getWorkItem().getKey(), @@ -163,7 +163,7 @@ public static WindmillStateReader forWork(Work work) { work.setState(Work.State.READING); return () -> work.setState(Work.State.PROCESSING); }, - work::isFailed); + workItemIsFailed); } private Future stateFuture(StateTag stateTag, @Nullable Coder coder) { @@ -588,7 +588,8 @@ private KeyedGetDataRequest createRequest(Iterable> toFetch) { private void consumeResponse(KeyedGetDataResponse response, Set> toFetch) { bytesRead += response.getSerializedSize(); if (response.getFailed()) { - throw new KeyTokenInvalidException(key.toStringUtf8()); + // upper layers will fail the work on seeing this exception. + throw new WorkCancellingException(shardingKey); } if (!key.equals(response.getKey())) { diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ComputationWorkExecutorFactory.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ComputationWorkExecutorFactory.java index 0c3591102f9c..b51512252e37 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ComputationWorkExecutorFactory.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ComputationWorkExecutorFactory.java @@ -31,6 +31,7 @@ import org.apache.beam.runners.dataflow.worker.DataflowMapTaskExecutorFactory; import org.apache.beam.runners.dataflow.worker.HotKeyLogger; import org.apache.beam.runners.dataflow.worker.IntrinsicMapTaskExecutorFactory; +import org.apache.beam.runners.dataflow.worker.MultiKeyBundleOptions; import org.apache.beam.runners.dataflow.worker.ReaderCache; import org.apache.beam.runners.dataflow.worker.ReaderRegistry; import org.apache.beam.runners.dataflow.worker.SinkRegistry; @@ -86,6 +87,7 @@ final class ComputationWorkExecutorFactory { private final SinkRegistry sinkRegistry; private final DataflowExecutionStateSampler sampler; private final CounterSet pendingDeltaCounters; + private final SideInputStateFetcherFactory sideInputStateFetcherFactory; private final StreamingCounters streamingCounters; private final FailureTracker failureTracker; @@ -104,7 +106,7 @@ final class ComputationWorkExecutorFactory { private final StreamingGlobalConfigHandle globalConfigHandle; private final boolean throwExceptionOnLargeOutput; private final HotKeyLogger hotKeyLogger; - private final SideInputStateFetcherFactory sideInputStateFetcherFactory; + private final MultiKeyBundleOptions multiKeyBundleOptions; ComputationWorkExecutorFactory( DataflowWorkerHarnessOptions options, @@ -117,7 +119,8 @@ final class ComputationWorkExecutorFactory { IdGenerator idGenerator, StreamingGlobalConfigHandle globalConfigHandle, HotKeyLogger hotKeyLogger, - SideInputStateFetcherFactory sideInputStateFetcherFactory) { + SideInputStateFetcherFactory sideInputStateFetcherFactory, + MultiKeyBundleOptions multiKeyBundleOptions) { this.options = options; this.mapTaskExecutorFactory = mapTaskExecutorFactory; this.readerCache = readerCache; @@ -139,6 +142,7 @@ final class ComputationWorkExecutorFactory { hasExperiment(options, THROW_EXCEPTIONS_ON_LARGE_OUTPUT_EXPERIMENT); this.hotKeyLogger = hotKeyLogger; this.sideInputStateFetcherFactory = sideInputStateFetcherFactory; + this.multiKeyBundleOptions = multiKeyBundleOptions; } private static Nodes.ParallelInstructionNode extractReadNode( @@ -297,6 +301,7 @@ private StreamingModeExecutionContext createExecutionContext( streamingCounters, failureTracker, computationState.sourceBytesProcessCounterName(), + multiKeyBundleOptions, sideInputStateFetcherFactory); } 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 990393f51f8f..74783b9e4751 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 @@ -21,6 +21,7 @@ import com.google.api.services.dataflow.model.MapTask; import com.google.auto.value.AutoValue; +import java.util.ArrayList; import java.util.List; import java.util.Map; import java.util.concurrent.ConcurrentMap; @@ -34,6 +35,7 @@ import org.apache.beam.runners.dataflow.worker.DataflowExecutionStateSampler; import org.apache.beam.runners.dataflow.worker.DataflowMapTaskExecutorFactory; import org.apache.beam.runners.dataflow.worker.HotKeyLogger; +import org.apache.beam.runners.dataflow.worker.MultiKeyBundleOptions; import org.apache.beam.runners.dataflow.worker.ReaderCache; import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext; import org.apache.beam.runners.dataflow.worker.StreamingModeExecutionContext.KeyTransitionListener; @@ -55,12 +57,12 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill.LatencyAttribution; import org.apache.beam.runners.dataflow.worker.windmill.client.commits.Commit; import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache; -import org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateReader; import org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.FailureTracker; import org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.WorkFailureProcessor; import org.apache.beam.sdk.annotations.Internal; import org.apache.beam.sdk.fn.IdGenerator; import org.apache.beam.vendor.grpc.v1p69p0.com.google.protobuf.ByteString; +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.checkerframework.checker.nullness.qual.Nullable; import org.joda.time.Instant; @@ -86,6 +88,7 @@ public class StreamingWorkScheduler { private final ConcurrentMap stageInfoMap; private final DataflowExecutionStateSampler sampler; private final BoundedQueueExecutor workExecutor; + private final MultiKeyBundleOptions multiKeyBundleOptions; public StreamingWorkScheduler( Supplier clock, @@ -95,7 +98,8 @@ public StreamingWorkScheduler( StreamingCommitFinalizer commitFinalizer, StreamingCounters streamingCounters, ConcurrentMap stageInfoMap, - DataflowExecutionStateSampler sampler) { + DataflowExecutionStateSampler sampler, + MultiKeyBundleOptions multiKeyBundleOptions) { this.clock = clock; this.workExecutor = workExecutor; this.computationWorkExecutorFactory = computationWorkExecutorFactory; @@ -104,10 +108,12 @@ public StreamingWorkScheduler( this.streamingCounters = streamingCounters; this.stageInfoMap = stageInfoMap; this.sampler = sampler; + this.multiKeyBundleOptions = multiKeyBundleOptions; } public static StreamingWorkScheduler create( DataflowWorkerHarnessOptions options, + MultiKeyBundleOptions multiKeyBundleOptions, Supplier clock, ReaderCache readerCache, DataflowMapTaskExecutorFactory mapTaskExecutorFactory, @@ -137,7 +143,8 @@ public static StreamingWorkScheduler create( idGenerator, globalConfigHandle, hotKeyLogger, - sideInputStateFetcherFactory); + sideInputStateFetcherFactory, + multiKeyBundleOptions); return new StreamingWorkScheduler( clock, @@ -147,7 +154,8 @@ public static StreamingWorkScheduler create( StreamingCommitFinalizer.create(workExecutor, commitFinalizerCleanupExecutor), streamingCounters, stageInfoMap, - sampler); + sampler, + multiKeyBundleOptions); } private static long computeShuffleBytesRead(Windmill.WorkItem workItem) { @@ -167,12 +175,6 @@ private static Windmill.WorkItemCommitRequest.Builder initializeOutputBuilder( .setCacheToken(workItem.getCacheToken()); } - /** Sets the stage name and workId of the Thread executing the {@link Work} for logging. */ - private static void setUpWorkLoggingContext(String workLatencyTrackingId, String computationId) { - setLoggingContextWorkId(workLatencyTrackingId); - setLoggingContextComputation(computationId); - } - private static void setLoggingContextComputation(@Nullable String computationId) { DataflowWorkerLoggingMDC.setStageName(computationId); } @@ -198,8 +200,14 @@ public void scheduleWork( computationState.activateWork( ExecutableWork.create( Work.create( - workItem, serializedWorkItemSize, watermarks, processingContext, drainMode, clock), - (work, handle) -> processWork(computationState, work, getWorkStreamLatencies, handle))); + workItem, + serializedWorkItemSize, + watermarks, + processingContext, + drainMode, + clock, + getWorkStreamLatencies), + (work, handle) -> processWork(computationState, work, handle))); } /** Adds any applied finalize ids to the commit finalizer to have their callbacks executed. */ @@ -213,25 +221,20 @@ public void queueAppliedFinalizeIds(ImmutableList appliedFinalizeIds) { * internally if processing fails due to uncaught {@link Exception}(s). * * @implNote This will block the calling thread during execution of user DoFns. - * @param handle handled to pass to BoundedQueueExecutor.pollWork, currently unused + * @param handle handled to pass to BoundedQueueExecutor.pollWork */ - private void processWork( - ComputationState computationState, - Work work, - ImmutableList getWorkStreamLatencies, - BoundedQueueExecutorWorkHandle handle) { - work.recordGetWorkStreamLatencies(getWorkStreamLatencies); - processWork(computationState, work, handle); - } - private void processWork( ComputationState computationState, Work work, BoundedQueueExecutorWorkHandle handle) { Windmill.WorkItem workItem = work.getWorkItem(); String computationId = computationState.getComputationId(); - work.setProcessingThreadName(Thread.currentThread().getName()); - work.setState(Work.State.PROCESSING); - setUpWorkLoggingContext(work.getLatencyTrackingId(), computationId); LOG.debug("Starting processing for {}:\n{}", computationId, work); + setLoggingContextComputation(computationId); + 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); @@ -248,7 +251,8 @@ private void processWork( } // Execute the user code for the Work batch. - ExecuteWorkResult executeWorkResult = executeWork(work, stageInfo, computationState, handle); + ExecuteWorkResult executeWorkResult = + executeWork(work, stageInfo, computationState, handle, keyTransitionListener); workBatch = executeWorkResult.workBatch(); List workItemCommits = executeWorkResult.workItemCommits(); @@ -259,21 +263,7 @@ private void processWork( recordProcessingStats(workBatch, workItemCommits, executeWorkResult.stateBytesRead()); LOG.debug("Processing done for work batch size: {}", workBatch.size()); } catch (Throwable t) { - // OutOfMemoryError that are caught will be rethrown and trigger jvm termination. - try { - workFailureProcessor.logAndProcessFailure( - computationId, - ExecutableWork.create(work, (retry, h) -> processWork(computationState, retry, h)), - t, - invalidWork -> - computationState.completeWorkAndScheduleNextWorkForKey( - invalidWork.getShardedKey(), invalidWork.id())); - } catch (OutOfMemoryError oom) { - throw oom; - } catch (Throwable t2) { - LOG.warn("Failed to process work failure safely for work {}", work.id(), t2); - throw ExceptionUtils.safeWrapThrowableAsException(t2); - } + handleProcessWorkFailure(computationState, handle.getWorkBatch(), computationId, work, t); } finally { List processedWorkBatch = workBatch != null ? workBatch : ImmutableList.of(work); // Update total processing time counters. Updating in finally clause ensures that @@ -319,7 +309,8 @@ private ExecuteWorkResult executeWork( Work work, StageInfo stageInfo, ComputationState computationState, - BoundedQueueExecutorWorkHandle handle) + BoundedQueueExecutorWorkHandle handle, + KeyTransitionListener keyTransitionListener) throws Exception { ComputationWorkExecutor computationWorkExecutor = computationState @@ -330,19 +321,16 @@ private ExecuteWorkResult executeWork( stageInfo, computationState, work.getLatencyTrackingId())); try { - WindmillStateReader stateReader = work.createWindmillStateReader(); + StreamingModeExecutionContext context = computationWorkExecutor.context(); - KeyTransitionListener keyTransitionListener = createKeyTransitionListener(); + // Blocks while executing work. + computationWorkExecutor.executeWork(work, workExecutor, handle, keyTransitionListener); List workBatch; List workItemCommits; Map> finalizationCallbacks; long stateBytesRead; { - // Blocks while executing work. - StreamingModeExecutionContext context = - computationWorkExecutor.executeWork( - work, stateReader, workExecutor, handle, keyTransitionListener); if (context.workIsFailed()) { throw new WorkItemCancelledException(work.getWorkItem().getShardingKey()); } @@ -398,9 +386,54 @@ private void commitWorkBatch( ComputationState computationState, List workBatch, List workItemCommits) { - checkState(workBatch.size() == 1, "Expected single-key work batch, got: " + workBatch.size()); - checkState(workBatch.size() == workItemCommits.size()); - commitSingleKeyWork(computationState, workBatch.get(0), workItemCommits.get(0)); + if (workBatch.isEmpty()) { + return; + } + if (workBatch.size() > 1 || multiKeyBundleOptions.multiKeyBundleEnabled()) { + commitMultiKeyWorkBatch(computationState, workBatch, workItemCommits); + } else { + commitSingleKeyWork(computationState, workBatch.get(0), workItemCommits.get(0)); + } + } + + private void commitMultiKeyWorkBatch( + ComputationState computationState, + List workBatch, + List workItemCommits) { + Preconditions.checkState(!workBatch.isEmpty()); + Preconditions.checkState(workBatch.size() == workItemCommits.size()); + + Windmill.MultiKeyWorkItemCommitRequest.Builder multiKeyBuilder = + Windmill.MultiKeyWorkItemCommitRequest.newBuilder(); + + Work primaryWork = workBatch.get(0); + Work.KeyGroup keyGroup = primaryWork.getKeyGroup(); + multiKeyBuilder.setKeyGroup( + Windmill.Uint128Proto.newBuilder().setHigh(keyGroup.high()).setLow(keyGroup.low()).build()); + + for (int i = 0; i < workBatch.size(); i++) { + Windmill.WorkItemCommitRequest commit = workItemCommits.get(i); + Work w = workBatch.get(i); + multiKeyBuilder.addRequests( + commit + .toBuilder() + .addAllPerWorkItemLatencyAttributions(w.getLatencyAttributions(sampler)) + .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); + } + + // Package and submit the commit batch transactionally + primaryWork + .workCommitter() + .accept( + Commit.createMultiKey( + multiKeyCommitRequest, computationState, ImmutableList.copyOf(workBatch))); } private void commitSingleKeyWork( @@ -414,12 +447,40 @@ private void commitSingleKeyWork( work.queueCommit(commitRequestWithAttributions, computationState); } + private void handleProcessWorkFailure( + ComputationState computationState, + List failedBatch, + String computationId, + Work primaryWork, + Throwable t) { + try { + List executableWorks = new ArrayList<>(); + for (Work w : failedBatch) { + executableWorks.add( + ExecutableWork.create(w, (retry, h) -> processWork(computationState, retry, h))); + } + + workFailureProcessor.logAndProcessFailureBatch( + computationId, + executableWorks, + t, + invalidWork -> + computationState.completeWorkAndScheduleNextWorkForKey( + invalidWork.getShardedKey(), invalidWork.id())); + } catch (OutOfMemoryError oom) { + throw oom; + } catch (Throwable t2) { + LOG.warn("Failed to process work failure safely for work {}", primaryWork.id(), t2); + throw ExceptionUtils.safeWrapThrowableAsException(t2); + } + } + private void recordProcessingTime( - StageInfo stageInfo, List worksToCleanup, long processingStartTimeNanos) { + StageInfo stageInfo, List workBatch, long processingStartTimeNanos) { long processingTimeMsecs = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - processingStartTimeNanos); stageInfo.totalProcessingMsecs().addValue(processingTimeMsecs); - if (anyWorkHasTimers(worksToCleanup)) { + if (anyWorkHasTimers(workBatch)) { // Attribute all the processing to timers if the work item contains any timers. // Tests show that work items rarely contain both timers and message bundles. It should // be a fairly close approximation. @@ -435,9 +496,15 @@ private static boolean anyWorkHasTimers(List works) { private KeyTransitionListener createKeyTransitionListener() { return (oldWork, newWork) -> { + newWork.recordGetWorkStreamLatencies(); + newWork.setState(Work.State.PROCESSING); setLoggingContextWorkId(newWork.getLatencyTrackingId()); - newWork.setProcessingThreadName(oldWork.getProcessingThreadName()); - oldWork.setProcessingThreadName(""); + if (oldWork != null) { + newWork.setProcessingThreadName(oldWork.getProcessingThreadName()); + oldWork.setProcessingThreadName(""); + } else { + newWork.setProcessingThreadName(Thread.currentThread().getName()); + } }; } @@ -461,4 +528,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 18c8e9b8d83c..2f28f19fa465 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 @@ -17,17 +17,17 @@ */ package org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures; +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.KeyTokenInvalidException; -import org.apache.beam.runners.dataflow.worker.WorkItemCancelledException; 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.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; @@ -99,28 +99,41 @@ private static boolean isOutOfMemoryError(@Nullable Throwable t) { return false; } - /** - * Processes failures caused by thrown exceptions that occur during execution of {@link Work}. May - * attempt to retry execution of the {@link Work} or drop it if it is invalid. - */ - public void logAndProcessFailure( + public void logAndProcessFailureBatch( String computationId, - ExecutableWork executableWork, + List executableWorks, Throwable t, Consumer onInvalidWork) throws Throwable { - switch (evaluateRetry(computationId, executableWork.work(), t)) { - 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()); - break; - case RETRY_LOCALLY: - // Try again after some delay and at the end of the queue to avoid a tight loop. - executeWithDelay(retryLocallyDelayMs, executableWork); - break; - case RETHROW_THROWABLE: - throw t; + List worksToRetryLocally = new java.util.ArrayList<>(); + + for (ExecutableWork executableWork : executableWorks) { + switch (evaluateRetry(computationId, executableWork.work(), t)) { + 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()); + break; + case RETRY_LOCALLY: + // Try again after some delay and at the end of the queue to avoid a tight loop. + worksToRetryLocally.add(executableWork); + break; + case RETHROW_THROWABLE: + throw t; + } + } + + executeWithDelay(worksToRetryLocally); + } + + private void executeWithDelay(List worksToRetryLocally) { + if (!worksToRetryLocally.isEmpty()) { + // Sleep ONCE for the entire batch delay to avoid sequential thread blocks + Uninterruptibles.sleepUninterruptibly(retryLocallyDelayMs, TimeUnit.MILLISECONDS); + for (ExecutableWork ew : worksToRetryLocally) { + workUnitExecutor.forceExecute(ew, ew.work().getSerializedWorkItemSize()); + } } } @@ -131,12 +144,6 @@ private String tryToDumpHeap() { .orElseGet(() -> "not written"); } - private void executeWithDelay(long delayMs, ExecutableWork executableWork) { - Uninterruptibles.sleepUninterruptibly(delayMs, TimeUnit.MILLISECONDS); - workUnitExecutor.forceExecute( - executableWork, executableWork.work().getSerializedWorkItemSize()); - } - private enum RetryEvaluation { DO_NOT_RETRY, RETRY_LOCALLY, @@ -144,23 +151,24 @@ private enum RetryEvaluation { } private RetryEvaluation evaluateRetry(String computationId, Work work, Throwable t) { - @Nullable final Throwable cause = t.getCause(); - Throwable parsedException = (t instanceof UserCodeException && cause != null) ? cause : t; - if (KeyTokenInvalidException.isKeyTokenInvalidException(parsedException)) { + if (work.isFailed()) { LOG.debug( - "Execution of work for computation '{}' on sharding key '{}' failed due to token expiration. " - + "Work will not be retried locally.", + "Execution of work for computation '{}' on sharding key '{}' failed. " + + "Work is already marked as failed, not retrying locally.", computationId, work.getWorkItem().getShardingKey()); return RetryEvaluation.DO_NOT_RETRY; } - if (WorkItemCancelledException.isWorkItemCancelledException(parsedException)) { - LOG.debug( - "Execution of work for computation '{}' on sharding key '{}' failed. " - + "Work will not be retried locally.", + @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.DO_NOT_RETRY; + return RetryEvaluation.RETRY_LOCALLY; } LastExceptionDataProvider.reportException(parsedException); diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/KeyTokenInvalidExceptionTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/KeyTokenInvalidExceptionTest.java deleted file mode 100644 index 1eb2871e8cd3..000000000000 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/KeyTokenInvalidExceptionTest.java +++ /dev/null @@ -1,39 +0,0 @@ -/* - * 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; - -import static org.junit.Assert.assertFalse; -import static org.junit.Assert.assertTrue; - -import org.junit.Test; -import org.junit.runner.RunWith; -import org.junit.runners.JUnit4; - -/** Tests for {@link KeyTokenInvalidException}. */ -@RunWith(JUnit4.class) -public final class KeyTokenInvalidExceptionTest { - @Test - public void testIsKeyTokenInvalidException() throws Exception { - KeyTokenInvalidException exception = new KeyTokenInvalidException("test"); - RuntimeException keyTokenCauseException = new RuntimeException("key token cause", exception); - assertTrue(KeyTokenInvalidException.isKeyTokenInvalidException(exception)); - assertTrue(KeyTokenInvalidException.isKeyTokenInvalidException(keyTokenCauseException)); - assertFalse( - KeyTokenInvalidException.isKeyTokenInvalidException(new RuntimeException("non key token"))); - } -} 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 1e7f3b01f005..eb3c32f63cae 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 @@ -141,6 +141,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; @@ -155,6 +156,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; @@ -291,6 +293,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(); @@ -372,7 +378,8 @@ private static ExecutableWork createMockWork( Work.createProcessingContext( computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now), + Instant::now, + ImmutableList.of()), (work, handle) -> { processWorkFn.accept(work); }); @@ -901,7 +908,7 @@ private ByteString addPaneTag(PaneInfo paneInfo, byte[] windowBytes) throws IOEx } private DataflowWorkerHarnessOptions createTestingPipelineOptions(String... args) { - List argsList = Lists.newArrayList(args); + List argsList = new ArrayList<>(Arrays.asList(args)); if (streamingEngine) { argsList.add("--experiments=enable_streaming_engine"); } @@ -1252,9 +1259,8 @@ public void testNumberOfWorkerHarnessThreadsIsHonored() throws Exception { } @Test - public void testKeyTokenInvalidException() throws Exception { - if (streamingEngine) { - // TODO: This test needs to be adapted to work with streamingEngine=true. + public void testMultiKeyCommit_success() throws Exception { + if (!streamingEngine) { return; } KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); @@ -1262,30 +1268,530 @@ public void testKeyTokenInvalidException() throws Exception { List instructions = Arrays.asList( makeSourceInstruction(kvCoder), - makeDoFnInstruction(new KeyTokenInvalidFn(), 0, kvCoder), + makeDoFnInstruction(new WorkDoFn(), 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) + .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: 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\"" + + " }" + + " }" + + " }" + + " work {" + + " key: \"key3\"" + + " sharding_key: 3" + + " 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 batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + server - .whenGetWorkCalled() - .thenReturn(makeInput(0, 0, DEFAULT_KEY_STRING, DEFAULT_SHARDING_KEY)); + .whenGetDataCalled() + .answerByDefault( + request -> { + Windmill.GetDataResponse.Builder builder = Windmill.GetDataResponse.newBuilder(); + for (ComputationGetDataRequest compRequest : request.getRequestsList()) { + ComputationGetDataResponse.Builder compBuilder = + builder.addDataBuilder().setComputationId(compRequest.getComputationId()); + for (KeyedGetDataRequest keyRequest : compRequest.getRequestsList()) { + KeyedGetDataResponse.Builder keyBuilder = + compBuilder + .addDataBuilder() + .setKey(keyRequest.getKey()) + .setShardingKey(keyRequest.getShardingKey()); + keyBuilder.addAllValues(keyRequest.getValuesToFetchList()); + keyBuilder.addAllBags(keyRequest.getBagsToFetchList()); + keyBuilder.addAllWatermarkHolds(keyRequest.getWatermarkHoldsToFetchList()); + } + } + return builder.build(); + }); + + server.whenGetWorkCalled().thenReturn(batchInput); + + Map result = server.waitForAndGetCommits(3); + + assertEquals(3, result.size()); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(1, multiKeyCommits.size()); + Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = multiKeyCommits.get(0); + assertEquals(3, multiKeyCommit.getRequestsCount()); + assertEquals(1, multiKeyCommit.getRequests(0).getWorkToken()); + assertEquals(2, multiKeyCommit.getRequests(1).getWorkToken()); + assertEquals(3, multiKeyCommit.getRequests(2).getWorkToken()); + + worker.stop(); + } + + @Test + public void testMultiKeyCommit_batchLimitExceeded() 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().setInstructions(instructions).publishCounters().build()); + makeWorker( + defaultWorkerParams( + "--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000", + "--numberOfWorkerHarnessThreads=1") + .setInstructions(instructions) + .setStreamingGlobalConfig( + StreamingGlobalConfig.builder() + .setOperationalLimits( + OperationalLimits.builder().setMaxWorkItemCommitBytes(1000).build()) + .build()) + .build()); worker.start(); - server.waitForEmptyWorkQueue(); + Windmill.Uint128Proto keyGroup = + Windmill.Uint128Proto.newBuilder().setHigh(0).setLow(1).build(); + String batchInputText = + " 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)); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertTrue(multiKeyCommits.isEmpty()); + + expectedWorkSchedulerLogs.verifyWarn("Windmill Multi-key commit batch size"); + + worker.stop(); + } + + @Test + public void testMultiKeyCommit_validationException_succeedsOnIndividualRetries() + 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") + .setInstructions(instructions) + .setStreamingGlobalConfig( + StreamingGlobalConfig.builder() + .setOperationalLimits( + OperationalLimits.builder().setMaxWorkItemCommitBytes(800).build()) + .build()) + .build()); + worker.start(); + + String batchInputText = + " 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)); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertTrue(multiKeyCommits.isEmpty()); + + worker.stop(); + } + + @Test + public void testMultiKeyCommit_elementFailure() throws Exception { + if (!streamingEngine) { + return; + } + KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); + + List instructions = + Arrays.asList( + makeSourceInstruction(kvCoder), + makeDoFnInstruction(new WorkDoFn(), 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) + .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: 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\"" + + " }" + + " }" + + " }" + + " work {" + + " key: \"key3\"" + + " sharding_key: 3" + + " 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 batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); server - .whenGetWorkCalled() - .thenReturn(makeInput(1, 0, DEFAULT_KEY_STRING, DEFAULT_SHARDING_KEY)); + .whenGetDataCalled() + .answerByDefault( + request -> { + Windmill.GetDataResponse.Builder builder = Windmill.GetDataResponse.newBuilder(); + for (ComputationGetDataRequest compRequest : request.getRequestsList()) { + ComputationGetDataResponse.Builder compBuilder = + builder.addDataBuilder().setComputationId(compRequest.getComputationId()); + for (KeyedGetDataRequest keyRequest : compRequest.getRequestsList()) { + KeyedGetDataResponse.Builder keyBuilder = + compBuilder + .addDataBuilder() + .setKey(keyRequest.getKey()) + .setShardingKey(keyRequest.getShardingKey()); + if (keyRequest.getWorkToken() == 2) { + keyBuilder.setFailed(true); + } else { + keyBuilder.addAllValues(keyRequest.getValuesToFetchList()); + keyBuilder.addAllBags(keyRequest.getBagsToFetchList()); + keyBuilder.addAllWatermarkHolds(keyRequest.getWatermarkHoldsToFetchList()); + } + } + } + return builder.build(); + }); + + server.whenGetWorkCalled().thenReturn(batchInput); + + 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(3, multiKeyCommit.getRequests(0).getWorkToken()); + assertEquals(1, multiKeyCommit.getRequests(1).getWorkToken()); + + worker.stop(); + } + + @Test + public void testCompleteCommit_retryableFailureTriggersReExecution() throws Exception { + if (!streamingEngine) { + return; + } + KvCoder kvCoder = KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()); + + List instructions = + Arrays.asList( + makeSourceInstruction(kvCoder), + makeDoFnInstruction(new WorkDoFn(), 0, kvCoder), + makeSinkInstruction(kvCoder, 1)); + + StreamingDataflowWorker worker = + makeWorker( + defaultWorkerParams( + "--experiments=unstable_enable_multi_key_bundle,max_key_group_batch_time_ms=5000", + "--numberOfWorkerHarnessThreads=1") + .setLocalRetryTimeoutMs(100) + .setInstructions(instructions) + .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: 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\"" + + " }" + + " }" + + " }" + + "}"; + Windmill.GetWorkResponse batchInput = + buildInput( + batchInputText, + CoderUtils.encodeToByteArray( + CollectionCoder.of(IntervalWindow.getCoder()), + Collections.singletonList(DEFAULT_WINDOW))); + + server + .whenGetDataCalled() + .answerByDefault( + request -> { + Windmill.GetDataResponse.Builder builder = Windmill.GetDataResponse.newBuilder(); + for (ComputationGetDataRequest compRequest : request.getRequestsList()) { + ComputationGetDataResponse.Builder compBuilder = + builder.addDataBuilder().setComputationId(compRequest.getComputationId()); + for (KeyedGetDataRequest keyRequest : compRequest.getRequestsList()) { + KeyedGetDataResponse.Builder keyBuilder = + compBuilder + .addDataBuilder() + .setKey(keyRequest.getKey()) + .setShardingKey(keyRequest.getShardingKey()); + if (keyRequest.getWorkToken() == 2) { + keyBuilder.setFailed(true); + } else { + keyBuilder.addAllValues(keyRequest.getValuesToFetchList()); + keyBuilder.addAllBags(keyRequest.getBagsToFetchList()); + keyBuilder.addAllWatermarkHolds(keyRequest.getWatermarkHoldsToFetchList()); + } + } + } + return builder.build(); + }); + + server.whenGetWorkCalled().thenReturn(batchInput); Map result = server.waitForAndGetCommits(1); - assertEquals( - makeExpectedOutput(1, 0, DEFAULT_KEY_STRING, DEFAULT_SHARDING_KEY, DEFAULT_KEY_STRING) - .build(), - removeDynamicFields(result.get(1L))); - assertEquals(1, result.size()); + assertTrue(result.containsKey(1L)); + assertFalse(result.containsKey(2L)); + + List multiKeyCommits = + server.getMultiKeyCommitsReceived(); + assertEquals(1, multiKeyCommits.size()); + Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = multiKeyCommits.get(0); + assertEquals(1, multiKeyCommit.getRequestsCount()); + assertEquals(1, multiKeyCommit.getRequests(0).getWorkToken()); worker.stop(); } @@ -3717,7 +4223,8 @@ public void testLatencyAttributionProtobufsPopulated() { ignored -> {}, mock(HeartbeatSender.class)), false, - clock); + clock, + ImmutableList.of()); clock.sleep(Duration.millis(10)); work.setState(Work.State.PROCESSING); @@ -4566,18 +5073,19 @@ public void evaluate() throws Throwable { } } - static class KeyTokenInvalidFn extends DoFn, KV> { - - static boolean thrown = false; + static class WorkDoFn extends DoFn, KV> { + @StateId("state") + private final StateSpec> stateSpec = StateSpecs.value(StringUtf8Coder.of()); @ProcessElement - public void processElement(ProcessContext c) { - if (!thrown) { - thrown = true; - throw new KeyTokenInvalidException("key"); - } else { - c.output(c.element()); + public void processElement(ProcessContext c, @StateId("state") ValueState state) { + try { + Thread.sleep(1000); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); } + state.read(); + c.output(c.element()); } } @@ -4597,6 +5105,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 3d11e5097463..54509fb5973e 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 @@ -24,6 +24,7 @@ import static org.hamcrest.Matchers.equalTo; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertThrows; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; @@ -58,6 +59,8 @@ 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.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.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; import org.apache.beam.runners.dataflow.worker.streaming.config.FakeGlobalConfigHandle; @@ -65,19 +68,20 @@ import org.apache.beam.runners.dataflow.worker.streaming.config.StreamingGlobalConfigHandle; 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.util.BoundedQueueExecutor; import org.apache.beam.runners.dataflow.worker.util.common.worker.WorkExecutor; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; 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.state.WindmillStateReader; 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; import org.apache.beam.sdk.coders.CoderException; import org.apache.beam.sdk.metrics.MetricsContainer; +import org.apache.beam.sdk.options.ExperimentalOptions; import org.apache.beam.sdk.options.PipelineOptionsFactory; import org.apache.beam.sdk.state.TimeDomain; import org.apache.beam.sdk.transforms.Create; @@ -86,6 +90,7 @@ import org.apache.beam.sdk.values.CausedByDrain; import org.apache.beam.sdk.values.PCollectionView; 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; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Lists; import org.hamcrest.Matchers; import org.joda.time.Duration; @@ -104,7 +109,7 @@ public class StreamingModeExecutionContextTest { @Rule public transient Timeout globalTimeout = Timeout.seconds(600); - @Mock private WindmillStateReader stateReader; + @Mock private WorkExecutor workExecutor; private static final String COMPUTATION_ID = "computationId"; @@ -116,7 +121,7 @@ public class StreamingModeExecutionContextTest { private FakeGlobalConfigHandle globalConfigHandle; private StreamingModeExecutionContext createExecutionContext( - StreamingGlobalConfigHandle configHandle) { + DataflowWorkerHarnessOptions options, StreamingGlobalConfigHandle configHandle) { CounterSet counterSet = new CounterSet(); ConcurrentHashMap stateNameMap = new ConcurrentHashMap<>(); stateNameMap.put(NameContextsForTests.nameContextForTest().userName(), "testStateFamily"); @@ -146,8 +151,9 @@ private StreamingModeExecutionContext createExecutionContext( /*stepName=*/ "stepName", /*systemName=*/ "systemName", StreamingCounters.create(), - mock(FailureTracker.class), + StreamingEngineFailureTracker.create(10, 10), "sourceBytesProcessCounterName", + MultiKeyBundleOptions.fromOptions(options), SideInputStateFetcherFactory.fromOptions(options)); } @@ -155,8 +161,11 @@ private StreamingModeExecutionContext createExecutionContext( public void setUp() { MockitoAnnotations.initMocks(this); options = PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + options + .as(ExperimentalOptions.class) + .setExperiments(Arrays.asList("unstable_enable_multi_key_bundle")); globalConfigHandle = new FakeGlobalConfigHandle(StreamingGlobalConfig.builder().build()); - executionContext = createExecutionContext(globalConfigHandle); + executionContext = createExecutionContext(options, globalConfigHandle); } private static Work createMockWork(Windmill.WorkItem workItem, Watermarks watermarks) { @@ -167,7 +176,8 @@ private static Work createMockWork(Windmill.WorkItem workItem, Watermarks waterm Work.createProcessingContext( COMPUTATION_ID, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); } private void start(Work work) { @@ -186,7 +196,6 @@ private void start(StreamingModeExecutionContext context, Work work, Coder ke try { context.start( work, - stateReader, workExecutor, /* workQueueExecutor= */ null, /* budgetHandle= */ null, @@ -454,7 +463,7 @@ public void testStateTagEncodingBasedOnConfig() { FakeGlobalConfigHandle configHandle = new FakeGlobalConfigHandle( StreamingGlobalConfig.builder().setEnableStateTagEncodingV2(isV2Encoding).build()); - StreamingModeExecutionContext context = createExecutionContext(configHandle); + StreamingModeExecutionContext context = createExecutionContext(options, configHandle); assertEquals(expectedEncoding, context.getWindmillTagEncoding().getClass()); } } @@ -508,6 +517,300 @@ public void testStart_internalKeyDecoding() throws Exception { assertEquals("decodedKey", executionContext.getKey()); } + @Test + public void testAdvance_success() 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()); + + Windmill.WorkItem workItem2 = + Windmill.WorkItem.newBuilder() + .setKey(ByteString.copyFromUtf8("key2")) + .setWorkToken(2L) + .setKeyGroup(keyGroup) + .build(); + Work work2 = + createMockWork( + workItem2, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); + ExecutableWork executableWork2 = ExecutableWork.create(work2, (w, h) -> {}); + + org.mockito.Mockito.when( + mockExecutor.pollWork( + org.mockito.Mockito.eq(COMPUTATION_ID), + org.mockito.Mockito.eq(work1.getKeyGroup()), + org.mockito.Mockito.eq(mockHandle))) + .thenReturn(executableWork2); + + executionContext.start( + work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertTrue(executionContext.advance()); + assertEquals("key2", executionContext.getSerializedKey().toStringUtf8()); + } + + @Test + public void testAdvance_noMoreWork() 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()); + + org.mockito.Mockito.when( + mockExecutor.pollWork( + org.mockito.Mockito.eq(COMPUTATION_ID), + org.mockito.Mockito.eq(work1.getKeyGroup()), + org.mockito.Mockito.eq(mockHandle))) + .thenReturn(null); + + executionContext.start( + work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(executionContext.advance()); + } + + @Test + public void testAdvance_respectsMaxBatchSize() throws Exception { + DataflowWorkerHarnessOptions optionsWithBatchSize = + PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + optionsWithBatchSize + .as(ExperimentalOptions.class) + .setExperiments(Arrays.asList("windmill_max_key_group_batch_size=1")); + StreamingModeExecutionContext context = + createExecutionContext(optionsWithBatchSize, globalConfigHandle); + + 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()); + + context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(context.advance()); + org.mockito.Mockito.verifyNoInteractions(mockExecutor); + } + + @Test + public void testAdvance_respectsMaxBatchTime() throws Exception { + DataflowWorkerHarnessOptions optionsWithBatchTime = + PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + optionsWithBatchTime + .as(ExperimentalOptions.class) + .setExperiments(Arrays.asList("windmill_max_key_group_batch_time_ms=0")); + StreamingModeExecutionContext context = + createExecutionContext(optionsWithBatchTime, globalConfigHandle); + + 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()); + + context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(context.advance()); + org.mockito.Mockito.verifyNoInteractions(mockExecutor); + } + + @Test + public void testAdvance_workFailed() 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()); + + executionContext.start( + work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + work1.setFailed(); + + assertThrows(WorkItemCancelledException.class, () -> executionContext.advance()); + } + + @Test + public void testAdvance_defaultKeyGroup() throws Exception { + BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class); + BoundedQueueExecutorWorkHandle mockHandle = mock(BoundedQueueExecutorWorkHandle.class); + + Windmill.WorkItem workItem1 = + Windmill.WorkItem.newBuilder() + .setKey(ByteString.copyFromUtf8("key1")) + .setWorkToken(1L) + .build(); + Work work1 = + createMockWork( + workItem1, Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build()); + + executionContext.start( + work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(executionContext.advance()); + org.mockito.Mockito.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()); + org.mockito.Mockito.verifyNoInteractions(mockExecutor); + } + + @Test + public void testAdvance_experimentDisabled() throws Exception { + DataflowWorkerHarnessOptions optionsDisabled = + PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + StreamingModeExecutionContext context = + createExecutionContext(optionsDisabled, globalConfigHandle); + + 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()); + + context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + assertFalse(context.advance()); + org.mockito.Mockito.verifyNoInteractions(mockExecutor); + } + + @Test + public void testAdvance_respectsMaxBatchSinkBytes() throws Exception { + DataflowWorkerHarnessOptions optionsWithSinkBytes = + PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + optionsWithSinkBytes + .as(ExperimentalOptions.class) + .setExperiments( + Arrays.asList( + "unstable_enable_multi_key_bundle", "windmill_max_key_group_batch_sink_bytes=100")); + StreamingModeExecutionContext context = + createExecutionContext(optionsWithSinkBytes, globalConfigHandle); + + 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()); + + context.start(work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, newWork) -> {}); + + context.reportBytesSinked(50); + assertFalse(context.advance()); + org.mockito.Mockito.verify(mockExecutor) + .pollWork(COMPUTATION_ID, work1.getKeyGroup(), mockHandle); + + org.mockito.Mockito.reset(mockExecutor); + + context.reportBytesSinked(60); + assertFalse(context.advance()); + org.mockito.Mockito.verifyNoInteractions(mockExecutor); + } + + @Test + public void testExperimentParsingWithInvalidValues() { + DataflowWorkerHarnessOptions optionsInvalid = + PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class); + optionsInvalid + .as(ExperimentalOptions.class) + .setExperiments( + Arrays.asList( + "windmill_max_key_group_batch_size=invalid_size", + "windmill_max_key_group_batch_time_ms=invalid_time", + "windmill_max_key_group_batch_sink_bytes=invalid_bytes")); + + // This should not throw NumberFormatException + StreamingModeExecutionContext context = + createExecutionContext(optionsInvalid, globalConfigHandle); + + org.junit.Assert.assertNotNull(context); + } + @Test public void testInternalsPoisonedAfterFlushState() throws Exception { NameContext nameContext = NameContextsForTests.nameContextForTest(); @@ -559,4 +862,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/WindmillReaderIteratorBaseTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WindmillReaderIteratorBaseTest.java index b45e0de6447c..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 @@ -37,6 +37,7 @@ import org.apache.beam.sdk.values.WindowedValue; import org.apache.beam.sdk.values.WindowedValues; 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; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.JUnit4; @@ -269,6 +270,7 @@ private static Work createMockWork(Windmill.WorkItem workItem) { org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender .class)), false, - org.joda.time.Instant::now); + org.joda.time.Instant::now, + ImmutableList.of()); } } 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 2e7c80330cf0..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 @@ -93,7 +93,8 @@ private static Work createMockWork(Windmill.WorkItem workItem) { Work.createProcessingContext( "computationId", new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); } private static ByteString encodeMetadata(List windows) throws IOException { 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 2da52bea8eba..679227a11dc0 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 @@ -103,7 +103,6 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill; 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.state.WindmillStateReader; import org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.FailureTracker; import org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender; import org.apache.beam.sdk.Pipeline; @@ -210,14 +209,14 @@ private static Work createMockWork(Windmill.WorkItem workItem, Watermarks waterm Work.createProcessingContext( COMPUTATION_ID, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); } private void startContext(StreamingModeExecutionContext context, Work work) { try { context.start( work, - mock(WindmillStateReader.class), mock(WorkExecutor.class), /* workQueueExecutor= */ null, /* budgetHandle= */ null, @@ -647,6 +646,7 @@ public void testReadUnboundedReader() throws Exception { StreamingCounters.create(), mock(FailureTracker.class), "sourceBytesProcessCounterName", + MultiKeyBundleOptions.fromOptions(options), SideInputStateFetcherFactory.fromOptions(options)); options.setNumWorkers(5); @@ -1023,6 +1023,7 @@ public void testFailedWorkItemsAbort() throws Exception { StreamingCounters.create(), mock(FailureTracker.class), "sourceBytesProcessCounterName", + MultiKeyBundleOptions.fromOptions(options), SideInputStateFetcherFactory.fromOptions(options)); options.setNumWorkers(5); @@ -1050,7 +1051,8 @@ public void testFailedWorkItemsAbort() throws Exception { ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); startContext(context, dummyWork); @SuppressWarnings({"unchecked", "rawtypes"}) 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 0f14efdd0c0b..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 @@ -20,6 +20,8 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertNull; import static org.junit.Assert.assertSame; import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; @@ -72,7 +74,8 @@ private static ExecutableWork createWork(Windmill.WorkItem workItem) { Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), createWorkProcessingContext(), false, - Instant::now), + Instant::now, + ImmutableList.of()), (work, handle) -> {}); } @@ -84,7 +87,8 @@ private static ExecutableWork expiredWork(Windmill.WorkItem workItem) { Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), createWorkProcessingContext(), false, - () -> Instant.EPOCH), + () -> Instant.EPOCH, + ImmutableList.of()), (work, handle) -> {}); } @@ -565,6 +569,31 @@ public void testFailWork_batchFail() { } } + @Test + public void testGetActiveWork() { + ShardedKey shardedKey = shardedKey("someKey", 1L); + ExecutableWork work = createWork(createWorkItem(1L, 1L, shardedKey)); + + // Initially empty + assertNull(activeWorkState.getActiveWork(shardedKey, work.id())); + + // Activate work + activeWorkState.activateWorkForKey(work); + + // Should find it now + ExecutableWork activeWork = activeWorkState.getActiveWork(shardedKey, work.id()); + assertNotNull(activeWork); + assertSame(work, activeWork); + + // Should not find it with different workId + assertNull(activeWorkState.getActiveWork(shardedKey, workId(2L, 1L))); + assertNull(activeWorkState.getActiveWork(shardedKey, workId(1L, 2L))); + + // Should not find it with different shardedKey + ShardedKey otherShardedKey = shardedKey("otherKey", 2L); + assertNull(activeWorkState.getActiveWork(otherShardedKey, work.id())); + } + private static ExecutableWork firstValue(Map map) { Iterator> iterator = map.entrySet().iterator(); if (iterator.hasNext()) { 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 30ad97140e1e..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 @@ -41,6 +41,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender; import org.apache.beam.sdk.fn.IdGenerators; 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; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableSet; import org.joda.time.Instant; @@ -76,7 +77,8 @@ private static ExecutableWork createWork(ShardedKey shardedKey, long workToken, ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now), + Instant::now, + ImmutableList.of()), (work, handle) -> {}); } 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 new file mode 100644 index 000000000000..22ddc8e4de5b --- /dev/null +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/streaming/ComputationStateTest.java @@ -0,0 +1,114 @@ +/* + * 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.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.verifyNoMoreInteractions; + +import com.google.api.services.dataflow.model.MapTask; +import java.util.Collections; +import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor; +import org.apache.beam.runners.dataflow.worker.windmill.Windmill; +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; +import org.joda.time.Instant; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public class ComputationStateTest { + + private final BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class); + private final WindmillStateCache.ForComputation mockStateCache = + mock(WindmillStateCache.ForComputation.class); + private final HeartbeatSender mockHeartbeatSender = mock(HeartbeatSender.class); + + private ComputationState computationState; + + private static ShardedKey shardedKey(String str, long shardKey) { + return ShardedKey.create(ByteString.copyFromUtf8(str), shardKey); + } + + private ExecutableWork createWork(Windmill.WorkItem workItem) { + return ExecutableWork.create( + Work.create( + workItem, + workItem.getSerializedSize(), + Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build(), + Work.createProcessingContext( + "computationId", new FakeGetDataClient(), ignored -> {}, mockHeartbeatSender), + false, + Instant::now, + ImmutableList.of()), + (work, handle) -> {}); + } + + private static Windmill.WorkItem createWorkItem( + long workToken, long cacheToken, ShardedKey shardedKey) { + return Windmill.WorkItem.newBuilder() + .setShardingKey(shardedKey.shardingKey()) + .setKey(shardedKey.key()) + .setWorkToken(workToken) + .setCacheToken(cacheToken) + .build(); + } + + @Before + public void setUp() { + MapTask mapTask = new MapTask(); + mapTask.setStageName("stage"); + mapTask.setSystemName("system"); + computationState = + new ComputationState( + "computationId", mapTask, mockExecutor, Collections.emptyMap(), mockStateCache); + } + + @Test + public void testReexecuteActiveWork_workNotActive() { + ShardedKey shardedKey = shardedKey("key", 1L); + WorkId workId = WorkId.builder().setWorkToken(1L).setCacheToken(1L).build(); + + computationState.reexecuteActiveWork(shardedKey, workId); + + verifyNoInteractions(mockExecutor); + } + + @Test + public void testReexecuteActiveWork_workActive() { + ShardedKey shardedKey = shardedKey("key", 1L); + Windmill.WorkItem workItem = createWorkItem(1L, 1L, shardedKey); + ExecutableWork work = createWork(workItem); + + // Activate work first. This will execute it once. + computationState.activateWork(work); + verify(mockExecutor).execute(work, work.work().getSerializedWorkItemSize()); + + // Now re-execute + computationState.reexecuteActiveWork(shardedKey, work.id()); + verify(mockExecutor).forceExecute(work, work.work().getSerializedWorkItemSize()); + + verifyNoMoreInteractions(mockExecutor); + } +} 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 80ca91da462f..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 @@ -30,6 +30,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.Windmill; 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; import org.joda.time.Instant; import org.junit.Test; import org.junit.runner.RunWith; @@ -57,7 +58,8 @@ private static Work createTestWork() { commit -> {}, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); } @Test 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..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 @@ -30,7 +30,10 @@ import java.util.Collection; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +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.ExecutableWork; import org.apache.beam.runners.dataflow.worker.streaming.Watermarks; import org.apache.beam.runners.dataflow.worker.streaming.Work; @@ -40,6 +43,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; 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; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.ThreadFactoryBuilder; import org.joda.time.Instant; import org.junit.Before; @@ -82,6 +86,14 @@ private static ExecutableWork createWorkWithCompId( private static ExecutableWork createWorkWithCompIdAndKeyGroup( String computationId, Work.KeyGroup keyGroup, Consumer executeWorkFn) { + return createWorkWithHandle( + computationId, keyGroup, (work, handle) -> executeWorkFn.accept(work)); + } + + private static ExecutableWork createWorkWithHandle( + String computationId, + Work.KeyGroup keyGroup, + BiConsumer executeWorkFn) { WorkItem workItem = WorkItem.newBuilder() .setKey(ByteString.EMPTY) @@ -102,10 +114,9 @@ private static ExecutableWork createWorkWithCompIdAndKeyGroup( Work.createProcessingContext( computationId, new FakeGetDataClient(), ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now), - (work, handle) -> { - executeWorkFn.accept(work); - }); + Instant::now, + ImmutableList.of()), + executeWorkFn); } private ExecutableWork createSleepProcessWork(CountDownLatch start, CountDownLatch stop) { @@ -406,18 +417,25 @@ public void testRunnableExceptionPropagationDecrementsCounters() throws Exceptio @Test public void testHandleMerge() throws Exception { - BoundedQueueExecutorWorkHandleImpl handle1 = executor.createBudgetHandle(1, 100L); - BoundedQueueExecutorWorkHandleImpl handle2 = executor.createBudgetHandle(2, 200L); + Work work1 = createWork(ignored -> {}).work(); + Work work2 = createWork(ignored -> {}).work(); + Work work3 = createWork(ignored -> {}).work(); + BoundedQueueExecutorWorkHandleImpl handle1 = executor.createBudgetHandle(work1, 100L); + BoundedQueueExecutorWorkHandleImpl handle2 = executor.createBudgetHandle(work2, 200L); + handle2.merge(executor.createBudgetHandle(work3, 0L)); handle1.merge(handle2); // Verify that handle2 has 0 budget and is closed. - assertEquals(0, handle2.elements()); + assertEquals(0, handle2.getWorkBatch().size()); assertEquals(0, handle2.bytes()); assertTrue(handle2.isClosed()); // Verify that handle1 has the combined budget and is not closed. - assertEquals(3, handle1.elements()); + assertEquals(3, handle1.getWorkBatch().size()); + assertTrue(handle1.getWorkBatch().contains(work1)); + assertTrue(handle1.getWorkBatch().contains(work2)); + assertTrue(handle1.getWorkBatch().contains(work3)); assertEquals(300L, handle1.bytes()); assertFalse(handle1.isClosed()); } @@ -449,11 +467,13 @@ public void testPollWork() throws Exception { // 1. Create blocker task to occupy the worker thread CountDownLatch blockerStart = new CountDownLatch(1); CountDownLatch blockerStop = new CountDownLatch(1); + AtomicReference blockerHandleRef = new AtomicReference<>(); ExecutableWork blockerWork = - createWorkWithCompIdAndKeyGroup( + createWorkWithHandle( "blockerComp", DEFAULT_KEY_GROUP, - ignored -> { + (work, handle) -> { + blockerHandleRef.set(handle); blockerStart.countDown(); try { blockerStop.await(); @@ -464,6 +484,9 @@ public void testPollWork() throws Exception { testExecutor.execute(blockerWork, 0); blockerStart.await(); + BoundedQueueExecutorWorkHandleImpl stealHandle = + (BoundedQueueExecutorWorkHandleImpl) blockerHandleRef.get(); + assertNotNull(stealHandle); // 2. Create two distinct key groups Work.KeyGroup keyGroup1 = Work.KeyGroup.create(1, 1); @@ -488,22 +511,18 @@ public void testPollWork() throws Exception { assertEquals(3, testExecutor.elementsOutstanding()); // Steal work2 using pollWork with compA and keyGroup2 - try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup2, stealHandle); - assertNotNull(stolen); - assertEquals(work2, stolen); - - // Run the stolen task - stolen.run(stealHandle); - targetStart.await(); - } + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup2, stealHandle); + assertNotNull(stolen); + assertEquals(work2, stolen); + + // Run the stolen task + stolen.run(stealHandle); + targetStart.await(); // Steal work1 using pollWork with compA and keyGroup1 - try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup1, stealHandle); - assertNotNull(stolen); - assertEquals(work1, stolen); - } + ExecutableWork stolen1 = testExecutor.pollWork("compA", keyGroup1, stealHandle); + assertNotNull(stolen1); + assertEquals(work1, stolen1); // Unblock the blocker and shut down blockerStop.countDown(); @@ -525,11 +544,13 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { CountDownLatch blockerStart = new CountDownLatch(1); CountDownLatch blockerStop = new CountDownLatch(1); + AtomicReference blockerHandleRef = new AtomicReference<>(); ExecutableWork blockerWork = - createWorkWithCompIdAndKeyGroup( + createWorkWithHandle( "blockerComp", DEFAULT_KEY_GROUP, - ignored -> { + (work, handle) -> { + blockerHandleRef.set(handle); blockerStart.countDown(); try { blockerStop.await(); @@ -540,15 +561,16 @@ public void testPollWorkWithLinkedBlockingQueue() throws Exception { testExecutor.execute(blockerWork, 0); blockerStart.await(); + BoundedQueueExecutorWorkHandleImpl stealHandle = + (BoundedQueueExecutorWorkHandleImpl) blockerHandleRef.get(); + assertNotNull(stealHandle); Work.KeyGroup keyGroup = Work.KeyGroup.create(1, 1); ExecutableWork work = createWorkWithCompIdAndKeyGroup("compA", keyGroup, ignored -> {}); testExecutor.execute(work, 100); - try (BoundedQueueExecutorWorkHandleImpl stealHandle = testExecutor.createBudgetHandle(0, 0L)) { - ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup, stealHandle); - assertNull(stolen); - } + ExecutableWork stolen = testExecutor.pollWork("compA", keyGroup, stealHandle); + assertNull(stolen); 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 994aa2030f3f..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 @@ -45,6 +45,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; 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; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.ThreadFactoryBuilder; import org.checkerframework.checker.nullness.qual.Nullable; import org.joda.time.Instant; @@ -63,7 +64,6 @@ public static Iterable data() { } @Parameterized.Parameter public boolean fairQueue; - private BoundedQueueExecutor executor; @Before @@ -114,9 +114,10 @@ private QueuedWork createQueuedWork( ignored -> {}, mock(HeartbeatSender.class)), false, - Instant::now), + Instant::now, + ImmutableList.of()), (w, h) -> {}); - return new QueuedWork(work, executor.createBudgetHandle(1, workBytes)); + return new QueuedWork(work, executor.createBudgetHandle(work.work(), workBytes)); } private static class NoOpRunnable implements Runnable { @@ -312,7 +313,6 @@ public String toString() { } })); } - // Start producers for (int i = 0; i < producerThreads; i++) { futures.add( @@ -470,7 +470,6 @@ public void testPollWorkWithKeyGroup() { QueuedWork polledNotExist = queue.pollWork("compA", keyGroupNotExist); assertNull(polledNotExist); assertEquals(2, queue.size()); - // Poll with keyGroup2 first - should return workA2 QueuedWork polledA2 = queue.pollWork("compA", keyGroup2); assertNotNull(polledA2); @@ -485,12 +484,32 @@ public void testPollWorkWithKeyGroup() { assertNotNull(polledA1); assertEquals(workA1, polledA1); assertTrue(queue.isEmpty()); - polledNotExist = queue.pollWork("compA", keyGroupNotExist); assertNull(polledNotExist); 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/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 0596210a0270..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 @@ -39,6 +39,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; 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; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap; import org.joda.time.Instant; import org.junit.After; @@ -75,7 +76,8 @@ private static Work createMockWork(long workToken) { }, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); } private static ComputationState createComputationState(String computationId) { 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 881cf620e8d8..3e1e4ae54f06 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 @@ -125,7 +125,8 @@ private static Work createMockWork(long workToken) { }, mock(HeartbeatSender.class)), false, - Instant::now); + Instant::now, + ImmutableList.of()); } private static ComputationState createComputationState(String computationId) { @@ -140,7 +141,8 @@ private static ComputationState createComputationState(String computationId) { private static CompleteCommit asCompleteCommit( String computationId, Work work, Windmill.CommitStatus status) { Windmill.CommitStatus finalStatus = work.isFailed() ? Windmill.CommitStatus.ABORTED : status; - return CompleteCommit.create(computationId, work.getShardedKey(), work.id(), finalStatus); + return CompleteCommit.create( + computationId, work.getShardedKey(), work.id(), finalStatus, /* retryableFailure= */ false); } @Before @@ -397,6 +399,7 @@ public void shutdown() {} assertThat(commits.size()).isEqualTo(completeCommits.size()); for (CompleteCommit completeCommit : completeCommits) { assertThat(completeCommit.status()).isEqualTo(Windmill.CommitStatus.ABORTED); + assertThat(completeCommit.retryableFailure()).isFalse(); } for (Commit commit : commits) { @@ -568,11 +571,23 @@ public void testCommit_multiKeyCommitSuccess() { assertThat(completeCommits) .containsExactly( CompleteCommit.create( - "computationId", workA.getShardedKey(), workA.id(), CommitStatus.OK), + "computationId", + workA.getShardedKey(), + workA.id(), + CommitStatus.OK, + /* retryableFailure= */ false), CompleteCommit.create( - "computationId", workB.getShardedKey(), workB.id(), CommitStatus.OK), + "computationId", + workB.getShardedKey(), + workB.id(), + CommitStatus.OK, + /* retryableFailure= */ false), CompleteCommit.create( - "computationId", workC.getShardedKey(), workC.id(), CommitStatus.OK)); + "computationId", + workC.getShardedKey(), + workC.id(), + CommitStatus.OK, + /* retryableFailure= */ false)); // There should be no more commits in the queue assertEquals(0, workCommitter.currentActiveCommitBytes()); @@ -632,11 +647,23 @@ public void testCommit_multiKeyCommitFailedWork() { assertThat(completeCommits) .containsExactly( CompleteCommit.create( - "computationId", workA.getShardedKey(), workA.id(), CommitStatus.ABORTED), + "computationId", + workA.getShardedKey(), + workA.id(), + CommitStatus.ABORTED, + /* retryableFailure= */ true), CompleteCommit.create( - "computationId", workB.getShardedKey(), workB.id(), CommitStatus.ABORTED), + "computationId", + workB.getShardedKey(), + workB.id(), + CommitStatus.ABORTED, + /* retryableFailure= */ false), CompleteCommit.create( - "computationId", workC.getShardedKey(), workC.id(), CommitStatus.ABORTED)); + "computationId", + workC.getShardedKey(), + workC.id(), + CommitStatus.ABORTED, + /* retryableFailure= */ true)); // There should be no more commits in the queue assertEquals(0, workCommitter.currentActiveCommitBytes()); @@ -703,11 +730,23 @@ public void testCommit_multiKeyCommitStatusNotOK() { assertThat(completeCommits) .containsExactly( CompleteCommit.create( - "computationId", workA.getShardedKey(), workA.id(), CommitStatus.NOT_FOUND), + "computationId", + workA.getShardedKey(), + workA.id(), + CommitStatus.NOT_FOUND, + /* retryableFailure= */ false), CompleteCommit.create( - "computationId", workB.getShardedKey(), workB.id(), CommitStatus.NOT_FOUND), + "computationId", + workB.getShardedKey(), + workB.id(), + CommitStatus.NOT_FOUND, + /* retryableFailure= */ false), CompleteCommit.create( - "computationId", workC.getShardedKey(), workC.id(), CommitStatus.NOT_FOUND)); + "computationId", + workC.getShardedKey(), + workC.id(), + CommitStatus.NOT_FOUND, + /* retryableFailure= */ false)); // There should be no more commits in the queue assertEquals(0, workCommitter.currentActiveCommitBytes()); diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java index 1e995f4047c3..147f053fd8ae 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java @@ -1219,6 +1219,194 @@ private void testMultiKeyCommit(CommitStatus commitStatus) throws Exception { assertThat(commitStatusFuture.get()).isEqualTo(commitStatus); } + @Test + public void testCommit_multiKeyCommit_multichunk() throws Exception { + GrpcCommitWorkStream commitWorkStream = createCommitWorkStream(); + FakeWindmillGrpcService.CommitStreamInfo streamInfo = waitForConnectionAndConsumeHeader(); + + CompletableFuture commitStatusFuture = new CompletableFuture<>(); + + long shardingKey1 = 101L; + long workToken1 = 201L; + long cacheToken1 = 301L; + long shardingKey2 = 102L; + long workToken2 = 202L; + long cacheToken2 = 302L; + + Windmill.WorkItemCommitRequest request1 = + Windmill.WorkItemCommitRequest.newBuilder() + .setKey(ByteString.copyFromUtf8("key1")) + .setShardingKey(shardingKey1) + .setWorkToken(workToken1) + .setCacheToken(cacheToken1) + .addBagUpdates(Windmill.TagBag.newBuilder().setTag(LARGE_BYTE_STRING).build()) + .build(); + + Windmill.WorkItemCommitRequest request2 = + Windmill.WorkItemCommitRequest.newBuilder() + .setKey(ByteString.copyFromUtf8("key2")) + .setShardingKey(shardingKey2) + .setWorkToken(workToken2) + .setCacheToken(cacheToken2) + .build(); + + Windmill.MultiKeyWorkItemCommitRequest multiKeyRequest = + Windmill.MultiKeyWorkItemCommitRequest.newBuilder() + .addRequests(request1) + .addRequests(request2) + .build(); + + try (WindmillStream.CommitWorkStream.RequestBatcher batcher = commitWorkStream.batcher()) { + assertTrue( + batcher.commitMultiKeyWorkItem( + COMPUTATION_ID, multiKeyRequest, commitStatusFuture::complete)); + } + + Windmill.StreamingCommitWorkRequest requestChunk1 = streamInfo.requests.take(); + assertThat(requestChunk1.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk chunk1 = requestChunk1.getCommitChunk(0); + + assertThat(chunk1.getCommitType()) + .isEqualTo(Windmill.StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_MULTI_KEY); + assertThat(chunk1.getShardingKey()).isEqualTo(request1.getShardingKey()); + assertThat(chunk1.getRemainingBytesForWorkItem()).isGreaterThan(0); + + Windmill.StreamingCommitWorkRequest requestChunk2 = streamInfo.requests.take(); + assertThat(requestChunk2.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk chunk2 = requestChunk2.getCommitChunk(0); + + assertThat(chunk2.getCommitType()) + .isEqualTo(Windmill.StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_MULTI_KEY); + assertThat(chunk2.getShardingKey()).isEqualTo(request1.getShardingKey()); + assertThat(chunk2.getRemainingBytesForWorkItem()).isEqualTo(0); + + ByteString reconstructedBytes = + chunk1.getSerializedWorkItemCommit().concat(chunk2.getSerializedWorkItemCommit()); + Windmill.MultiKeyWorkItemCommitRequest parsedRequest = + Windmill.MultiKeyWorkItemCommitRequest.parseFrom(reconstructedBytes); + assertThat(parsedRequest).isEqualTo(multiKeyRequest); + + long requestId = chunk1.getRequestId(); + assertThat(chunk2.getRequestId()).isEqualTo(requestId); + + streamInfo.responseObserver.onNext( + Windmill.StreamingCommitResponse.newBuilder().addRequestId(requestId).build()); + + assertThat(commitStatusFuture.get()).isEqualTo(Windmill.CommitStatus.OK); + } + + @Test + public void testCommitMultiKeyWorkItem_retryOnNewStream() throws Exception { + GrpcCommitWorkStream commitWorkStream = createCommitWorkStream(); + FakeWindmillGrpcService.CommitStreamInfo streamInfo = waitForConnectionAndConsumeHeader(); + + CompletableFuture commitStatusFuture = new CompletableFuture<>(); + + long shardingKey1 = 101L; + long workToken1 = 201L; + long cacheToken1 = 301L; + Windmill.WorkItemCommitRequest request1 = + Windmill.WorkItemCommitRequest.newBuilder() + .setKey(ByteString.copyFromUtf8("key1")) + .setShardingKey(shardingKey1) + .setWorkToken(workToken1) + .setCacheToken(cacheToken1) + .build(); + Windmill.MultiKeyWorkItemCommitRequest multiKeyRequest = + Windmill.MultiKeyWorkItemCommitRequest.newBuilder().addRequests(request1).build(); + + try (WindmillStream.CommitWorkStream.RequestBatcher batcher = commitWorkStream.batcher()) { + assertTrue( + batcher.commitMultiKeyWorkItem( + COMPUTATION_ID, multiKeyRequest, commitStatusFuture::complete)); + } + + Windmill.StreamingCommitWorkRequest request = streamInfo.requests.take(); + assertThat(request.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk chunk = request.getCommitChunk(0); + assertThat(chunk.getCommitType()) + .isEqualTo(Windmill.StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_MULTI_KEY); + long requestId = chunk.getRequestId(); + + streamInfo.responseObserver.onError(new IOException("test error")); + + FakeWindmillGrpcService.CommitStreamInfo reconnectStreamInfo = + waitForConnectionAndConsumeHeader(); + Windmill.StreamingCommitWorkRequest reconnectRequest = reconnectStreamInfo.requests.take(); + assertThat(reconnectRequest.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk reconnectChunk = reconnectRequest.getCommitChunk(0); + assertThat(reconnectChunk.getCommitType()) + .isEqualTo(Windmill.StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_MULTI_KEY); + assertThat(reconnectChunk.getRequestId()).isEqualTo(requestId); + + Windmill.MultiKeyWorkItemCommitRequest parsedRequest = + Windmill.MultiKeyWorkItemCommitRequest.parseFrom( + reconnectChunk.getSerializedWorkItemCommit()); + assertThat(parsedRequest).isEqualTo(multiKeyRequest); + + reconnectStreamInfo.responseObserver.onNext( + Windmill.StreamingCommitResponse.newBuilder().addRequestId(requestId).build()); + assertThat(commitStatusFuture.get()).isEqualTo(Windmill.CommitStatus.OK); + } + + @Test + public void testCommitWorkItem_retryOnNewStream_multichunk() throws Exception { + GrpcCommitWorkStream commitWorkStream = createCommitWorkStream(); + FakeWindmillGrpcService.CommitStreamInfo streamInfo = waitForConnectionAndConsumeHeader(); + + CompletableFuture commitStatusFuture = new CompletableFuture<>(); + + Windmill.WorkItemCommitRequest largeRequest = + workItemCommitRequest(1) + .toBuilder() + .addBagUpdates(Windmill.TagBag.newBuilder().setTag(LARGE_BYTE_STRING).build()) + .build(); + + try (WindmillStream.CommitWorkStream.RequestBatcher batcher = commitWorkStream.batcher()) { + assertTrue( + batcher.commitWorkItem(COMPUTATION_ID, largeRequest, commitStatusFuture::complete)); + } + + Windmill.StreamingCommitWorkRequest requestChunk1 = streamInfo.requests.take(); + assertThat(requestChunk1.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk chunk1 = requestChunk1.getCommitChunk(0); + long requestId = chunk1.getRequestId(); + assertThat(chunk1.getRemainingBytesForWorkItem()).isGreaterThan(0); + + Windmill.StreamingCommitWorkRequest requestChunk2 = streamInfo.requests.take(); + assertThat(requestChunk2.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk chunk2 = requestChunk2.getCommitChunk(0); + assertThat(chunk2.getRequestId()).isEqualTo(requestId); + assertThat(chunk2.getRemainingBytesForWorkItem()).isEqualTo(0); + + streamInfo.responseObserver.onError(new IOException("test error")); + + FakeWindmillGrpcService.CommitStreamInfo reconnectStreamInfo = + waitForConnectionAndConsumeHeader(); + + Windmill.StreamingCommitWorkRequest reconnectChunk1 = reconnectStreamInfo.requests.take(); + assertThat(reconnectChunk1.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk reconChunk1 = reconnectChunk1.getCommitChunk(0); + assertThat(reconChunk1.getRequestId()).isEqualTo(requestId); + assertThat(reconChunk1.getRemainingBytesForWorkItem()).isGreaterThan(0); + + Windmill.StreamingCommitWorkRequest reconnectChunk2 = reconnectStreamInfo.requests.take(); + assertThat(reconnectChunk2.getCommitChunkCount()).isEqualTo(1); + Windmill.StreamingCommitRequestChunk reconChunk2 = reconnectChunk2.getCommitChunk(0); + assertThat(reconChunk2.getRequestId()).isEqualTo(requestId); + assertThat(reconChunk2.getRemainingBytesForWorkItem()).isEqualTo(0); + + ByteString reconstructedBytes = + reconChunk1.getSerializedWorkItemCommit().concat(reconChunk2.getSerializedWorkItemCommit()); + Windmill.WorkItemCommitRequest parsedRequest = + Windmill.WorkItemCommitRequest.parseFrom(reconstructedBytes); + assertThat(parsedRequest).isEqualTo(largeRequest); + + reconnectStreamInfo.responseObserver.onNext( + Windmill.StreamingCommitResponse.newBuilder().addRequestId(requestId).build()); + assertThat(commitStatusFuture.get()).isEqualTo(Windmill.CommitStatus.OK); + } + @Test public void testCommitWorkItem_stopsRetriesAfterDuration() throws Exception { int numCommits = 1; diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReaderTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReaderTest.java index 1611fdac25dc..e38388ff566b 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReaderTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/state/WindmillStateReaderTest.java @@ -35,9 +35,9 @@ import java.util.Map; import java.util.Optional; import java.util.concurrent.Future; -import org.apache.beam.runners.dataflow.worker.KeyTokenInvalidException; import org.apache.beam.runners.dataflow.worker.WindmillStateTestUtils; import org.apache.beam.runners.dataflow.worker.WindmillTimeUtils; +import org.apache.beam.runners.dataflow.worker.WorkCancellingException; import org.apache.beam.runners.dataflow.worker.windmill.Windmill; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.KeyedGetDataRequest; import org.apache.beam.runners.dataflow.worker.windmill.Windmill.SortedListEntry; @@ -1572,16 +1572,16 @@ public void testKeyTokenInvalid() throws Exception { try { watermarkFuture.get(); - fail("Expected KeyTokenInvalidException"); + fail("Expected WorkCancelingException"); } catch (Exception e) { - assertTrue(KeyTokenInvalidException.isKeyTokenInvalidException(e)); + assertTrue(WorkCancellingException.isWorkCancellingException(e)); } try { bagFuture.get(); - fail("Expected KeyTokenInvalidException"); + fail("Expected WorkCancelingException"); } catch (Exception e) { - assertTrue(KeyTokenInvalidException.isKeyTokenInvalidException(e)); + assertTrue(WorkCancellingException.isWorkCancellingException(e)); } } 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 0610ed44c27f..02f9c4d5a3d4 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 @@ -21,15 +21,15 @@ import static org.junit.Assert.assertThrows; import static org.mockito.Mockito.mock; +import java.util.Arrays; 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.function.Consumer; import java.util.function.Supplier; -import org.apache.beam.runners.dataflow.worker.KeyTokenInvalidException; -import org.apache.beam.runners.dataflow.worker.WorkItemCancelledException; 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; @@ -39,6 +39,7 @@ import org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient; 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; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.ThreadFactoryBuilder; import org.joda.time.Duration; import org.joda.time.Instant; @@ -98,7 +99,8 @@ private static ExecutableWork createWork(Supplier clock, Consumer ignored -> {}, mock(HeartbeatSender.class)), false, - clock), + clock, + ImmutableList.of()), (work, handle) -> { processWorkFn.accept(work); }); @@ -109,38 +111,22 @@ private static ExecutableWork createWork(Consumer processWorkFn) { } @Test - public void logAndProcessFailure_doesNotRetryKeyTokenInvalidException() throws Throwable { + public void logAndProcessFailureBatch_doesNotRetryFailedWork() throws Throwable { Set executedWork = new HashSet<>(); ExecutableWork work = createWork(executedWork::add); + work.work().setFailed(); WorkFailureProcessor workFailureProcessor = createWorkFailureProcessor(streamingEngineFailureReporter()); Set invalidWork = new HashSet<>(); - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, work, new KeyTokenInvalidException("key"), invalidWork::add); + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, List.of(work), new RuntimeException(), invalidWork::add); assertThat(executedWork).isEmpty(); assertThat(invalidWork).containsExactly(work.work()); } @Test - public void logAndProcessFailure_doesNotRetryWhenWorkItemCancelled() throws Throwable { - Set executedWork = new HashSet<>(); - ExecutableWork work = createWork(executedWork::add); - WorkFailureProcessor workFailureProcessor = - createWorkFailureProcessor(streamingEngineFailureReporter()); - Set invalidWork = new HashSet<>(); - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, - work, - new WorkItemCancelledException(work.getWorkItem().getShardingKey()), - invalidWork::add); - - assertThat(executedWork).isEmpty(); - assertThat(invalidWork).containsExactly(work.work()); - } - - @Test - public void logAndProcessFailure_doesNotRetryOOM() { + public void logAndProcessFailureBatch_doesNotRetryOOM() { Set executedWork = new HashSet<>(); ExecutableWork work = createWork(executedWork::add); WorkFailureProcessor workFailureProcessor = @@ -149,69 +135,141 @@ public void logAndProcessFailure_doesNotRetryOOM() { assertThrows( OutOfMemoryError.class, () -> - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, work, new OutOfMemoryError(), invalidWork::add)); + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, + Arrays.asList(work), + new OutOfMemoryError(), + invalidWork::add)); assertThat(executedWork).isEmpty(); } @Test - public void logAndProcessFailure_doesNotRetryWhenFailureReporterMarksAsNonRetryable() + public void logAndProcessFailureBatch_doesNotRetryWhenFailureReporterMarksAsNonRetryable() throws Throwable { Set executedWork = new HashSet<>(); ExecutableWork work = createWork(executedWork::add); WorkFailureProcessor workFailureProcessor = createWorkFailureProcessor(streamingApplianceFailureReporter(true)); Set invalidWork = new HashSet<>(); - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, work, new RuntimeException(), invalidWork::add); + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, Arrays.asList(work), new RuntimeException(), invalidWork::add); assertThat(executedWork).isEmpty(); assertThat(invalidWork).containsExactly(work.work()); } @Test - public void logAndProcessFailure_doesNotRetryAfterLocalRetryTimeout() throws Throwable { + public void logAndProcessFailureBatch_doesNotRetryAfterLocalRetryTimeout() throws Throwable { Set executedWork = new HashSet<>(); ExecutableWork veryOldWork = createWork(() -> Instant.now().minus(Duration.standardDays(30)), executedWork::add); WorkFailureProcessor workFailureProcessor = createWorkFailureProcessor(streamingEngineFailureReporter()); Set invalidWork = new HashSet<>(); - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, veryOldWork, new RuntimeException(), invalidWork::add); + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, + Arrays.asList(veryOldWork), + new RuntimeException(), + invalidWork::add); assertThat(executedWork).isEmpty(); assertThat(invalidWork).contains(veryOldWork.work()); } @Test - public void logAndProcessFailure_retriesOnUncaughtUnhandledException_streamingEngine() + public void logAndProcessFailureBatch_retriesOnUncaughtUnhandledException_streamingEngine() throws Throwable { CountDownLatch runWork = new CountDownLatch(1); ExecutableWork work = createWork(ignored -> runWork.countDown()); WorkFailureProcessor workFailureProcessor = createWorkFailureProcessor(streamingEngineFailureReporter()); Set invalidWork = new HashSet<>(); - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, work, new RuntimeException(), invalidWork::add); + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, Arrays.asList(work), new RuntimeException(), invalidWork::add); runWork.await(); assertThat(invalidWork).isEmpty(); } @Test - public void logAndProcessFailure_retriesOnUncaughtUnhandledException_streamingAppliance() + public void logAndProcessFailureBatch_retriesOnUncaughtUnhandledException_streamingAppliance() throws Throwable { CountDownLatch runWork = new CountDownLatch(1); ExecutableWork work = createWork(ignored -> runWork.countDown()); WorkFailureProcessor workFailureProcessor = createWorkFailureProcessor(streamingApplianceFailureReporter(false)); Set invalidWork = new HashSet<>(); - workFailureProcessor.logAndProcessFailure( - DEFAULT_COMPUTATION_ID, work, new RuntimeException(), invalidWork::add); + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, Arrays.asList(work), new RuntimeException(), invalidWork::add); + + runWork.await(); + assertThat(invalidWork).isEmpty(); + } + + @Test + public void logAndProcessFailureBatch_retryAll() throws Throwable { + CountDownLatch runWork1 = new CountDownLatch(1); + CountDownLatch runWork2 = new CountDownLatch(1); + ExecutableWork work1 = createWork(ignored -> runWork1.countDown()); + ExecutableWork work2 = createWork(ignored -> runWork2.countDown()); + + WorkFailureProcessor workFailureProcessor = + createWorkFailureProcessor(streamingEngineFailureReporter()); + Set invalidWork = new HashSet<>(); + + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, + Arrays.asList(work1, work2), + new RuntimeException(), + invalidWork::add); + + runWork1.await(); + runWork2.await(); + assertThat(invalidWork).isEmpty(); + } + + @Test + public void logAndProcessFailureBatch_mixRetryAndAbort() throws Throwable { + CountDownLatch runWork1 = new CountDownLatch(1); + Set executedWork2 = new HashSet<>(); + ExecutableWork work1 = createWork(ignored -> runWork1.countDown()); + ExecutableWork work2 = createWork(executedWork2::add); + work2.work().setFailed(); + + WorkFailureProcessor workFailureProcessor = + createWorkFailureProcessor(streamingEngineFailureReporter()); + Set invalidWork = new HashSet<>(); + + workFailureProcessor.logAndProcessFailureBatch( + DEFAULT_COMPUTATION_ID, + Arrays.asList(work1, work2), + new RuntimeException(), + invalidWork::add); + + runWork1.await(); + 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, + Arrays.asList(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(); } } 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 f32282056e4f..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 @@ -54,6 +54,7 @@ import org.apache.beam.runners.direct.Clock; 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.HashBasedTable; +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.Table; import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.ThreadFactoryBuilder; import org.joda.time.Duration; @@ -137,7 +138,8 @@ private ExecutableWork createOldWork( Work.createProcessingContext( "computationId", new FakeGetDataClient(), ignored -> {}, heartbeatSender), false, - ActiveWorkRefresherTest::aLongTimeAgo), + ActiveWorkRefresherTest::aLongTimeAgo, + ImmutableList.of()), (work, handle) -> { processWork.accept(work); });