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;
+    }
+  }
+}

Reply via email to