This is an automated email from the ASF dual-hosted git repository.
Abacn 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 dfc4d667cd9 [Spark 4] Support stateful ParDo in Structured Streaming
via transformWithState (#40281)
dfc4d667cd9 is described below
commit dfc4d667cd947d955687333f9cc8bd064242c213
Author: Tobias Kaymak <[email protected]>
AuthorDate: Tue Oct 6 23:16:02 2026 +0200
[Spark 4] Support stateful ParDo in Structured Streaming via
transformWithState (#40281)
* [Spark 4] Support stateful ParDo in Structured Streaming via
transformWithState
Stateful ParDo runs through Spark transformWithState. User state reuses
SparkStateInternals over a MapState and timers reuse SparkTimerInternals
persisted in a ValueState with one Spark wake up per key. Unsupported
timer domains, window expiration, time sorted input and merging windows
fail at translation naming #36841.
---
...a_ValidatesRunner_SparkStructuredStreaming.json | 3 +-
.../translation/PipelineTranslatorStreaming.java | 65 +++-
.../StatefulParDoStreamingTranslator.java | 86 ++++++
.../streaming/state/BeamStatefulProcessor.java | 330 +++++++++++++++++++++
.../io/streaming/TestUnboundedSource.java | 16 +-
.../PipelineTranslatorStreamingTest.java | 97 +++++-
.../streaming/StatefulParDoStreamingTest.java | 242 +++++++++++++++
.../spark/stateful/SparkStateInternals.java | 78 ++++-
.../spark/stateful/SparkTimerInternals.java | 5 +
.../translation/batch/DoFnRunnerFactory.java | 15 +-
.../batch/StatefulDoFnGroupFunction.java | 118 +-------
.../translation/batch/StatefulTaskRunner.java | 145 +++++++++
12 files changed, 1061 insertions(+), 139 deletions(-)
diff --git
a/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
b/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
index ea9b041f3fc..2b28539c14f 100644
---
a/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
+++
b/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
@@ -10,5 +10,6 @@
"https://github.com/apache/beam/pull/35159": "moving WindowedValue and
making an interface",
"https://github.com/apache/beam/pull/39793": "noting that PR #39793 should
run this test",
"https://github.com/apache/beam/pull/40103": "noting that PR #40103 should
run this test",
- "https://github.com/apache/beam/issues/40427": "Spark ValidatesRunner tests
run serially"
+ "https://github.com/apache/beam/issues/40427": "Spark ValidatesRunner tests
run serially",
+ "https://github.com/apache/beam/pull/40281": "noting that PR #40281 should
run this test"
}
diff --git
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreaming.java
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreaming.java
index 3529daaff31..5a18b5da6fd 100644
---
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreaming.java
+++
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreaming.java
@@ -21,8 +21,12 @@ import java.util.Collection;
import org.apache.beam.runners.spark.SparkCommonPipelineOptions;
import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.PipelineTranslatorCommon;
import
org.apache.beam.runners.spark.structuredstreaming.translation.streaming.ReadUnboundedTranslator;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.streaming.StatefulParDoStreamingTranslator;
import org.apache.beam.sdk.annotations.Internal;
+import org.apache.beam.sdk.state.TimeDomain;
+import org.apache.beam.sdk.state.TimerSpec;
import org.apache.beam.sdk.transforms.Combine;
+import org.apache.beam.sdk.transforms.DoFn;
import org.apache.beam.sdk.transforms.GroupByKey;
import org.apache.beam.sdk.transforms.Impulse;
import org.apache.beam.sdk.transforms.PTransform;
@@ -30,8 +34,11 @@ import org.apache.beam.sdk.transforms.ParDo;
import org.apache.beam.sdk.transforms.reflect.DoFnSignature;
import org.apache.beam.sdk.transforms.reflect.DoFnSignatures;
import org.apache.beam.sdk.util.construction.SplittableParDo;
+import org.apache.beam.sdk.values.KV;
+import org.apache.beam.sdk.values.PCollection;
import org.apache.beam.sdk.values.PInput;
import org.apache.beam.sdk.values.POutput;
+import org.apache.beam.sdk.values.WindowingStrategy;
import org.apache.spark.sql.SparkSession;
import org.checkerframework.checker.nullness.qual.Nullable;
@@ -47,6 +54,18 @@ public class PipelineTranslatorStreaming extends
PipelineTranslatorCommon {
@SuppressWarnings("rawtypes")
private static final TransformTranslator READ_UNBOUNDED = new
ReadUnboundedTranslator<>();
+ @SuppressWarnings({"rawtypes", "unchecked"})
+ private static final TransformTranslator STATEFUL_PAR_DO =
+ new StatefulParDoStreamingTranslator<Object, Object, Object>() {
+ @Override
+ protected void translate(
+ ParDo.MultiOutput<KV<Object, Object>, Object> transform, Context
cxt) {
+ PCollection<?> input = (PCollection<?>) cxt.getInput();
+ checkNotMergingWindows(input.getWindowingStrategy());
+ super.translate(transform, cxt);
+ }
+ };
+
private static final String NOT_SUPPORTED =
" is not supported by the Spark 4 streaming runner yet, see"
+ " https://github.com/apache/beam/issues/36841";
@@ -83,10 +102,6 @@ public class PipelineTranslatorStreaming extends
PipelineTranslatorCommon {
if (transform instanceof ParDo.MultiOutput) {
ParDo.MultiOutput<?, ?> parDo = (ParDo.MultiOutput<?, ?>) transform;
DoFnSignature signature = DoFnSignatures.signatureForDoFn(parDo.getFn());
- if (signature.usesState() || signature.usesTimers()) {
- throw new UnsupportedOperationException(
- "Stateful ParDo (" + signature.fnClass().getName() + ")" +
NOT_SUPPORTED);
- }
if (!parDo.getSideInputs().isEmpty()) {
throw new UnsupportedOperationException(
"ParDo with side inputs (" + signature.fnClass().getName() + ")" +
NOT_SUPPORTED);
@@ -98,11 +113,53 @@ public class PipelineTranslatorStreaming extends
PipelineTranslatorCommon {
+ ")"
+ NOT_SUPPORTED);
}
+ if (signature.processElement().requiresTimeSortedInput()) {
+ throw new UnsupportedOperationException(
+ "@RequiresTimeSortedInput (" + signature.fnClass().getName() + ")"
+ NOT_SUPPORTED);
+ }
+ if (signature.onWindowExpiration() != null) {
+ throw new UnsupportedOperationException(
+ "@OnWindowExpiration (" + signature.fnClass().getName() + ")" +
NOT_SUPPORTED);
+ }
+ checkNoProcessingTimeTimers(parDo.getFn(), signature);
+
+ if (signature.usesState() || signature.usesTimers()) {
+ return STATEFUL_PAR_DO;
+ }
}
return super.getTransformTranslator(transform);
}
+ private static void checkNoProcessingTimeTimers(DoFn<?, ?> doFn,
DoFnSignature signature) {
+ for (DoFnSignature.TimerDeclaration timer :
signature.timerDeclarations().values()) {
+ TimerSpec spec = DoFnSignatures.getTimerSpecOrThrow(timer, doFn);
+ if (spec.getTimeDomain() != TimeDomain.EVENT_TIME) {
+ throw new UnsupportedOperationException(
+ spec.getTimeDomain() + " timer @TimerId(\"" + timer.id() + "\")" +
NOT_SUPPORTED);
+ }
+ }
+ for (DoFnSignature.TimerFamilyDeclaration family :
+ signature.timerFamilyDeclarations().values()) {
+ TimerSpec spec = DoFnSignatures.getTimerFamilySpecOrThrow(family, doFn);
+ if (spec.getTimeDomain() != TimeDomain.EVENT_TIME) {
+ throw new UnsupportedOperationException(
+ spec.getTimeDomain()
+ + " timer family @TimerFamily(\""
+ + family.id()
+ + "\")"
+ + NOT_SUPPORTED);
+ }
+ }
+ }
+
+ private static void checkNotMergingWindows(WindowingStrategy<?, ?>
windowingStrategy) {
+ if (windowingStrategy.needsMerge()) {
+ throw new UnsupportedOperationException(
+ "Stateful ParDo over merging windows" + NOT_SUPPORTED);
+ }
+ }
+
@Override
protected EvaluationContext createEvaluationContext(
Collection<? extends EvaluationContext.NamedDataset<?>> leaves,
diff --git
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/StatefulParDoStreamingTranslator.java
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/StatefulParDoStreamingTranslator.java
new file mode 100644
index 00000000000..1709ec57003
--- /dev/null
+++
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/StatefulParDoStreamingTranslator.java
@@ -0,0 +1,86 @@
+/*
+ * 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.spark.structuredstreaming.translation.streaming;
+
+import
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.TransformTranslator;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.DoFnRunnerFactory;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.functions.SparkSideInputReader;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.streaming.state.BeamStatefulProcessor;
+import org.apache.beam.sdk.annotations.Internal;
+import org.apache.beam.sdk.coders.Coder;
+import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.transforms.DoFn;
+import org.apache.beam.sdk.transforms.ParDo;
+import org.apache.beam.sdk.values.KV;
+import org.apache.beam.sdk.values.PCollection;
+import org.apache.beam.sdk.values.PCollectionTuple;
+import org.apache.beam.sdk.values.TupleTag;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.beam.sdk.values.WindowingStrategy;
+import org.apache.spark.api.java.function.MapFunction;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Encoder;
+import org.apache.spark.sql.streaming.OutputMode;
+import org.apache.spark.sql.streaming.TimeMode;
+
+/** Translates a stateful {@link ParDo.MultiOutput} for Spark 4 Structured
Streaming. */
+@Internal
+public class StatefulParDoStreamingTranslator<K, V, OutputT>
+ extends TransformTranslator<
+ PCollection<? extends KV<K, V>>, PCollectionTuple,
ParDo.MultiOutput<KV<K, V>, OutputT>> {
+
+ public StatefulParDoStreamingTranslator() {
+ super(0.2f);
+ }
+
+ @Override
+ protected void translate(ParDo.MultiOutput<KV<K, V>, OutputT> transform,
Context cxt) {
+ @SuppressWarnings("unchecked")
+ PCollection<KV<K, V>> input = (PCollection<KV<K, V>>) cxt.getInput();
+ WindowingStrategy<?, ?> windowing = input.getWindowingStrategy();
+ DoFn<KV<K, V>, OutputT> doFn = transform.getFn();
+
+ KvCoder<K, V> inputCoder = (KvCoder<K, V>) input.getCoder();
+ Encoder<K> keyEnc = cxt.keyEncoderOf(inputCoder);
+ TupleTag<OutputT> mainOutputTag = transform.getMainOutputTag();
+ PCollection<OutputT> output = cxt.getOutput(mainOutputTag);
+ Coder<OutputT> outputCoder = output.getCoder();
+
+ DoFnRunnerFactory<KV<K, V>, OutputT> runnerFactory =
+ DoFnRunnerFactory.simple(
+ cxt.getCurrentTransform(), input, SparkSideInputReader.empty(),
false);
+ MetricsAccumulator metrics =
MetricsAccumulator.getInstance(cxt.getSparkSession());
+
+ BeamStatefulProcessor<K, V, OutputT> processor =
+ new BeamStatefulProcessor<>(
+ doFn, runnerFactory, inputCoder, windowing,
cxt.getOptionsSupplier(), metrics);
+
+ MapFunction<WindowedValue<KV<K, V>>, K> keyFn = v -> v.getValue().getKey();
+ Dataset<WindowedValue<OutputT>> result =
+ cxt.getDataset(input)
+ .groupByKey(keyFn, keyEnc)
+ .transformWithState(
+ processor,
+ TimeMode.EventTime(),
+ OutputMode.Append(),
+ cxt.windowedEncoder(outputCoder));
+
+ cxt.putDataset(output, result);
+ }
+}
diff --git
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/state/BeamStatefulProcessor.java
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/state/BeamStatefulProcessor.java
new file mode 100644
index 00000000000..c307ae6f37b
--- /dev/null
+++
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/state/BeamStatefulProcessor.java
@@ -0,0 +1,330 @@
+/*
+ * 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.spark.structuredstreaming.translation.streaming.state;
+
+import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;
+
+import java.util.ArrayList;
+import java.util.Collection;
+import java.util.Comparator;
+import java.util.HashSet;
+import java.util.List;
+import java.util.Set;
+import java.util.function.Supplier;
+import org.apache.beam.runners.core.DoFnRunner;
+import org.apache.beam.runners.core.DoFnRunners;
+import org.apache.beam.runners.core.StatefulDoFnRunner;
+import org.apache.beam.runners.core.TimerInternals.TimerData;
+import org.apache.beam.runners.core.TimerInternals.TimerDataCoderV2;
+import org.apache.beam.runners.spark.coders.CoderHelpers;
+import org.apache.beam.runners.spark.stateful.SparkStateInternals;
+import org.apache.beam.runners.spark.stateful.SparkStateInternals.StateCells;
+import org.apache.beam.runners.spark.stateful.SparkTimerInternals;
+import
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.DoFnRunnerFactory;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.DoFnRunnerFactory.DoFnRunnerWithTeardown;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.StatefulTaskRunner;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.utils.ScalaInterop;
+import org.apache.beam.sdk.coders.Coder;
+import org.apache.beam.sdk.coders.ListCoder;
+import org.apache.beam.sdk.options.PipelineOptions;
+import org.apache.beam.sdk.state.TimeDomain;
+import org.apache.beam.sdk.transforms.DoFn;
+import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
+import org.apache.beam.sdk.transforms.windowing.GlobalWindow;
+import org.apache.beam.sdk.util.WindowedValueMultiReceiver;
+import org.apache.beam.sdk.values.KV;
+import org.apache.beam.sdk.values.TupleTag;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.beam.sdk.values.WindowingStrategy;
+import org.apache.spark.sql.Encoders;
+import org.apache.spark.sql.streaming.ExpiredTimerInfo;
+import org.apache.spark.sql.streaming.MapState;
+import org.apache.spark.sql.streaming.OutputMode;
+import org.apache.spark.sql.streaming.StatefulProcessor;
+import org.apache.spark.sql.streaming.StatefulProcessorHandle;
+import org.apache.spark.sql.streaming.TTLConfig;
+import org.apache.spark.sql.streaming.TimeMode;
+import org.apache.spark.sql.streaming.TimerValues;
+import org.apache.spark.sql.streaming.ValueState;
+import org.checkerframework.checker.nullness.qual.Nullable;
+import org.joda.time.Instant;
+import scala.collection.Iterator;
+
+/** Runs a stateful {@link DoFn} within Spark 4 Structured Streaming {@code
transformWithState}. */
+public class BeamStatefulProcessor<K, V, OutputT>
+ extends StatefulProcessor<K, WindowedValue<KV<K, V>>,
WindowedValue<OutputT>> {
+
+ private static final String BEAM_STATE = "beamState";
+ private static final String BEAM_TIMERS = "beamTimers";
+
+ private final DoFn<KV<K, V>, OutputT> doFn;
+ private final DoFnRunnerFactory<KV<K, V>, OutputT> runnerFactory;
+ private final Coder<KV<K, V>> inputCoder;
+ private final WindowingStrategy<?, ?> windowingStrategy;
+ private final Supplier<PipelineOptions> optionsSupplier;
+ private final MetricsAccumulator metrics;
+ private final Coder<BoundedWindow> windowCoder;
+ private final TimerDataCoderV2 timerDataCoder;
+ private final StatefulTaskRunner<KV<K, V>, OutputT> taskRunner = new
StatefulTaskRunner<>();
+
+ private transient @Nullable MapState<String, byte[]> beamState;
+ private transient @Nullable ValueState<byte[]> beamTimers;
+ private transient @Nullable List<WindowedValue<OutputT>> currentOutputs;
+ private transient boolean needsBundleStart;
+
+ public BeamStatefulProcessor(
+ DoFn<KV<K, V>, OutputT> doFn,
+ DoFnRunnerFactory<KV<K, V>, OutputT> runnerFactory,
+ Coder<KV<K, V>> inputCoder,
+ WindowingStrategy<?, ?> windowingStrategy,
+ Supplier<PipelineOptions> optionsSupplier,
+ MetricsAccumulator metrics) {
+ this.doFn = doFn;
+ this.runnerFactory = runnerFactory;
+ this.inputCoder = inputCoder;
+ this.windowingStrategy = windowingStrategy;
+ this.optionsSupplier = optionsSupplier;
+ this.metrics = metrics;
+ @SuppressWarnings("unchecked")
+ Coder<BoundedWindow> windowCoder =
+ (Coder<BoundedWindow>) windowingStrategy.getWindowFn().windowCoder();
+ this.windowCoder = windowCoder;
+ this.timerDataCoder = TimerDataCoderV2.of(windowCoder);
+ }
+
+ @Override
+ public void init(OutputMode outputMode, TimeMode timeMode) {
+ beamState =
+ getHandle().getMapState(BEAM_STATE, Encoders.STRING(),
Encoders.BINARY(), TTLConfig.NONE());
+ beamTimers = getHandle().getValueState(BEAM_TIMERS, Encoders.BINARY(),
TTLConfig.NONE());
+ }
+
+ @Override
+ public Iterator<WindowedValue<OutputT>> handleInputRows(
+ K key, Iterator<WindowedValue<KV<K, V>>> rows, TimerValues timerValues) {
+ List<WindowedValue<OutputT>> outputs = new ArrayList<>();
+ SparkTimerInternals timerInternals =
+ SparkTimerInternals.forWatermark(new
Instant(timerValues.getCurrentWatermarkInMs()));
+ restoreTimers(timerInternals);
+
+ DoFnRunner<KV<K, V>, OutputT> runner = keyRunner(key, timerInternals,
outputs);
+ while (rows.hasNext()) {
+ runner.processElement(rows.next());
+ }
+ // The next bundle opens in keyRunner.
+ runner.finishBundle();
+ needsBundleStart = true;
+
+ persistTimers(timerInternals);
+ reconcileWakeupTimer(timerInternals, null);
+ return ScalaInterop.scalaIterator(outputs);
+ }
+
+ @Override
+ public Iterator<WindowedValue<OutputT>> handleExpiredTimer(
+ K key, TimerValues timerValues, ExpiredTimerInfo expiredTimerInfo) {
+ List<WindowedValue<OutputT>> outputs = new ArrayList<>();
+ Instant watermark = new Instant(timerValues.getCurrentWatermarkInMs());
+ SparkTimerInternals timerInternals =
SparkTimerInternals.forWatermark(watermark);
+ restoreTimers(timerInternals);
+
+ DoFnRunner<KV<K, V>, OutputT> runner = keyRunner(key, timerInternals,
outputs);
+ fireDueTimers(key, watermark, timerInternals, runner);
+ // The next bundle opens in keyRunner.
+ runner.finishBundle();
+ needsBundleStart = true;
+
+ persistTimers(timerInternals);
+ reconcileWakeupTimer(timerInternals, expiredTimerInfo.getExpiryTimeInMs());
+ return ScalaInterop.scalaIterator(outputs);
+ }
+
+ @Override
+ public void close() {
+ taskRunner.teardownOnce();
+ }
+
+ /** Creates the task runner on first use, {@code create} opens the first
bundle. */
+ private DoFnRunnerWithTeardown<KV<K, V>, OutputT> baseRunner(
+ SparkStateInternals<K> stateInternals, SparkTimerInternals
timerInternals) {
+ return taskRunner.getOrCreate(
+ ctx -> {
+ ctx.set(stateInternals, timerInternals);
+ return runnerFactory.create(
+ optionsSupplier.get(),
+ metrics,
+ new WindowedValueMultiReceiver() {
+ @Override
+ public <T> void output(TupleTag<T> tag, WindowedValue<T>
output) {
+ @SuppressWarnings("unchecked")
+ WindowedValue<OutputT> out = (WindowedValue<OutputT>) output;
+ checkStateNotNull(currentOutputs, "currentOutputs not
initialized").add(out);
+ }
+ },
+ ctx);
+ });
+ }
+
+ private DoFnRunner<KV<K, V>, OutputT> keyRunner(
+ K key, SparkTimerInternals timerInternals, List<WindowedValue<OutputT>>
outputs) {
+ this.currentOutputs = outputs;
+ MapState<String, byte[]> state = checkStateNotNull(beamState);
+ SparkStateInternals<K> stateInternals =
+ SparkStateInternals.forKey(key, new MapStateAdapter(state));
+ DoFnRunnerWithTeardown<KV<K, V>, OutputT> base =
baseRunner(stateInternals, timerInternals);
+ taskRunner.stepContext().set(stateInternals, timerInternals);
+
+ // Closed in handleInputRows and handleExpiredTimer.
+ if (needsBundleStart) {
+ needsBundleStart = false;
+ base.startBundle();
+ }
+
+ StatefulDoFnRunner.CleanupTimer<KV<K, V>> cleanupTimer =
+ new StatefulDoFnRunner.TimeInternalsCleanupTimer<KV<K, V>>(
+ timerInternals, windowingStrategy) {
+ @Override
+ public void setForWindow(KV<K, V> input, BoundedWindow window) {
+ // GlobalWindow state is never garbage collected, as in the Flink
runner.
+ if (!window.equals(GlobalWindow.INSTANCE)) {
+ super.setForWindow(input, window);
+ }
+ }
+ };
+ StatefulDoFnRunner.StateCleaner<BoundedWindow> stateCleaner =
+ new StatefulDoFnRunner.StateInternalsStateCleaner<>(doFn,
stateInternals, windowCoder);
+
+ return DoFnRunners.defaultStatefulDoFnRunner(
+ doFn,
+ inputCoder,
+ base,
+ taskRunner.stepContext(),
+ windowingStrategy,
+ cleanupTimer,
+ stateCleaner);
+ }
+
+ private void restoreTimers(SparkTimerInternals timerInternals) {
+ ValueState<byte[]> timersState = checkStateNotNull(beamTimers);
+ if (timersState.exists()) {
+ byte[] timerBytes = timersState.get();
+ if (timerBytes != null) {
+ List<TimerData> timers =
+ CoderHelpers.fromByteArray(timerBytes,
ListCoder.of(timerDataCoder));
+ timerInternals.addTimers(timers.iterator());
+ }
+ }
+ }
+
+ private void persistTimers(SparkTimerInternals timerInternals) {
+ ValueState<byte[]> timersState = checkStateNotNull(beamTimers);
+ Collection<TimerData> timers = timerInternals.getTimers();
+ if (timers.isEmpty()) {
+ if (timersState.exists()) {
+ timersState.clear();
+ }
+ } else {
+ timersState.update(
+ CoderHelpers.toByteArray(new ArrayList<>(timers),
ListCoder.of(timerDataCoder)));
+ }
+ }
+
+ private void fireDueTimers(
+ K key,
+ Instant watermark,
+ SparkTimerInternals timerInternals,
+ DoFnRunner<KV<K, V>, OutputT> runner) {
+ while (true) {
+ TimerData next =
+ timerInternals.getTimers().stream()
+ .filter(
+ t ->
+ t.getDomain().equals(TimeDomain.EVENT_TIME)
+ && watermark.isAfter(t.getTimestamp()))
+ .min(Comparator.comparing(TimerData::getTimestamp))
+ .orElse(null);
+ if (next == null) {
+ break;
+ }
+ timerInternals.deleteTimer(next);
+ StatefulTaskRunner.fireTimer(runner, key, next);
+ }
+ }
+
+ private void reconcileWakeupTimer(
+ SparkTimerInternals timerInternals, @Nullable Long firedExpiryMs) {
+ Long nextWakeupMs = null;
+ TimerData earliest =
+ timerInternals.getTimers().stream()
+ .filter(t -> t.getDomain().equals(TimeDomain.EVENT_TIME))
+ .min(Comparator.comparing(TimerData::getTimestamp))
+ .orElse(null);
+ if (earliest != null) {
+ nextWakeupMs = earliest.getTimestamp().getMillis() + 1;
+ }
+
+ StatefulProcessorHandle handle = getHandle();
+ Iterator<Object> it = handle.listTimers();
+ Set<Long> registered = new HashSet<>();
+ while (it.hasNext()) {
+ registered.add(((Number) it.next()).longValue());
+ }
+
+ for (Long expiry : registered) {
+ if (expiry.equals(firedExpiryMs)) {
+ continue;
+ }
+ if (nextWakeupMs != null && expiry.equals(nextWakeupMs)) {
+ continue;
+ }
+ handle.deleteTimer(expiry);
+ }
+
+ if (nextWakeupMs != null && !registered.contains(nextWakeupMs)) {
+ handle.registerTimer(nextWakeupMs);
+ }
+ }
+
+ private static class MapStateAdapter implements StateCells {
+ private final MapState<String, byte[]> mapState;
+
+ MapStateAdapter(MapState<String, byte[]> mapState) {
+ this.mapState = mapState;
+ }
+
+ @Override
+ public byte @Nullable [] get(String namespace, String stateId) {
+ return mapState.getValue(key(namespace, stateId));
+ }
+
+ @Override
+ public void put(String namespace, String stateId, byte[] value) {
+ mapState.updateValue(key(namespace, stateId), value);
+ }
+
+ @Override
+ public void remove(String namespace, String stateId) {
+ mapState.removeKey(key(namespace, stateId));
+ }
+
+ private static String key(String namespace, String stateId) {
+ return namespace.length() + ":" + namespace + stateId;
+ }
+ }
+}
diff --git
a/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/TestUnboundedSource.java
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/TestUnboundedSource.java
index 4a71dee5ac8..98f62501f37 100644
---
a/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/TestUnboundedSource.java
+++
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/TestUnboundedSource.java
@@ -56,6 +56,7 @@ public final class TestUnboundedSource extends
UnboundedSource<String, TestUnbou
private static final ConcurrentMap<String, List<Integer>> FINALIZED = new
ConcurrentHashMap<>();
private static final ConcurrentMap<String, AtomicInteger> CREATED = new
ConcurrentHashMap<>();
+ private static final ConcurrentMap<String, Integer> EXTENDED = new
ConcurrentHashMap<>();
private final String tag;
private final int shard;
@@ -66,6 +67,10 @@ public final class TestUnboundedSource extends
UnboundedSource<String, TestUnbou
this(tag, -1, shards, count / shards);
}
+ public static void extend(String tag, int count) {
+ EXTENDED.put(tag, count);
+ }
+
private TestUnboundedSource(String tag, int shard, int shards, int perShard)
{
this.tag = tag;
this.shard = shard;
@@ -114,6 +119,7 @@ public final class TestUnboundedSource extends
UnboundedSource<String, TestUnbou
public static void forget(String tag) {
FINALIZED.keySet().removeIf(key -> key.startsWith(tag + "/"));
CREATED.remove(tag);
+ EXTENDED.remove(tag);
}
private static String key(String tag, int shard) {
@@ -205,9 +211,14 @@ public final class TestUnboundedSource extends
UnboundedSource<String, TestUnbou
return advance();
}
+ private int limit() {
+ Integer extended = EXTENDED.get(source.tag);
+ return extended != null ? extended / source.shards : source.perShard;
+ }
+
@Override
public boolean advance() {
- if (next < source.perShard) {
+ if (next < limit()) {
current = next++;
return true;
}
@@ -227,8 +238,7 @@ public final class TestUnboundedSource extends
UnboundedSource<String, TestUnbou
if (current < 0) {
throw new NoSuchElementException();
}
- return new Instant(
- BASE_MILLIS + (source.shard * source.perShard + current) *
INTERVAL_MILLIS);
+ return new Instant(BASE_MILLIS + (source.shard * limit() + current) *
INTERVAL_MILLIS);
}
@Override
diff --git
a/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreamingTest.java
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreamingTest.java
index 4ad37f78f4a..39992b76fe3 100644
---
a/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreamingTest.java
+++
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/PipelineTranslatorStreamingTest.java
@@ -30,6 +30,9 @@ import org.apache.beam.sdk.coders.VarIntCoder;
import org.apache.beam.sdk.io.Read;
import org.apache.beam.sdk.state.StateSpec;
import org.apache.beam.sdk.state.StateSpecs;
+import org.apache.beam.sdk.state.TimeDomain;
+import org.apache.beam.sdk.state.TimerSpec;
+import org.apache.beam.sdk.state.TimerSpecs;
import org.apache.beam.sdk.transforms.Combine;
import org.apache.beam.sdk.transforms.Create;
import org.apache.beam.sdk.transforms.DoFn;
@@ -37,6 +40,7 @@ import org.apache.beam.sdk.transforms.GroupByKey;
import org.apache.beam.sdk.transforms.ParDo;
import org.apache.beam.sdk.transforms.Sum;
import org.apache.beam.sdk.transforms.windowing.FixedWindows;
+import org.apache.beam.sdk.transforms.windowing.Sessions;
import org.apache.beam.sdk.transforms.windowing.Window;
import org.apache.beam.sdk.values.KV;
import org.apache.beam.sdk.values.PCollection;
@@ -60,7 +64,8 @@ public class PipelineTranslatorStreamingTest implements
Serializable {
@Rule public transient TemporaryFolder temp = new TemporaryFolder();
private PCollection<KV<String, Integer>> kv(String tag) throws Exception {
- SparkStructuredStreamingPipelineOptions o =
StreamingTestUtils.streamingOptions(temp);
+ SparkStructuredStreamingPipelineOptions o =
+
StreamingTestUtils.streamingOptions(temp.newFolder(tag).getAbsolutePath());
return Pipeline.create(o)
.apply(Read.from(new TestUnboundedSource(tag, 1, 1)))
.apply(Window.into(FixedWindows.of(Duration.millis(1))))
@@ -70,7 +75,9 @@ public class PipelineTranslatorStreamingTest implements
Serializable {
private static void assertUnsupported(Pipeline pipeline, String expected) {
Throwable thrown = assertThrows(Exception.class, () ->
StreamingTestUtils.run(pipeline));
for (Throwable t = thrown; t != null; t = t.getCause()) {
- if (t instanceof UnsupportedOperationException &&
t.getMessage().contains(expected)) {
+ if (t instanceof UnsupportedOperationException
+ && t.getMessage().contains(expected)
+ && t.getMessage().contains("36841")) {
return;
}
}
@@ -105,8 +112,36 @@ public class PipelineTranslatorStreamingTest implements
Serializable {
}
@Test
- public void rejectsStatefulParDo() throws Exception {
- assertUnsupported(kv("s").apply(ParDo.of(new
StatefulDoFn())).getPipeline(), "Stateful ParDo");
+ public void rejectsUnsupportedStatefulFeatures() throws Exception {
+ Object[][] cases = {
+ {kv("pt").apply(ParDo.of(new ProcessingTimeTimerDoFn())).getPipeline(),
"PROCESSING_TIME"},
+ {
+ kv("spt").apply(ParDo.of(new
SyncProcessingTimeTimerDoFn())).getPipeline(),
+ "SYNCHRONIZED_PROCESSING_TIME"
+ },
+ {
+ kv("ptf").apply(ParDo.of(new
ProcessingTimeTimerFamilyDoFn())).getPipeline(),
+ "PROCESSING_TIME"
+ },
+ {
+ kv("owe").apply(ParDo.of(new OnWindowExpirationDoFn())).getPipeline(),
"@OnWindowExpiration"
+ },
+ {
+ kv("rts").apply(ParDo.of(new
RequiresTimeSortedInputDoFn())).getPipeline(),
+ "@RequiresTimeSortedInput"
+ },
+ {
+ kv("mw")
+ .apply(Window.into(Sessions.withGapDuration(Duration.millis(10))))
+ .apply(ParDo.of(new StatefulDoFn()))
+ .getPipeline(),
+ "merging windows"
+ }
+ };
+
+ for (Object[] testCase : cases) {
+ assertUnsupported((Pipeline) testCase[0], (String) testCase[1]);
+ }
}
@Test
@@ -139,4 +174,58 @@ public class PipelineTranslatorStreamingTest implements
Serializable {
@ProcessElement
public void process() {}
}
+
+ private static final class ProcessingTimeTimerDoFn extends DoFn<KV<String,
Integer>, Integer> {
+ @TimerId("pt")
+ final TimerSpec timer = TimerSpecs.timer(TimeDomain.PROCESSING_TIME);
+
+ @ProcessElement
+ public void process() {}
+
+ @OnTimer("pt")
+ public void onTimer() {}
+ }
+
+ private static final class SyncProcessingTimeTimerDoFn
+ extends DoFn<KV<String, Integer>, Integer> {
+ @TimerId("spt")
+ final TimerSpec timer =
TimerSpecs.timer(TimeDomain.SYNCHRONIZED_PROCESSING_TIME);
+
+ @ProcessElement
+ public void process() {}
+
+ @OnTimer("spt")
+ public void onTimer() {}
+ }
+
+ private static final class ProcessingTimeTimerFamilyDoFn
+ extends DoFn<KV<String, Integer>, Integer> {
+ @TimerFamily("ptf")
+ @SuppressWarnings("unused") // Reflected by DoFnSignatures
+ final TimerSpec timer = TimerSpecs.timerMap(TimeDomain.PROCESSING_TIME);
+
+ @ProcessElement
+ public void process() {}
+
+ @OnTimerFamily("ptf")
+ public void onTimer() {}
+ }
+
+ private static final class OnWindowExpirationDoFn extends DoFn<KV<String,
Integer>, Integer> {
+ @DoFn.StateId("state")
+ final StateSpec<?> spec = StateSpecs.value(VarIntCoder.of());
+
+ @ProcessElement
+ public void process() {}
+
+ @OnWindowExpiration
+ public void onWindowExpiration() {}
+ }
+
+ private static final class RequiresTimeSortedInputDoFn
+ extends DoFn<KV<String, Integer>, Integer> {
+ @RequiresTimeSortedInput
+ @ProcessElement
+ public void process() {}
+ }
}
diff --git
a/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/StatefulParDoStreamingTest.java
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/StatefulParDoStreamingTest.java
new file mode 100644
index 00000000000..26e900eba8b
--- /dev/null
+++
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/translation/streaming/StatefulParDoStreamingTest.java
@@ -0,0 +1,242 @@
+/*
+ * 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.spark.structuredstreaming.translation.streaming;
+
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertFalse;
+import static org.junit.Assert.assertTrue;
+
+import java.io.Serializable;
+import java.util.Arrays;
+import java.util.Set;
+import org.apache.beam.runners.spark.StreamingTest;
+import org.apache.beam.runners.spark.structuredstreaming.SparkSessionRule;
+import
org.apache.beam.runners.spark.structuredstreaming.SparkStructuredStreamingPipelineOptions;
+import
org.apache.beam.runners.spark.structuredstreaming.io.streaming.BeamReaderCache;
+import
org.apache.beam.runners.spark.structuredstreaming.io.streaming.TestUnboundedSource;
+import org.apache.beam.sdk.Pipeline;
+import org.apache.beam.sdk.PipelineResult;
+import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.coders.StringUtf8Coder;
+import org.apache.beam.sdk.io.Read;
+import org.apache.beam.sdk.state.StateSpec;
+import org.apache.beam.sdk.state.StateSpecs;
+import org.apache.beam.sdk.state.TimeDomain;
+import org.apache.beam.sdk.state.Timer;
+import org.apache.beam.sdk.state.TimerSpec;
+import org.apache.beam.sdk.state.TimerSpecs;
+import org.apache.beam.sdk.state.ValueState;
+import org.apache.beam.sdk.transforms.DoFn;
+import org.apache.beam.sdk.transforms.ParDo;
+import org.apache.beam.sdk.values.KV;
+import org.joda.time.Duration;
+import org.joda.time.Instant;
+import org.junit.After;
+import org.junit.ClassRule;
+import org.junit.Rule;
+import org.junit.Test;
+import org.junit.experimental.categories.Category;
+import org.junit.rules.TemporaryFolder;
+import org.junit.runner.RunWith;
+import org.junit.runners.JUnit4;
+
+/** End to end tests for stateful ParDo on Spark 4 Structured Streaming. */
+@RunWith(JUnit4.class)
+@Category(StreamingTest.class)
+public class StatefulParDoStreamingTest implements Serializable {
+
+ @ClassRule public static final SparkSessionRule SESSION = new
SparkSessionRule();
+
+ @Rule public transient TemporaryFolder tempFolder = new TemporaryFolder();
+
+ private static final String TAG_BATCHES = "stateful-batches";
+ private static final String TAG_RESTART = "stateful-restart";
+
+ @After
+ public void tearDown() {
+ TestUnboundedSource.forget(TAG_BATCHES);
+ TestUnboundedSource.forget(TAG_RESTART);
+ }
+
+ private static Pipeline pipeline(
+ SparkStructuredStreamingPipelineOptions options,
+ String tag,
+ int count,
+ DoFn<KV<String, String>, String> fn,
+ String collectorId) {
+ options.setMaxRecordsPerBatch(1L);
+ Pipeline p = Pipeline.create(options);
+ p.apply("Read", Read.from(new TestUnboundedSource(tag, 1, count)))
+ .apply("KeySelector", ParDo.of(new KeySelectorFn()))
+ .setCoder(KvCoder.of(StringUtf8Coder.of(), StringUtf8Coder.of()))
+ .apply("StatefulParDo", ParDo.of(fn))
+ .apply("Collect", ParDo.of(new
StreamingTestUtils.CollectDoFn<>(collectorId)));
+ return p;
+ }
+
+ @Test
+ public void testStateAndTimersAcrossBatches() throws Exception {
+ String collectorId = StreamingTestUtils.newCollectorId(TAG_BATCHES);
+ Pipeline pipeline =
+ pipeline(
+ StreamingTestUtils.streamingOptions(tempFolder),
+ TAG_BATCHES,
+ 10,
+ new StateAndTimerTestFn(),
+ collectorId);
+ assertEquals(PipelineResult.State.DONE,
StreamingTestUtils.run(pipeline).getState());
+
+ Set<String> collected = StreamingTestUtils.collected(collectorId);
+ assertTrue(
+ collected.containsAll(
+ Arrays.asList("k1-seen-1", "k1-seen-2", "k2-seen-1", "k1-timer",
"k2-early-1")));
+ assertFalse(collected.contains("k2-early-2"));
+ assertFalse(collected.contains("k2-late"));
+ }
+
+ @Test
+ public void testCheckpointRestartPreservesStateAndTimers() throws Exception {
+ String checkpoint = tempFolder.newFolder("cp").getAbsolutePath();
+ String c1 = StreamingTestUtils.newCollectorId("r1");
+ String c2 = StreamingTestUtils.newCollectorId("r2");
+
+ Pipeline p1 =
+ pipeline(
+ StreamingTestUtils.streamingOptions(checkpoint),
+ TAG_RESTART,
+ 3,
+ new RestartTestFn(),
+ c1);
+ assertEquals(PipelineResult.State.DONE,
StreamingTestUtils.run(p1).getState());
+ BeamReaderCache.invalidateAll();
+
+ TestUnboundedSource.extend(TAG_RESTART, 10);
+
+ Pipeline p2 =
+ pipeline(
+ StreamingTestUtils.streamingOptions(checkpoint),
+ TAG_RESTART,
+ 10,
+ new RestartTestFn(),
+ c2);
+ assertEquals(PipelineResult.State.DONE,
StreamingTestUtils.run(p2).getState());
+
+ Set<String> collected = StreamingTestUtils.collected(c2);
+ assertTrue(collected.contains("k1-resumed-saved-val") &&
collected.contains("k1-timer"));
+ assertFalse(collected.contains("k1-started"));
+ }
+
+ private static final class KeySelectorFn extends DoFn<String, KV<String,
String>> {
+ @ProcessElement
+ public void process(@Element String element, OutputReceiver<KV<String,
String>> out) {
+ int i = TestUnboundedSource.indexOf(element);
+ out.output(KV.of((i == 0 || i == 4) ? "k1" : (i == 1 || i == 3) ? "k2" :
"k_other", element));
+ }
+ }
+
+ private static final class StateAndTimerTestFn extends DoFn<KV<String,
String>, String> {
+ @StateId("count")
+ private final StateSpec<ValueState<Integer>> countSpec =
StateSpecs.value();
+
+ @StateId("earlyFired")
+ private final StateSpec<ValueState<Integer>> earlyFiredSpec =
StateSpecs.value();
+
+ @TimerId("timer1")
+ private final TimerSpec timer1Spec =
TimerSpecs.timer(TimeDomain.EVENT_TIME);
+
+ @TimerId("timerEarly")
+ private final TimerSpec timerEarlySpec =
TimerSpecs.timer(TimeDomain.EVENT_TIME);
+
+ @TimerId("timerLate")
+ private final TimerSpec timerLateSpec =
TimerSpecs.timer(TimeDomain.EVENT_TIME);
+
+ @ProcessElement
+ public void process(
+ @Element KV<String, String> element,
+ @Timestamp Instant ts,
+ @StateId("count") ValueState<Integer> countState,
+ @TimerId("timer1") Timer timer1,
+ @TimerId("timerEarly") Timer timerEarly,
+ @TimerId("timerLate") Timer timerLate,
+ OutputReceiver<String> out) {
+ String key = element.getKey();
+ int count = (countState.read() == null ? 0 : countState.read()) + 1;
+ countState.write(count);
+ out.output(key + "-seen-" + count);
+
+ if ("k1".equals(key) && count == 1) {
+ timer1.set(ts.plus(Duration.millis(3000)));
+ } else if ("k2".equals(key) && count == 1) {
+ timerEarly.set(ts.plus(Duration.millis(1000)));
+ timerLate.set(ts.plus(Duration.millis(1500)));
+ }
+ }
+
+ @OnTimer("timer1")
+ public void onTimer1(OutputReceiver<String> out) {
+ out.output("k1-timer");
+ }
+
+ @OnTimer("timerEarly")
+ public void onTimerEarly(
+ @TimerId("timerLate") Timer timerLate,
+ @StateId("earlyFired") ValueState<Integer> earlyFiredState,
+ OutputReceiver<String> out) {
+ int count = (earlyFiredState.read() == null ? 0 :
earlyFiredState.read()) + 1;
+ earlyFiredState.write(count);
+ out.output("k2-early-" + count);
+ timerLate.clear();
+ }
+
+ @OnTimer("timerLate")
+ public void onTimerLate(OutputReceiver<String> out) {
+ out.output("k2-late");
+ }
+ }
+
+ private static final class RestartTestFn extends DoFn<KV<String, String>,
String> {
+ @StateId("stored")
+ private final StateSpec<ValueState<String>> storedSpec =
StateSpecs.value();
+
+ @TimerId("timer")
+ private final TimerSpec timerSpec =
TimerSpecs.timer(TimeDomain.EVENT_TIME);
+
+ @ProcessElement
+ public void process(
+ @Element KV<String, String> element,
+ @Timestamp Instant ts,
+ @StateId("stored") ValueState<String> storedState,
+ @TimerId("timer") Timer timer,
+ OutputReceiver<String> out) {
+ String key = element.getKey();
+ String stored = storedState.read();
+ if (stored == null) {
+ storedState.write("saved-val");
+ timer.set(ts.plus(Duration.millis(4000)));
+ out.output(key + "-started");
+ } else {
+ out.output(key + "-resumed-" + stored);
+ }
+ }
+
+ @OnTimer("timer")
+ public void onTimer(OutputReceiver<String> out) {
+ out.output("k1-timer");
+ }
+ }
+}
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkStateInternals.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkStateInternals.java
index 4f744ab3ab1..8e1d3411a98 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkStateInternals.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkStateInternals.java
@@ -30,6 +30,7 @@ import org.apache.beam.runners.core.StateInternals;
import org.apache.beam.runners.core.StateNamespace;
import org.apache.beam.runners.core.StateTag;
import org.apache.beam.runners.spark.coders.CoderHelpers;
+import org.apache.beam.sdk.annotations.Internal;
import org.apache.beam.sdk.coders.Coder;
import org.apache.beam.sdk.coders.InstantCoder;
import org.apache.beam.sdk.coders.ListCoder;
@@ -66,17 +67,15 @@ import org.joda.time.Instant;
public class SparkStateInternals<K> implements StateInternals {
private final K key;
- // Serializable state for internals (namespace to state tag to coded value).
- private final Table<String, String, byte[]> stateTable;
+ private final StateCells cells;
private SparkStateInternals(K key) {
- this.key = key;
- this.stateTable = HashBasedTable.create();
+ this(key, new TableCells(HashBasedTable.create()));
}
- private SparkStateInternals(K key, Table<String, String, byte[]> stateTable)
{
+ private SparkStateInternals(K key, StateCells cells) {
this.key = key;
- this.stateTable = stateTable;
+ this.cells = cells;
}
public static <K> SparkStateInternals<K> forKey(K key) {
@@ -85,11 +84,21 @@ public class SparkStateInternals<K> implements
StateInternals {
public static <K> SparkStateInternals<K> forKeyAndState(
K key, Table<String, String, byte[]> stateTable) {
- return new SparkStateInternals<>(key, stateTable);
+ return new SparkStateInternals<>(key, new TableCells(stateTable));
+ }
+
+ /** Creates state internals for key backed by cells. */
+ @Internal
+ public static <K> SparkStateInternals<K> forKey(K key, StateCells cells) {
+ return new SparkStateInternals<>(key, cells);
}
public Table<String, String, byte[]> getState() {
- return stateTable;
+ if (cells instanceof TableCells) {
+ return ((TableCells) cells).getTable();
+ }
+ throw new IllegalStateException(
+ "getState() is not supported on non-table-backed
SparkStateInternals.");
}
@Override
@@ -192,7 +201,7 @@ public class SparkStateInternals<K> implements
StateInternals {
}
T readValue() {
- byte[] buf = stateTable.get(namespace.stringKey(), id);
+ byte[] buf = cells.get(namespace.stringKey(), id);
if (buf != null) {
return CoderHelpers.fromByteArray(buf, coder);
}
@@ -200,11 +209,11 @@ public class SparkStateInternals<K> implements
StateInternals {
}
void writeValue(T input) {
- stateTable.put(namespace.stringKey(), id,
CoderHelpers.toByteArray(input, coder));
+ cells.put(namespace.stringKey(), id, CoderHelpers.toByteArray(input,
coder));
}
public void clear() {
- stateTable.remove(namespace.stringKey(), id);
+ cells.remove(namespace.stringKey(), id);
}
@Override
@@ -289,7 +298,7 @@ public class SparkStateInternals<K> implements
StateInternals {
@Override
public Boolean read() {
- return stateTable.get(namespace.stringKey(), id) == null;
+ return cells.get(namespace.stringKey(), id) == null;
}
};
}
@@ -350,7 +359,7 @@ public class SparkStateInternals<K> implements
StateInternals {
@Override
public Boolean read() {
- return stateTable.get(namespace.stringKey(), id) == null;
+ return cells.get(namespace.stringKey(), id) == null;
}
};
}
@@ -502,7 +511,7 @@ public class SparkStateInternals<K> implements
StateInternals {
return new ReadableState<Boolean>() {
@Override
public Boolean read() {
- return stateTable.get(namespace.stringKey(), id) == null;
+ return cells.get(namespace.stringKey(), id) == null;
}
@Override
@@ -565,7 +574,7 @@ public class SparkStateInternals<K> implements
StateInternals {
return new ReadableState<Boolean>() {
@Override
public Boolean read() {
- return stateTable.get(namespace.stringKey(), id) == null;
+ return cells.get(namespace.stringKey(), id) == null;
}
@Override
@@ -625,9 +634,46 @@ public class SparkStateInternals<K> implements
StateInternals {
@Override
public Boolean read() {
- return stateTable.get(namespace.stringKey(), id) == null;
+ return cells.get(namespace.stringKey(), id) == null;
}
};
}
}
+
+ private static class TableCells implements StateCells {
+ private final Table<String, String, byte[]> stateTable;
+
+ TableCells(Table<String, String, byte[]> stateTable) {
+ this.stateTable = stateTable;
+ }
+
+ Table<String, String, byte[]> getTable() {
+ return stateTable;
+ }
+
+ @Override
+ public byte @Nullable [] get(String namespace, String stateId) {
+ return stateTable.get(namespace, stateId);
+ }
+
+ @Override
+ public void put(String namespace, String stateId, byte[] value) {
+ stateTable.put(namespace, stateId, value);
+ }
+
+ @Override
+ public void remove(String namespace, String stateId) {
+ stateTable.remove(namespace, stateId);
+ }
+ }
+
+ /** Key-value cells abstraction addressed by namespace and state id. */
+ @Internal
+ public interface StateCells {
+ byte @Nullable [] get(String namespace, String stateId);
+
+ void put(String namespace, String stateId, byte[] value);
+
+ void remove(String namespace, String stateId);
+ }
}
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkTimerInternals.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkTimerInternals.java
index ca7ebce2f19..4b904218b5a 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkTimerInternals.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkTimerInternals.java
@@ -54,6 +54,11 @@ public class SparkTimerInternals implements TimerInternals {
this.synchronizedProcessingTime = synchronizedProcessingTime;
}
+ /** Build a {@link TimerInternals} initialized with a given watermark. */
+ public static SparkTimerInternals forWatermark(Instant watermark) {
+ return new SparkTimerInternals(watermark, watermark, new Instant(0));
+ }
+
/** Build the {@link TimerInternals} according to the feeding streams. */
public static SparkTimerInternals forStreamFromSources(
List<Integer> sourceIds, Map<Integer, SparkWatermarks> watermarks) {
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/DoFnRunnerFactory.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/DoFnRunnerFactory.java
index ce4155ee8e1..988a7e6027d 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/DoFnRunnerFactory.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/DoFnRunnerFactory.java
@@ -29,6 +29,7 @@ import org.apache.beam.runners.core.StepContext;
import
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.functions.CachedSideInputReader;
import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.functions.NoOpStepContext;
+import org.apache.beam.sdk.annotations.Internal;
import org.apache.beam.sdk.coders.Coder;
import org.apache.beam.sdk.options.PipelineOptions;
import org.apache.beam.sdk.runners.AppliedPTransform;
@@ -57,9 +58,11 @@ import org.joda.time.Instant;
* Factory to create a {@link DoFnRunner}. The factory supports fusing
multiple {@link DoFnRunner
* runners} into a single one.
*/
-abstract class DoFnRunnerFactory<InT, T> implements Serializable {
+@Internal
+public abstract class DoFnRunnerFactory<InT, T> implements Serializable {
- interface DoFnRunnerWithTeardown<InT, T> extends DoFnRunner<InT, T> {
+ @Internal
+ public interface DoFnRunnerWithTeardown<InT, T> extends DoFnRunner<InT, T> {
void teardown();
}
@@ -77,7 +80,8 @@ abstract class DoFnRunnerFactory<InT, T> implements
Serializable {
*
* <p>Only supported for a single, unfused {@link DoFn}: a fused runner
cannot drive timers.
*/
- DoFnRunnerWithTeardown<InT, T> create(
+ @Internal
+ public DoFnRunnerWithTeardown<InT, T> create(
PipelineOptions options,
MetricsAccumulator metrics,
WindowedValueMultiReceiver output,
@@ -92,7 +96,8 @@ abstract class DoFnRunnerFactory<InT, T> implements
Serializable {
*/
abstract <T2> DoFnRunnerFactory<InT, T2> fuse(DoFnRunnerFactory<T, T2> next);
- static <InT, T> DoFnRunnerFactory<InT, T> simple(
+ @Internal
+ public static <InT, T> DoFnRunnerFactory<InT, T> simple(
AppliedPTransform<PCollection<? extends InT>, ?, ParDo.MultiOutput<InT,
T>> appliedPT,
PCollection<InT> input,
SideInputReader sideInputReader,
@@ -147,7 +152,7 @@ abstract class DoFnRunnerFactory<InT, T> implements
Serializable {
}
@Override
- DoFnRunnerWithTeardown<InT, T> create(
+ public DoFnRunnerWithTeardown<InT, T> create(
PipelineOptions options,
MetricsAccumulator metrics,
WindowedValueMultiReceiver output,
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulDoFnGroupFunction.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulDoFnGroupFunction.java
index b5c2d602407..4196bda0e82 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulDoFnGroupFunction.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulDoFnGroupFunction.java
@@ -25,12 +25,7 @@ import java.util.Iterator;
import java.util.Map;
import java.util.function.Supplier;
import javax.annotation.CheckForNull;
-import org.apache.beam.runners.core.InMemoryStateInternals;
import org.apache.beam.runners.core.InMemoryTimerInternals;
-import org.apache.beam.runners.core.StateInternals;
-import org.apache.beam.runners.core.StateNamespaces;
-import org.apache.beam.runners.core.StepContext;
-import org.apache.beam.runners.core.TimerInternals;
import org.apache.beam.runners.core.TimerInternals.TimerData;
import
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.DoFnRunnerFactory.DoFnRunnerWithTeardown;
@@ -38,14 +33,11 @@ import org.apache.beam.sdk.options.PipelineOptions;
import org.apache.beam.sdk.transforms.DoFn;
import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
import org.apache.beam.sdk.util.WindowedValueMultiReceiver;
-import org.apache.beam.sdk.values.CausedByDrain;
import org.apache.beam.sdk.values.KV;
import org.apache.beam.sdk.values.TupleTag;
import org.apache.beam.sdk.values.WindowedValue;
import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.AbstractIterator;
-import org.apache.spark.TaskContext;
import org.apache.spark.api.java.function.FlatMapGroupsFunction;
-import org.apache.spark.util.TaskCompletionListener;
import org.checkerframework.checker.nullness.qual.Nullable;
import scala.Tuple2;
@@ -76,12 +68,10 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
private final Supplier<PipelineOptions> options;
private final MetricsAccumulator metrics;
private final DoFnRunnerFactory<InT, ?> factory;
+ private final StatefulTaskRunner<InT, ?> taskRunner = new
StatefulTaskRunner<>();
private transient @Nullable Deque<OutT> buffer;
- private transient @Nullable MutableStepContext stepContext;
- private transient @Nullable DoFnRunnerWithTeardown<InT, ?> doFnRunner;
private transient boolean needsBundleStart;
- private transient boolean isTornDown;
private StatefulDoFnGroupFunction(
Supplier<PipelineOptions> options,
@@ -122,7 +112,7 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
public Iterator<OutT> call(K key, Iterator<WindowedValue<InT>> values) {
DoFnRunnerWithTeardown<InT, ?> runner = runner();
// Fresh state and timers for this key; the DoFn instance itself is
untouched.
- stepContext().reset(key);
+ taskRunner.stepContext().reset(key);
if (needsBundleStart) {
needsBundleStart = false;
runner.startBundle();
@@ -135,39 +125,12 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
* and opens the first bundle, so this happens exactly once per task rather
than once per key.
*/
private DoFnRunnerWithTeardown<InT, ?> runner() {
- DoFnRunnerWithTeardown<InT, ?> runner = doFnRunner;
- if (runner == null) {
- MutableStepContext ctx = new MutableStepContext();
- Deque<OutT> buf = new ArrayDeque<>();
- buffer = buf;
- stepContext = ctx;
- runner = factory.create(options.get(), metrics, outputManager(buf), ctx);
- doFnRunner = runner;
- // Spark is free to abandon an iterator part way through (a downstream
limit, a task kill, an
- // exception elsewhere in the stage). Tearing down from the task
completion listener is the
- // only way to guarantee @Teardown runs and DoFn resources are released.
- TaskContext taskContext = TaskContext.get();
- if (taskContext != null) {
- // An explicit listener rather than a lambda: TaskContext overloads
this for both the Scala
- // function and the Java interface, so a lambda is ambiguous.
- taskContext.addTaskCompletionListener(
- new TaskCompletionListener() {
- @Override
- public void onTaskCompletion(TaskContext context) {
- teardownOnce();
- }
- });
- }
- }
- return runner;
- }
-
- private MutableStepContext stepContext() {
- MutableStepContext ctx = stepContext;
- if (ctx == null) {
- throw new IllegalStateException("StepContext requested before the runner
was created");
- }
- return ctx;
+ return taskRunner.getOrCreate(
+ ctx -> {
+ Deque<OutT> buf = new ArrayDeque<>();
+ buffer = buf;
+ return factory.create(options.get(), metrics, outputManager(buf),
ctx);
+ });
}
private Deque<OutT> buffer() {
@@ -178,14 +141,6 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
return buf;
}
- private void teardownOnce() {
- DoFnRunnerWithTeardown<InT, ?> runner = doFnRunner;
- if (runner != null && !isTornDown) {
- isTornDown = true;
- runner.teardown();
- }
- }
-
/** Output manager emitting outputs of type {@link OutT} to the buffer. */
abstract WindowedValueMultiReceiver outputManager(Deque<OutT> buffer);
@@ -246,45 +201,6 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
}
}
- /**
- * A {@link StepContext} whose state and timers are swapped per key, so that
one {@link DoFn} and
- * one {@link org.apache.beam.runners.core.DoFnRunner DoFnRunner} can serve
every key of a task.
- *
- * <p>{@code SimpleDoFnRunner} re-reads {@code stateInternals()} on each
access rather than
- * caching it, which is what makes rebinding safe.
- */
- private static class MutableStepContext implements StepContext {
- private @Nullable StateInternals stateInternals;
- private @Nullable InMemoryTimerInternals timerInternals;
-
- void reset(@Nullable Object key) {
- stateInternals = InMemoryStateInternals.forKey(key);
- timerInternals = new InMemoryTimerInternals();
- }
-
- InMemoryTimerInternals timers() {
- InMemoryTimerInternals timers = timerInternals;
- if (timers == null) {
- throw new IllegalStateException("StepContext used before reset");
- }
- return timers;
- }
-
- @Override
- public StateInternals stateInternals() {
- StateInternals state = stateInternals;
- if (state == null) {
- throw new IllegalStateException("StepContext used before reset");
- }
- return state;
- }
-
- @Override
- public TimerInternals timerInternals() {
- return timers();
- }
- }
-
private class StatefulGroupIt extends AbstractIterator<OutT> {
private final Iterator<WindowedValue<InT>> groupIt;
private final K key;
@@ -300,7 +216,7 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
this.key = key;
this.groupIt = groupIt;
this.runner = runner;
- this.timerInternals = stepContext().timers();
+ this.timerInternals = taskRunner.stepContext().timers();
}
@Override
@@ -326,10 +242,10 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
}
}
} catch (RuntimeException re) {
- teardownOnce();
+ taskRunner.teardownOnce();
throw re;
} catch (Exception e) {
- teardownOnce();
+ taskRunner.teardownOnce();
throw new RuntimeException(e);
}
}
@@ -375,17 +291,7 @@ abstract class StatefulDoFnGroupFunction<K, InT extends
KV<K, ?>, OutT>
}
private void fire(TimerData timer) {
- BoundedWindow window =
- ((StateNamespaces.WindowNamespace<?>)
timer.getNamespace()).getWindow();
- runner.onTimer(
- timer.getTimerId(),
- timer.getTimerFamilyId(),
- key,
- window,
- timer.getTimestamp(),
- timer.getOutputTimestamp(),
- timer.getDomain(),
- CausedByDrain.NORMAL);
+ StatefulTaskRunner.fireTimer(runner, key, timer);
}
}
}
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulTaskRunner.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulTaskRunner.java
new file mode 100644
index 00000000000..09bbe54e543
--- /dev/null
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/batch/StatefulTaskRunner.java
@@ -0,0 +1,145 @@
+/*
+ * 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.spark.structuredstreaming.translation.batch;
+
+import java.io.Serializable;
+import java.util.function.Function;
+import org.apache.beam.runners.core.DoFnRunner;
+import org.apache.beam.runners.core.InMemoryStateInternals;
+import org.apache.beam.runners.core.InMemoryTimerInternals;
+import org.apache.beam.runners.core.StateInternals;
+import org.apache.beam.runners.core.StateNamespaces;
+import org.apache.beam.runners.core.StepContext;
+import org.apache.beam.runners.core.TimerInternals;
+import org.apache.beam.runners.core.TimerInternals.TimerData;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.batch.DoFnRunnerFactory.DoFnRunnerWithTeardown;
+import org.apache.beam.sdk.annotations.Internal;
+import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
+import org.apache.beam.sdk.values.CausedByDrain;
+import org.apache.spark.TaskContext;
+import org.apache.spark.util.TaskCompletionListener;
+import org.checkerframework.checker.nullness.qual.Nullable;
+
+/**
+ * Manages the task scoped lifecycle of a stateful {@link DoFnRunner} and its
{@link StepContext}.
+ */
+@Internal
+public class StatefulTaskRunner<InT, T> implements Serializable {
+
+ private transient @Nullable DoFnRunnerWithTeardown<InT, ?> runner;
+ private transient @Nullable MutableStepContext stepContext;
+ private transient boolean isTornDown;
+
+ public <R extends DoFnRunnerWithTeardown<InT, ?>> R getOrCreate(
+ Function<MutableStepContext, R> createRunner) {
+ DoFnRunnerWithTeardown<InT, ?> r = runner;
+ if (r == null) {
+ MutableStepContext ctx = new MutableStepContext();
+ this.stepContext = ctx;
+ R created = createRunner.apply(ctx);
+ this.runner = created;
+ TaskContext taskContext = TaskContext.get();
+ if (taskContext != null) {
+ taskContext.addTaskCompletionListener(
+ new TaskCompletionListener() {
+ @Override
+ public void onTaskCompletion(TaskContext context) {
+ teardownOnce();
+ }
+ });
+ }
+ return created;
+ }
+ @SuppressWarnings("unchecked")
+ R existing = (R) r;
+ return existing;
+ }
+
+ public MutableStepContext stepContext() {
+ MutableStepContext ctx = stepContext;
+ if (ctx == null) {
+ throw new IllegalStateException("StepContext requested before the runner
was created");
+ }
+ return ctx;
+ }
+
+ public void teardownOnce() {
+ DoFnRunnerWithTeardown<InT, ?> r = runner;
+ if (r != null && !isTornDown) {
+ isTornDown = true;
+ r.teardown();
+ }
+ }
+
+ public static void fireTimer(DoFnRunner<?, ?> runner, @Nullable Object key,
TimerData timer) {
+ BoundedWindow window = ((StateNamespaces.WindowNamespace<?>)
timer.getNamespace()).getWindow();
+ runner.onTimer(
+ timer.getTimerId(),
+ timer.getTimerFamilyId(),
+ key,
+ window,
+ timer.getTimestamp(),
+ timer.getOutputTimestamp(),
+ timer.getDomain(),
+ CausedByDrain.NORMAL);
+ }
+
+ /** A mutable step context that can rebind state and timers per key. */
+ @Internal
+ public static class MutableStepContext implements StepContext {
+ private @Nullable StateInternals stateInternals;
+ private @Nullable TimerInternals timerInternals;
+
+ public void reset(@Nullable Object key) {
+ this.stateInternals = InMemoryStateInternals.forKey(key);
+ this.timerInternals = new InMemoryTimerInternals();
+ }
+
+ public void set(StateInternals stateInternals, TimerInternals
timerInternals) {
+ this.stateInternals = stateInternals;
+ this.timerInternals = timerInternals;
+ }
+
+ public InMemoryTimerInternals timers() {
+ TimerInternals timers = timerInternals();
+ if (timers instanceof InMemoryTimerInternals) {
+ return (InMemoryTimerInternals) timers;
+ }
+ throw new IllegalStateException(
+ "Expected InMemoryTimerInternals but was " +
timers.getClass().getName());
+ }
+
+ @Override
+ public StateInternals stateInternals() {
+ StateInternals state = stateInternals;
+ if (state == null) {
+ throw new IllegalStateException("StepContext used before reset");
+ }
+ return state;
+ }
+
+ @Override
+ public TimerInternals timerInternals() {
+ TimerInternals timers = timerInternals;
+ if (timers == null) {
+ throw new IllegalStateException("StepContext used before reset");
+ }
+ return timers;
+ }
+ }
+}