This is an automated email from the ASF dual-hosted git repository.
damccorm pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/master by this push:
new eda08d8a3c6 Changes SplittableDoFn to call TruncateRestriction on
drain (#39535)
eda08d8a3c6 is described below
commit eda08d8a3c60fe42a7848e162a10972f0bc4bb0c
Author: Andrew Crites <[email protected]>
AuthorDate: Wed Jul 29 12:11:16 2026 -0700
Changes SplittableDoFn to call TruncateRestriction on drain (#39535)
* Changes SplittableDoFn to call TruncateRestriction first when getting a
timer caused by drain. We then pass the residual restriction (if present) to
ProcessElement.
* Adds the rest of a comment string.
---
.../core/SplittableParDoViaKeyedWorkItems.java | 70 ++++-
.../runners/core/SplittableParDoProcessFnTest.java | 283 +++++++++++++++++++--
2 files changed, 333 insertions(+), 20 deletions(-)
diff --git
a/runners/core-java/src/main/java/org/apache/beam/runners/core/SplittableParDoViaKeyedWorkItems.java
b/runners/core-java/src/main/java/org/apache/beam/runners/core/SplittableParDoViaKeyedWorkItems.java
index a750b01963f..4eea2c0509d 100644
---
a/runners/core-java/src/main/java/org/apache/beam/runners/core/SplittableParDoViaKeyedWorkItems.java
+++
b/runners/core-java/src/main/java/org/apache/beam/runners/core/SplittableParDoViaKeyedWorkItems.java
@@ -469,8 +469,9 @@ public class SplittableParDoViaKeyedWorkItems {
restrictionState.readLater();
watermarkEstimatorState.readLater();
WindowedValue<InputT> read = elementState.read();
+ RestrictionT restriction = restrictionState.read();
if (timer.causedByDrain() == CausedByDrain.CAUSED_BY_DRAIN) {
- read =
+ WindowedValue<InputT> drainRead =
WindowedValues.of(
read.getValue(),
read.getTimestamp(),
@@ -481,8 +482,73 @@ public class SplittableParDoViaKeyedWorkItems {
CausedByDrain.CAUSED_BY_DRAIN,
read.getOpenTelemetryContext(),
read.getValueKind());
+ RestrictionTracker.TruncateResult<RestrictionT> truncateResult =
+ invoker.invokeTruncateRestriction(
+ new BaseArgumentProvider<InputT, OutputT>() {
+ @Override
+ public InputT element(DoFn<InputT, OutputT> doFn) {
+ return drainRead.getValue();
+ }
+
+ @Override
+ public Object restriction() {
+ return restriction;
+ }
+
+ @Override
+ public RestrictionTracker<?, ?> restrictionTracker() {
+ return invoker.invokeNewTracker(this);
+ }
+
+ @Override
+ public Instant timestamp(DoFn<InputT, OutputT> doFn) {
+ return drainRead.getTimestamp();
+ }
+
+ @Override
+ public PipelineOptions pipelineOptions() {
+ return c.getPipelineOptions();
+ }
+
+ @Override
+ public PaneInfo paneInfo(DoFn<InputT, OutputT> doFn) {
+ return drainRead.getPaneInfo();
+ }
+
+ @Override
+ public BoundedWindow window() {
+ return Iterables.getOnlyElement(drainRead.getWindows());
+ }
+
+ @Override
+ public Object sideInput(String tagId) {
+ PCollectionView<?> view = sideInputMapping.get(tagId);
+ if (view == null) {
+ throw new IllegalArgumentException(
+ "calling getSideInput() with unknown view");
+ }
+ return sideInputReader.get(
+ view,
view.getWindowMappingFn().getSideInputWindow(window()));
+ }
+
+ @Override
+ public String getErrorContext() {
+ return ProcessFn.class.getSimpleName() +
".invokeTruncateRestriction";
+ }
+ });
+ if (truncateResult == null) {
+ elementState.clear();
+ restrictionState.clear();
+ watermarkEstimatorState.clear();
+ holdState.clear();
+ return;
+ }
+ RestrictionT truncatedRestriction =
truncateResult.getTruncatedRestriction();
+ elementAndRestriction = KV.of(drainRead, truncatedRestriction);
+ restrictionState.write(truncatedRestriction);
+ } else {
+ elementAndRestriction = KV.of(read, restriction);
}
- elementAndRestriction = KV.of(read, restrictionState.read());
watermarkEstimatorStateT = watermarkEstimatorState.read();
}
diff --git
a/runners/core-java/src/test/java/org/apache/beam/runners/core/SplittableParDoProcessFnTest.java
b/runners/core-java/src/test/java/org/apache/beam/runners/core/SplittableParDoProcessFnTest.java
index 381e41c9870..247412a85dc 100644
---
a/runners/core-java/src/test/java/org/apache/beam/runners/core/SplittableParDoProcessFnTest.java
+++
b/runners/core-java/src/test/java/org/apache/beam/runners/core/SplittableParDoProcessFnTest.java
@@ -36,6 +36,7 @@ import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.List;
+import java.util.Map;
import java.util.NoSuchElementException;
import java.util.concurrent.Executors;
import org.apache.beam.runners.core.SplittableParDoViaKeyedWorkItems.ProcessFn;
@@ -48,12 +49,15 @@ import org.apache.beam.sdk.io.range.OffsetRange;
import org.apache.beam.sdk.state.TimeDomain;
import org.apache.beam.sdk.testing.ResetDateTimeProvider;
import org.apache.beam.sdk.testing.TestPipeline;
+import org.apache.beam.sdk.transforms.Create;
import org.apache.beam.sdk.transforms.DoFn;
import org.apache.beam.sdk.transforms.DoFnTester;
+import org.apache.beam.sdk.transforms.View;
import org.apache.beam.sdk.transforms.splittabledofn.HasDefaultTracker;
import org.apache.beam.sdk.transforms.splittabledofn.ManualWatermarkEstimator;
import org.apache.beam.sdk.transforms.splittabledofn.OffsetRangeTracker;
import org.apache.beam.sdk.transforms.splittabledofn.RestrictionTracker;
+import
org.apache.beam.sdk.transforms.splittabledofn.RestrictionTracker.IsBounded;
import org.apache.beam.sdk.transforms.splittabledofn.SplitResult;
import org.apache.beam.sdk.transforms.splittabledofn.WatermarkEstimators;
import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
@@ -152,6 +156,29 @@ public class SplittableParDoProcessFnTest {
int maxOutputsPerBundle,
Duration maxBundleDuration)
throws Exception {
+ this(
+ currentProcessingTime,
+ fn,
+ inputCoder,
+ restrictionCoder,
+ watermarkEstimatorStateCoder,
+ maxOutputsPerBundle,
+ maxBundleDuration,
+ Collections.emptyMap(),
+ NullSideInputReader.empty());
+ }
+
+ ProcessFnTester(
+ Instant currentProcessingTime,
+ final DoFn<InputT, OutputT> fn,
+ Coder<InputT> inputCoder,
+ Coder<RestrictionT> restrictionCoder,
+ Coder<WatermarkEstimatorStateT> watermarkEstimatorStateCoder,
+ int maxOutputsPerBundle,
+ Duration maxBundleDuration,
+ Map<String, PCollectionView<?>> sideInputMapping,
+ SideInputReader sideInputReader)
+ throws Exception {
// The exact windowing strategy doesn't matter in this test, but it
should be able to
// encode IntervalWindow's because that's what all tests here use.
WindowingStrategy<InputT, BoundedWindow> windowingStrategy =
@@ -163,35 +190,20 @@ public class SplittableParDoProcessFnTest {
restrictionCoder,
watermarkEstimatorStateCoder,
windowingStrategy,
- Collections.emptyMap());
+ sideInputMapping);
this.tester = DoFnTester.of(processFn);
this.timerInternals = new InMemoryTimerInternals();
this.stateInternals = new TestInMemoryStateInternals<>("dummy");
processFn.setStateInternalsFactory(key -> stateInternals);
processFn.setTimerInternalsFactory(key -> timerInternals);
- processFn.setSideInputReader(NullSideInputReader.empty());
+ processFn.setSideInputReader(sideInputReader);
processFn.setProcessElementInvoker(
new OutputAndTimeBoundedSplittableProcessElementInvoker<>(
fn,
tester.getPipelineOptions(),
new DoFnTesterWindowedValueReceiver(tester),
tester.getMainOutputTag(),
- new SideInputReader() {
- @Override
- public <T> T get(PCollectionView<T> view, BoundedWindow
window) {
- throw new NoSuchElementException();
- }
-
- @Override
- public <T> boolean contains(PCollectionView<T> view) {
- return false;
- }
-
- @Override
- public boolean isEmpty() {
- return true;
- }
- },
+ sideInputReader,
Executors.newSingleThreadScheduledExecutor(Executors.defaultThreadFactory()),
maxOutputsPerBundle,
maxBundleDuration,
@@ -790,4 +802,239 @@ public class SplittableParDoProcessFnTest {
assertEquals(7.0, backlogs.get(0), 0.001);
}
}
+
+ private static class TruncateFn extends DoFn<Integer, String> {
+ private final boolean truncateToNull;
+ private final List<String> calls = new ArrayList<>();
+
+ public TruncateFn(boolean truncateToNull) {
+ this.truncateToNull = truncateToNull;
+ }
+
+ @ProcessElement
+ public ProcessContinuation process(
+ ProcessContext c, RestrictionTracker<OffsetRange, Long> tracker) {
+ for (long i = tracker.currentRestriction().getFrom();
tracker.tryClaim(i); ++i) {
+ c.output(c.element() + ":" + i);
+ if (i == 2) {
+ return resume();
+ }
+ }
+ return stop();
+ }
+
+ @GetInitialRestriction
+ public OffsetRange getInitialRestriction() {
+ return new OffsetRange(0, 10);
+ }
+
+ @NewTracker
+ public OffsetRangeTracker newTracker(@Restriction OffsetRange range) {
+ return new OffsetRangeTracker(range);
+ }
+
+ @TruncateRestriction
+ public RestrictionTracker.TruncateResult<OffsetRange> truncate(
+ @Restriction OffsetRange restriction, @Element Integer element) {
+ calls.add("truncate:" + element + ":" + restriction);
+ if (truncateToNull) {
+ return null;
+ }
+ // Truncate so that we only process one more element.
+ return RestrictionTracker.TruncateResult.of(
+ new OffsetRange(restriction.getFrom(), restriction.getFrom() + 1));
+ }
+ }
+
+ @Test
+ public void testTruncateRestrictionOnDrain() throws Exception {
+ TruncateFn fn = new TruncateFn(false);
+ Instant base = Instant.now();
+
+ try (ProcessFnTester<Integer, String, OffsetRange, Long, Void> tester =
+ new ProcessFnTester<>(
+ base,
+ fn,
+ BigEndianIntegerCoder.of(),
+ SerializableCoder.of(OffsetRange.class),
+ VoidCoder.of(),
+ MAX_OUTPUTS_PER_BUNDLE,
+ MAX_BUNDLE_DURATION)) {
+ tester.startElement(42, new OffsetRange(0, 10));
+ assertThat(tester.takeOutputElements(), contains("42:0", "42:1",
"42:2"));
+
+ assertTrue(tester.advanceDrain());
+ assertThat(tester.takeOutputElements(), contains("42:3"));
+ assertEquals(Collections.singletonList("truncate:42:[3, 10)"), fn.calls);
+ assertEquals(null, tester.getWatermarkHold());
+ }
+ }
+
+ @Test
+ public void testTruncateRestrictionReturnsNullOnDrain() throws Exception {
+ TruncateFn fn = new TruncateFn(true);
+ Instant base = Instant.now();
+
+ try (ProcessFnTester<Integer, String, OffsetRange, Long, Void> tester =
+ new ProcessFnTester<>(
+ base,
+ fn,
+ BigEndianIntegerCoder.of(),
+ SerializableCoder.of(OffsetRange.class),
+ VoidCoder.of(),
+ MAX_OUTPUTS_PER_BUNDLE,
+ MAX_BUNDLE_DURATION)) {
+ tester.startElement(42, new OffsetRange(0, 10));
+ assertThat(tester.takeOutputElements(), contains("42:0", "42:1",
"42:2"));
+
+ assertTrue(tester.advanceDrain());
+ assertTrue(tester.takeOutputElements().isEmpty());
+ assertEquals(Collections.singletonList("truncate:42:[3, 10)"), fn.calls);
+ assertEquals(null, tester.getWatermarkHold());
+ }
+ }
+
+ private static class TruncateWithSideInputFn extends DoFn<Integer, String> {
+ private final List<String> calls = new ArrayList<>();
+
+ @ProcessElement
+ public ProcessContinuation process(
+ ProcessContext c, RestrictionTracker<OffsetRange, Long> tracker) {
+ for (long i = tracker.currentRestriction().getFrom();
tracker.tryClaim(i); ++i) {
+ c.output(c.element() + ":" + i);
+ if (i == 2) {
+ return resume();
+ }
+ }
+ return stop();
+ }
+
+ @GetInitialRestriction
+ public OffsetRange getInitialRestriction() {
+ return new OffsetRange(0, 10);
+ }
+
+ @NewTracker
+ public OffsetRangeTracker newTracker(@Restriction OffsetRange range) {
+ return new OffsetRangeTracker(range);
+ }
+
+ @TruncateRestriction
+ public RestrictionTracker.TruncateResult<OffsetRange> truncate(
+ @Restriction OffsetRange restriction,
+ @Element Integer element,
+ @SideInput("sideInput") String sideInput) {
+ calls.add("truncate:" + element + ":" + sideInput + ":" + restriction);
+ return RestrictionTracker.TruncateResult.of(
+ new OffsetRange(restriction.getFrom(), restriction.getFrom() + 1));
+ }
+ }
+
+ @Test
+ public void testTruncateRestrictionWithSideInputOnDrain() throws Exception {
+ TruncateWithSideInputFn fn = new TruncateWithSideInputFn();
+ Instant base = Instant.now();
+ PCollectionView<String> view =
+
TestPipeline.create().apply(Create.of("sideValue")).apply(View.asSingleton());
+ Map<String, PCollectionView<?>> sideInputMapping =
Collections.singletonMap("sideInput", view);
+ SideInputReader sideInputReader =
+ new SideInputReader() {
+ @Override
+ public <T> T get(PCollectionView<T> v, BoundedWindow window) {
+ if (v.equals(view)) {
+ return (T) "sideValue";
+ }
+ throw new NoSuchElementException();
+ }
+
+ @Override
+ public <T> boolean contains(PCollectionView<T> v) {
+ return v.equals(view);
+ }
+
+ @Override
+ public boolean isEmpty() {
+ return false;
+ }
+ };
+
+ try (ProcessFnTester<Integer, String, OffsetRange, Long, Void> tester =
+ new ProcessFnTester<>(
+ base,
+ fn,
+ BigEndianIntegerCoder.of(),
+ SerializableCoder.of(OffsetRange.class),
+ VoidCoder.of(),
+ MAX_OUTPUTS_PER_BUNDLE,
+ MAX_BUNDLE_DURATION,
+ sideInputMapping,
+ sideInputReader)) {
+ tester.startElement(42, new OffsetRange(0, 10));
+ assertThat(tester.takeOutputElements(), contains("42:0", "42:1",
"42:2"));
+
+ assertTrue(tester.advanceDrain());
+ assertThat(tester.takeOutputElements(), contains("42:3"));
+ assertEquals(Collections.singletonList("truncate:42:sideValue:[3, 10)"),
fn.calls);
+ }
+ }
+
+ private static class UnboundedOffsetRangeTracker extends OffsetRangeTracker {
+ public UnboundedOffsetRangeTracker(OffsetRange range) {
+ super(range);
+ }
+
+ @Override
+ public IsBounded isBounded() {
+ return IsBounded.UNBOUNDED;
+ }
+ }
+
+ // Tests that if we don't override TruncateRestriction, the default
TruncateRestriction
+ // implementation is used (which for unbounded restrictions stops processing
immediately).
+ private static class DefaultTruncateUnboundedFn extends DoFn<Integer,
String> {
+ @ProcessElement
+ public ProcessContinuation process(
+ ProcessContext c, RestrictionTracker<OffsetRange, Long> tracker) {
+ for (long i = tracker.currentRestriction().getFrom();
tracker.tryClaim(i); ++i) {
+ c.output(c.element() + ":" + i);
+ if (i == 2) {
+ return resume();
+ }
+ }
+ return stop();
+ }
+
+ @GetInitialRestriction
+ public OffsetRange getInitialRestriction() {
+ return new OffsetRange(0, 10);
+ }
+
+ @NewTracker
+ public RestrictionTracker<OffsetRange, Long> newTracker(@Restriction
OffsetRange range) {
+ return new UnboundedOffsetRangeTracker(range);
+ }
+ }
+
+ @Test
+ public void testDefaultTruncateRestrictionUnboundedStopsOnDrain() throws
Exception {
+ DefaultTruncateUnboundedFn fn = new DefaultTruncateUnboundedFn();
+ Instant base = Instant.now();
+
+ try (ProcessFnTester<Integer, String, OffsetRange, Long, Void> tester =
+ new ProcessFnTester<>(
+ base,
+ fn,
+ BigEndianIntegerCoder.of(),
+ SerializableCoder.of(OffsetRange.class),
+ VoidCoder.of(),
+ MAX_OUTPUTS_PER_BUNDLE,
+ MAX_BUNDLE_DURATION)) {
+ tester.startElement(42, new OffsetRange(0, 10));
+ assertThat(tester.takeOutputElements(), contains("42:0", "42:1",
"42:2"));
+
+ assertTrue(tester.advanceDrain());
+ assertTrue(tester.takeOutputElements().isEmpty());
+ assertEquals(null, tester.getWatermarkHold());
+ }
+ }
}