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 abe41005e7b [Spark][#36841] Add the DataSourceV2 unbounded source for 
the Spark 4 streaming runner (#39971)
abe41005e7b is described below

commit abe41005e7bca3a41d8e42e0d5511902dadac124
Author: Tobias Kaymak <[email protected]>
AuthorDate: Tue Sep 8 21:37:49 2026 +0200

    [Spark][#36841] Add the DataSourceV2 unbounded source for the Spark 4 
streaming runner (#39971)
    
    * [Spark 4] Add DataSourceV2 micro-batch source for Beam unbounded sources
    
    Exposes any Beam UnboundedSource as a Spark 4 DataSourceV2 streaming
    table with a fixed two column schema, encoded payload plus event
    timestamp. Offsets are opaque, strictly increasing epoch counters, so
    Spark keeps scheduling micro-batches and termination stays with the
    lifecycle owner.
    
    Recovery is durable under the query's checkpoint location: the source id
    derives deterministically from the read transform's full name, the first
    run pins its split list (Beam sources do not guarantee deterministic
    splitting), and every split persists its CheckpointMark per epoch with a
    retention of two, written atomically via temp file and rename. Executors
    cache live readers between micro-batches and fall back to the newest
    durable mark at or before the replayed epoch after a restart. Semantics
    are at least once, a crash between finishing a read and Spark's commit
    replays the last micro-batch.
    
    The batch cutoff honors maxRecordsPerBatch, values below 1, including
    the default, mean no limit and the batch ends on the duration deadline.
---
 .../io/streaming/BeamReaderCache.java              |  352 +++++++
 .../io/streaming/BeamSourceCheckpoint.java         |  227 +++++
 .../io/streaming/UnboundedSourceDataset.java       |  721 ++++++++++++++
 .../io/streaming/BeamMicroBatchSourceTest.java     | 1028 ++++++++++++++++++++
 .../SparkStructuredStreamingPipelineOptions.java   |    9 +
 5 files changed, 2337 insertions(+)

diff --git 
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamReaderCache.java
 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamReaderCache.java
new file mode 100644
index 00000000000..1eb31e766e7
--- /dev/null
+++ 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamReaderCache.java
@@ -0,0 +1,352 @@
+/*
+ * 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.io.streaming;
+
+import java.io.Closeable;
+import java.io.IOException;
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
+import java.util.concurrent.Executors;
+import java.util.concurrent.ScheduledExecutorService;
+import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicBoolean;
+import java.util.function.LongSupplier;
+import 
org.apache.beam.runners.spark.structuredstreaming.translation.helpers.CoderHelpers;
+import org.apache.beam.sdk.io.UnboundedSource;
+import org.apache.beam.sdk.io.UnboundedSource.CheckpointMark;
+import org.apache.beam.sdk.io.UnboundedSource.UnboundedReader;
+import org.apache.beam.sdk.options.PipelineOptions;
+import org.checkerframework.checker.nullness.qual.Nullable;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+/**
+ * Executor side cache of live Beam {@link UnboundedReader}s keyed by 
checkpoint location and split.
+ *
+ * <p>An entry records the epoch its reader is positioned at and the mark 
taken there. A batch
+ * starting at that epoch reuses the reader and finalizes the pending mark, 
the start epoch of a
+ * batch is always committed by Spark. Any other start epoch, or a reader that 
moved without
+ * completing its batch, closes the entry without finalizing and restores the 
reader from the
+ * durable mark at the start epoch.
+ *
+ * <p>A sweeper thread closes readers idle for longer than their idle timeout, 
finalizing marks
+ * whose epoch Spark committed, see {@link 
BeamSourceCheckpoint#readSparkCommittedEpoch()}, and
+ * dropping the others, the source redelivers. The timeout must exceed the 
longest gap between two
+ * micro-batches of one split. Under {@code 
spark.sql.streaming.asyncProgressTrackingEnabled} the
+ * commit log lags, idle readers then drop their marks. Speculative execution 
can leave a losing
+ * attempt's mark finalized on another executor, this source is not safe under 
{@code
+ * spark.speculation} with sources whose reads are not deterministic.
+ */
+public final class BeamReaderCache {
+
+  private static final Logger LOG = 
LoggerFactory.getLogger(BeamReaderCache.class);
+
+  private static final ConcurrentMap<String, CachedReader<?>> READERS = new 
ConcurrentHashMap<>();
+
+  /** One monitor per key, acquire serializes per split, not across splits. */
+  private static final ConcurrentMap<String, Object> LOCKS = new 
ConcurrentHashMap<>();
+
+  private static final long SWEEP_INTERVAL_MILLIS = 10_000L;
+
+  private static final AtomicBoolean SWEEPER_STARTED = new AtomicBoolean();
+
+  private BeamReaderCache() {}
+
+  public static String key(String checkpointLocation, int splitId) {
+    return checkpointLocation + '|' + splitId;
+  }
+
+  /** Supplies the durable coded mark at the start epoch of a batch, null if 
there is none. */
+  @FunctionalInterface
+  interface MarkRestorer {
+    byte @Nullable [] restore() throws IOException;
+  }
+
+  /**
+   * Returns the reader for {@code key} positioned at {@code startEpoch}, 
reusing the cached one if
+   * it is there, restoring from the durable mark otherwise. A zero length 
durable mark means a
+   * fresh start. {@code committedEpoch} supplies the epoch Spark last 
committed, -1 if unknown.
+   *
+   * @throws IllegalStateException if {@code startEpoch > 0} and no durable 
mark exists
+   */
+  public static <T> CachedReader<T> acquire(
+      String key,
+      long startEpoch,
+      UnboundedSource<T, ?> source,
+      PipelineOptions options,
+      long idleTimeoutMillis,
+      LongSupplier committedEpoch,
+      MarkRestorer restorer)
+      throws IOException {
+    startSweeper();
+    closeIdle();
+    synchronized (lock(key)) {
+      CachedReader<?> existing = READERS.get(key);
+      if (existing != null) {
+        if (existing.beginBatch(startEpoch)) {
+          existing.finalizePendingMark(key);
+          @SuppressWarnings("unchecked") // one source per key, its element 
type never changes
+          CachedReader<T> reused = (CachedReader<T>) existing;
+          return reused;
+        }
+        LOG.info(
+            "Cached Beam reader {} is at epoch {}, batch starts at {}, 
restoring from the durable"
+                + " mark.",
+            key,
+            existing.positionEpoch(),
+            startEpoch);
+        invalidate(key);
+      }
+      byte[] codedMark = restorer.restore();
+      if (codedMark == null && startEpoch > 0) {
+        throw new IllegalStateException(
+            "No durable checkpoint mark for Beam reader " + key + " at epoch " 
+ startEpoch);
+      }
+      if (codedMark != null && codedMark.length == 0) {
+        codedMark = null;
+      }
+      LOG.info(
+          "Creating Beam reader {} at epoch {} ({} mark).",
+          key,
+          startEpoch,
+          codedMark == null ? "no" : "restored");
+      CachedReader<T> created =
+          new CachedReader<>(
+              createReader(source, options, codedMark),
+              startEpoch,
+              codedMark,
+              idleTimeoutMillis,
+              committedEpoch);
+      created.beginBatch(startEpoch);
+      READERS.put(key, created);
+      return created;
+    }
+  }
+
+  private static <T, MarkT extends CheckpointMark> UnboundedReader<T> 
createReader(
+      UnboundedSource<T, MarkT> source, PipelineOptions options, byte 
@Nullable [] codedMark)
+      throws IOException {
+    MarkT mark =
+        codedMark == null
+            ? null
+            : CoderHelpers.fromByteArray(codedMark, 
source.getCheckpointMarkCoder());
+    return source.createReader(options, mark);
+  }
+
+  /** Closes and forgets the reader of {@code key}, nothing is finalized. */
+  public static void invalidate(String key) {
+    synchronized (lock(key)) {
+      CachedReader<?> removed = READERS.remove(key);
+      if (removed != null) {
+        close(key, removed);
+      }
+    }
+  }
+
+  /** Closes and forgets every cached reader. */
+  public static void invalidateAll() {
+    for (String key : READERS.keySet()) {
+      invalidate(key);
+    }
+  }
+
+  @SuppressWarnings("FutureReturnValueIgnored") // the sweep runs until the 
JVM exits
+  private static void startSweeper() {
+    if (!SWEEPER_STARTED.compareAndSet(false, true)) {
+      return;
+    }
+    ScheduledExecutorService sweeper =
+        Executors.newSingleThreadScheduledExecutor(
+            runnable -> {
+              Thread thread = new Thread(runnable, "beam-reader-idle-sweep");
+              thread.setDaemon(true);
+              return thread;
+            });
+    sweeper.scheduleWithFixedDelay(
+        BeamReaderCache::sweep,
+        SWEEP_INTERVAL_MILLIS,
+        SWEEP_INTERVAL_MILLIS,
+        TimeUnit.MILLISECONDS);
+  }
+
+  private static void sweep() {
+    try {
+      closeIdle(System.currentTimeMillis());
+    } catch (RuntimeException e) {
+      LOG.warn("Idle sweep of Beam readers failed.", e);
+    }
+  }
+
+  private static void closeIdle() {
+    closeIdle(System.currentTimeMillis());
+  }
+
+  /** Closes every reader idle at {@code nowMillis}, finalizing marks of 
committed epochs. */
+  static void closeIdle(long nowMillis) {
+    for (Map.Entry<String, CachedReader<?>> entry : READERS.entrySet()) {
+      String key = entry.getKey();
+      CachedReader<?> reader = entry.getValue();
+      synchronized (lock(key)) {
+        if (!reader.isIdleSince(nowMillis) || !READERS.remove(key, reader)) {
+          continue;
+        }
+        LOG.info("Closing idle Beam reader {}.", key);
+        reader.finalizeIfCommitted(key);
+        close(key, reader);
+      }
+    }
+  }
+
+  private static void close(String key, CachedReader<?> reader) {
+    try {
+      reader.close();
+    } catch (IOException | RuntimeException e) {
+      LOG.warn("Failed to close Beam reader {}.", key, e);
+    }
+  }
+
+  private static Object lock(String key) {
+    return LOCKS.computeIfAbsent(key, k -> new Object());
+  }
+
+  /** A live reader with the epoch it is positioned at and the coded mark 
taken there. */
+  public static final class CachedReader<T> implements Closeable {
+    private final UnboundedReader<T> reader;
+    private final long idleTimeoutMillis;
+    private final LongSupplier committedEpoch;
+    private boolean started;
+    private boolean inBatch;
+    private boolean moved;
+    private long positionEpoch;
+    private byte @Nullable [] positionMark;
+    private @Nullable CheckpointMark pendingMark;
+    private long lastUsedMillis;
+
+    CachedReader(
+        UnboundedReader<T> reader,
+        long positionEpoch,
+        byte @Nullable [] positionMark,
+        long idleTimeoutMillis,
+        LongSupplier committedEpoch) {
+      this.reader = reader;
+      this.positionEpoch = positionEpoch;
+      this.positionMark = positionMark;
+      this.idleTimeoutMillis = idleTimeoutMillis;
+      this.committedEpoch = committedEpoch;
+      this.lastUsedMillis = System.currentTimeMillis();
+    }
+
+    public UnboundedReader<T> reader() {
+      return reader;
+    }
+
+    public synchronized boolean startOrAdvance() throws IOException {
+      moved = true;
+      if (!started) {
+        started = true;
+        return reader.start();
+      }
+      return reader.advance();
+    }
+
+    /**
+     * Whether {@link #startOrAdvance()} was called at least once, only then 
may a mark be taken.
+     */
+    public synchronized boolean started() {
+      return started;
+    }
+
+    public synchronized long positionEpoch() {
+      return positionEpoch;
+    }
+
+    /** The coded mark of the current position, null for a fresh start. */
+    synchronized byte @Nullable [] positionMark() {
+      return positionMark;
+    }
+
+    /**
+     * Claims the reader for a batch starting at {@code epoch}, false if it 
cannot continue there.
+     */
+    synchronized boolean beginBatch(long epoch) {
+      if (positionEpoch != epoch || moved) {
+        return false;
+      }
+      inBatch = true;
+      lastUsedMillis = System.currentTimeMillis();
+      return true;
+    }
+
+    /** Records a completed batch, the reader is positioned at {@code 
endEpoch} from now on. */
+    synchronized void endBatch(long endEpoch, @Nullable CheckpointMark mark, 
byte[] codedMark) {
+      positionEpoch = endEpoch;
+      positionMark = codedMark;
+      pendingMark = mark;
+      moved = false;
+      inBatch = false;
+      lastUsedMillis = System.currentTimeMillis();
+    }
+
+    synchronized boolean isIdleSince(long nowMillis) {
+      return !inBatch && nowMillis - lastUsedMillis > idleTimeoutMillis;
+    }
+
+    /** Finalizes the pending mark if any, a failure is logged. */
+    synchronized void finalizePendingMark(String key) {
+      CheckpointMark mark = pendingMark;
+      pendingMark = null;
+      if (mark == null) {
+        return;
+      }
+      LOG.debug("Finalizing checkpoint mark of Beam reader {} at epoch {}.", 
key, positionEpoch);
+      try {
+        mark.finalizeCheckpoint();
+      } catch (IOException | RuntimeException e) {
+        LOG.warn(
+            "Failed to finalize checkpoint mark of Beam reader {} at epoch 
{}.",
+            key,
+            positionEpoch,
+            e);
+      }
+    }
+
+    /** Finalizes the pending mark if Spark committed its epoch, drops it 
otherwise. */
+    synchronized void finalizeIfCommitted(String key) {
+      if (pendingMark == null) {
+        return;
+      }
+      long committed = committedEpoch.getAsLong();
+      if (positionEpoch <= committed) {
+        finalizePendingMark(key);
+        return;
+      }
+      LOG.info(
+          "Dropping mark of Beam reader {} at epoch {}, Spark committed up to 
{}.",
+          key,
+          positionEpoch,
+          committed);
+      pendingMark = null;
+    }
+
+    @Override
+    public void close() throws IOException {
+      reader.close();
+    }
+  }
+}
diff --git 
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamSourceCheckpoint.java
 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamSourceCheckpoint.java
new file mode 100644
index 00000000000..d667976edce
--- /dev/null
+++ 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamSourceCheckpoint.java
@@ -0,0 +1,227 @@
+/*
+ * 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.io.streaming;
+
+import java.io.FileNotFoundException;
+import java.io.IOException;
+import java.nio.charset.StandardCharsets;
+import java.util.ArrayList;
+import java.util.Arrays;
+import java.util.List;
+import org.apache.beam.sdk.io.UnboundedSource;
+import org.apache.beam.sdk.util.SerializableUtils;
+import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.io.ByteStreams;
+import org.apache.hadoop.conf.Configuration;
+import org.apache.hadoop.fs.FSDataInputStream;
+import org.apache.hadoop.fs.FileStatus;
+import org.apache.hadoop.fs.Path;
+import org.apache.spark.sql.execution.streaming.CheckpointFileManager;
+import 
org.apache.spark.sql.execution.streaming.CheckpointFileManager.CancellableFSDataOutputStream;
+import org.checkerframework.checker.nullness.qual.Nullable;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+/**
+ * Durable state of one Beam unbounded source under the per source checkpoint 
location Spark hands
+ * to {@code toMicroBatchStream}.
+ *
+ * <p>{@code <location>/splits} pins the split list, written once by the 
driver. {@code
+ * <location>/marks/<epoch>/<splitId>} holds the coded checkpoint mark of a 
split at the end of the
+ * batch ending at that epoch. The epoch Spark last committed is read from 
Spark's own {@code
+ * commits} and {@code offsets} logs two levels up. All IO goes through 
Spark's {@link
+ * CheckpointFileManager}, writes are atomic renames.
+ */
+public final class BeamSourceCheckpoint {
+
+  private static final Logger LOG = 
LoggerFactory.getLogger(BeamSourceCheckpoint.class);
+
+  private static final String SPLITS_FILE = "splits";
+  private static final String MARKS_DIR = "marks";
+  private static final String SPARK_COMMITS_DIR = "commits";
+  private static final String SPARK_OFFSETS_DIR = "offsets";
+  private static final String SERIALIZED_VOID_OFFSET = "-";
+
+  /** A purge further than this above the last one lists the directory instead 
of probing epochs. */
+  private static final long MAX_BLIND_PURGE_RANGE = 1_000L;
+
+  private final String location;
+  private final CheckpointFileManager fm;
+  private final Path root;
+  private final Path splitsPath;
+  private final Path marksRoot;
+
+  /** Every mark epoch strictly below the value is known to be deleted, -1 for 
unknown. */
+  private volatile long purgeFloor = -1L;
+
+  public BeamSourceCheckpoint(String checkpointLocation, Configuration 
hadoopConf) {
+    this.location = checkpointLocation;
+    this.root = new Path(checkpointLocation);
+    this.fm = CheckpointFileManager.create(root, hadoopConf);
+    this.splitsPath = new Path(root, SPLITS_FILE);
+    this.marksRoot = new Path(root, MARKS_DIR);
+  }
+
+  public String location() {
+    return location;
+  }
+
+  /** The pinned split list, or null if none was pinned yet. */
+  public @Nullable List<UnboundedSource<?, ?>> readSplits() throws IOException 
{
+    if (!fm.exists(splitsPath)) {
+      return null;
+    }
+    @SuppressWarnings("unchecked") // written by writeSplits as an ArrayList 
of sources
+    List<UnboundedSource<?, ?>> splits =
+        (List<UnboundedSource<?, ?>>)
+            SerializableUtils.deserializeFromByteArray(read(splitsPath), 
"splits at " + splitsPath);
+    return splits;
+  }
+
+  /** Pins the split list, fails if one is pinned already. */
+  public void writeSplits(List<? extends UnboundedSource<?, ?>> splits) throws 
IOException {
+    fm.mkdirs(root);
+    if (fm.exists(splitsPath)) {
+      throw new IOException("Split list already pinned at " + splitsPath);
+    }
+    write(splitsPath, SerializableUtils.serializeToByteArray(new 
ArrayList<>(splits)), false);
+    LOG.info("Pinned {} split(s) at {}.", splits.size(), splitsPath);
+  }
+
+  /** Creates the mark directory of {@code epoch}, the driver calls this once 
per batch. */
+  public void prepareEpoch(long epoch) throws IOException {
+    fm.mkdirs(epochDir(epoch));
+  }
+
+  /**
+   * Writes the mark, creating the epoch directory if a manager without parent 
creation needs it.
+   */
+  public void writeMark(int splitId, long epoch, byte[] codedMark) throws 
IOException {
+    Path path = markPath(splitId, epoch);
+    try {
+      write(path, codedMark, true);
+    } catch (FileNotFoundException e) {
+      fm.mkdirs(epochDir(epoch));
+      write(path, codedMark, true);
+    }
+  }
+
+  /** The coded mark of a split at an epoch, or null if absent. */
+  public byte @Nullable [] readMark(int splitId, long epoch) throws 
IOException {
+    Path path = markPath(splitId, epoch);
+    if (!fm.exists(path)) {
+      return null;
+    }
+    return read(path);
+  }
+
+  /**
+   * The end epoch of this source in the last batch Spark committed, or -1 if 
there is none or the
+   * logs cannot be read. The location is {@code <root>/sources/<index>}, the 
batch id is the
+   * highest entry of {@code <root>/commits} and its epoch is line {@code 
index} after the version
+   * and metadata lines of {@code <root>/offsets/<id>}.
+   */
+  public long readSparkCommittedEpoch() {
+    try {
+      Path sparkRoot = root.getParent().getParent();
+      int sourceIndex = Integer.parseInt(root.getName());
+      Path commits = new Path(sparkRoot, SPARK_COMMITS_DIR);
+      if (!fm.exists(commits)) {
+        return -1L;
+      }
+      long batchId = -1L;
+      for (FileStatus status : fm.list(commits)) {
+        batchId = Math.max(batchId, parseEpoch(status.getPath().getName()));
+      }
+      if (batchId < 0) {
+        return -1L;
+      }
+      Path offsets = new Path(new Path(sparkRoot, SPARK_OFFSETS_DIR), 
Long.toString(batchId));
+      List<String> lines =
+          Arrays.asList(new String(read(offsets), 
StandardCharsets.UTF_8).split("\n", -1));
+      String line = lines.get(2 + sourceIndex).trim();
+      return line.equals(SERIALIZED_VOID_OFFSET) ? -1L : Long.parseLong(line);
+    } catch (IOException | RuntimeException e) {
+      LOG.warn("Failed to read the epoch Spark committed for {}.", location, 
e);
+      return -1L;
+    }
+  }
+
+  /**
+   * Deletes the marks of every epoch strictly below {@code epoch}, one 
recursive delete per epoch
+   * directory. Lists the marks directory once, later calls delete the range 
above the previous
+   * floor only. Idempotent.
+   */
+  public void purgeMarksBelow(long epoch) throws IOException {
+    long floor = purgeFloor;
+    if (floor >= 0 && epoch - floor > MAX_BLIND_PURGE_RANGE) {
+      floor = -1L;
+    }
+    if (floor < 0) {
+      if (fm.exists(marksRoot)) {
+        for (FileStatus status : fm.list(marksRoot)) {
+          long existing = parseEpoch(status.getPath().getName());
+          if (existing >= 0 && existing < epoch) {
+            fm.delete(status.getPath());
+          }
+        }
+      }
+      purgeFloor = epoch;
+      return;
+    }
+    for (long e = floor; e < epoch; e++) {
+      fm.delete(epochDir(e));
+    }
+    if (epoch > floor) {
+      purgeFloor = epoch;
+    }
+  }
+
+  private Path epochDir(long epoch) {
+    return new Path(marksRoot, Long.toString(epoch));
+  }
+
+  private Path markPath(int splitId, long epoch) {
+    return new Path(epochDir(epoch), Integer.toString(splitId));
+  }
+
+  private byte[] read(Path path) throws IOException {
+    try (FSDataInputStream in = fm.open(path)) {
+      return ByteStreams.toByteArray(in);
+    }
+  }
+
+  private void write(Path path, byte[] bytes, boolean overwrite) throws 
IOException {
+    CancellableFSDataOutputStream out = fm.createAtomic(path, overwrite);
+    try {
+      out.write(bytes);
+      out.close();
+    } catch (IOException | RuntimeException e) {
+      out.cancel();
+      throw e;
+    }
+  }
+
+  /** The epoch encoded in a mark directory name, or -1 for anything else. */
+  private static long parseEpoch(String name) {
+    try {
+      return Long.parseLong(name);
+    } catch (NumberFormatException e) {
+      return -1L;
+    }
+  }
+}
diff --git 
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/UnboundedSourceDataset.java
 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/UnboundedSourceDataset.java
new file mode 100644
index 00000000000..0368f95ba6d
--- /dev/null
+++ 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/UnboundedSourceDataset.java
@@ -0,0 +1,721 @@
+/*
+ * 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.io.streaming;
+
+import java.io.IOException;
+import java.io.UncheckedIOException;
+import java.util.ArrayList;
+import java.util.Arrays;
+import java.util.Collections;
+import java.util.List;
+import java.util.Set;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.atomic.AtomicBoolean;
+import java.util.concurrent.atomic.AtomicLong;
+import org.apache.beam.runners.core.construction.SerializablePipelineOptions;
+import 
org.apache.beam.runners.spark.structuredstreaming.SparkStructuredStreamingPipelineOptions;
+import 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.BeamReaderCache.CachedReader;
+import 
org.apache.beam.runners.spark.structuredstreaming.translation.helpers.CoderHelpers;
+import org.apache.beam.sdk.coders.Coder;
+import org.apache.beam.sdk.io.UnboundedSource;
+import org.apache.beam.sdk.io.UnboundedSource.CheckpointMark;
+import org.apache.beam.sdk.options.PipelineOptions;
+import org.apache.beam.sdk.util.BackOff;
+import org.apache.beam.sdk.util.BackOffUtils;
+import org.apache.beam.sdk.util.FluentBackoff;
+import org.apache.beam.sdk.util.Sleeper;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.beam.sdk.values.WindowedValues;
+import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableSet;
+import org.apache.hadoop.conf.Configuration;
+import org.apache.spark.TaskContext;
+import org.apache.spark.broadcast.Broadcast;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Row;
+import org.apache.spark.sql.SparkSession;
+import org.apache.spark.sql.catalyst.InternalRow;
+import org.apache.spark.sql.catalyst.expressions.GenericInternalRow;
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan;
+import org.apache.spark.sql.catalyst.streaming.StreamingRelationV2;
+import org.apache.spark.sql.catalyst.types.DataTypeUtils;
+import org.apache.spark.sql.classic.Dataset$;
+import org.apache.spark.sql.connector.catalog.SupportsRead;
+import org.apache.spark.sql.connector.catalog.Table;
+import org.apache.spark.sql.connector.catalog.TableCapability;
+import org.apache.spark.sql.connector.read.InputPartition;
+import org.apache.spark.sql.connector.read.PartitionReader;
+import org.apache.spark.sql.connector.read.PartitionReaderFactory;
+import org.apache.spark.sql.connector.read.Scan;
+import org.apache.spark.sql.connector.read.ScanBuilder;
+import org.apache.spark.sql.connector.read.streaming.MicroBatchStream;
+import org.apache.spark.sql.connector.read.streaming.Offset;
+import org.apache.spark.sql.types.DataTypes;
+import org.apache.spark.sql.types.StructType;
+import org.apache.spark.sql.util.CaseInsensitiveStringMap;
+import org.apache.spark.util.SerializableConfiguration;
+import org.checkerframework.checker.nullness.qual.Nullable;
+import org.joda.time.Duration;
+import org.joda.time.Instant;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+import scala.Option;
+import scala.reflect.ClassTag;
+
+/**
+ * Translator facing entry point turning a Beam {@link UnboundedSource} into a 
streaming Spark
+ * {@link Dataset} of rows, with the DataSourceV2 micro-batch glue as nested 
classes.
+ *
+ * <p>The dataset has two columns, {@value #COL_PAYLOAD} of type {@code 
BINARY} holding the element
+ * encoded with the supplied {@code WindowedValue} coder, and {@value 
#COL_EVENT_TS} of type {@code
+ * TIMESTAMP} holding the event timestamp of that element.
+ *
+ * <p>The event time watermark is declared here and only here. Spark 4 rejects 
a second {@code
+ * withWatermark} further down the plan, so downstream translators must never 
call it again.
+ */
+public final class UnboundedSourceDataset {
+
+  private static final Logger LOG = 
LoggerFactory.getLogger(UnboundedSourceDataset.class);
+
+  public static final String COL_PAYLOAD = "payload";
+
+  public static final String COL_EVENT_TS = "eventTimestamp";
+
+  public static final StructType SCHEMA =
+      new StructType()
+          .add(COL_PAYLOAD, DataTypes.BinaryType, false)
+          .add(COL_EVENT_TS, DataTypes.TimestampType, false);
+
+  private static final String SOURCE_NAME = "beam-unbounded";
+
+  private UnboundedSourceDataset() {}
+
+  /**
+   * Builds the streaming {@link Dataset} for {@code source} with the event 
time watermark applied.
+   *
+   * @param session the active Spark session
+   * @param source the Beam unbounded source to read
+   * @param windowedValueCoder the coder of the {@value #COL_PAYLOAD} column
+   * @param options the pipeline options, supplying the watermark delay and 
the micro-batch limits
+   * @param transformName the full name of the read transform, used for naming 
only
+   * @param <T> the element type of the source
+   * @param <CheckpointMarkT> the checkpoint mark type of the source
+   */
+  public static <T, CheckpointMarkT extends UnboundedSource.CheckpointMark> 
Dataset<Row> of(
+      SparkSession session,
+      UnboundedSource<T, CheckpointMarkT> source,
+      Coder<WindowedValue<T>> windowedValueCoder,
+      SparkStructuredStreamingPipelineOptions options,
+      String transformName) {
+    org.apache.spark.sql.classic.SparkSession classic =
+        (org.apache.spark.sql.classic.SparkSession) session;
+    Configuration hadoopConf = classic.sessionState().newHadoopConf();
+    BeamTable<T> table =
+        new BeamTable<>(
+            source,
+            windowedValueCoder,
+            broadcast(
+                session,
+                new SerializablePipelineOptions(options),
+                SerializablePipelineOptions.class),
+            broadcast(
+                session,
+                new SerializableConfiguration(hadoopConf),
+                SerializableConfiguration.class),
+            session.sparkContext().defaultParallelism(),
+            options.getMaxRecordsPerBatch(),
+            Math.max(1L, options.getMaxBatchDurationMillis()),
+            options.getReaderIdleTimeoutMillis(),
+            transformName);
+    LogicalPlan plan =
+        new StreamingRelationV2(
+            Option.empty(),
+            SOURCE_NAME,
+            table,
+            CaseInsensitiveStringMap.empty(),
+            DataTypeUtils.toAttributes(SCHEMA),
+            Option.empty(),
+            Option.empty(),
+            Option.empty());
+    Dataset<Row> rows = Dataset$.MODULE$.ofRows(classic, plan);
+    return rows.withWatermark(COL_EVENT_TS, options.getWatermarkDelayMillis() 
+ " milliseconds");
+  }
+
+  private static <T> Broadcast<T> broadcast(SparkSession session, T value, 
Class<T> type) {
+    return session.sparkContext().broadcast(value, ClassTag.apply(type));
+  }
+
+  /**
+   * Opaque epoch counter used as the Spark {@link Offset} of a Beam unbounded 
source.
+   *
+   * <p>The read position lives in Beam checkpoint marks on the executors, see 
{@link
+   * BeamSourceCheckpoint}. Equality is the base class comparison of {@link 
#json()}.
+   */
+  static class BeamOffset extends Offset {
+
+    public static final BeamOffset ZERO = new BeamOffset(0L);
+
+    private final long epoch;
+
+    public BeamOffset(long epoch) {
+      this.epoch = epoch;
+    }
+
+    public long epoch() {
+      return epoch;
+    }
+
+    @Override
+    public String json() {
+      return Long.toString(epoch);
+    }
+
+    public static BeamOffset fromJson(String json) {
+      try {
+        return new BeamOffset(Long.parseLong(json.trim()));
+      } catch (NumberFormatException e) {
+        throw new IllegalArgumentException("Not a valid BeamOffset: " + json, 
e);
+      }
+    }
+
+    @Override
+    public String toString() {
+      return json();
+    }
+  }
+
+  /** DataSourceV2 {@link Table} over one Beam unbounded source, micro-batch 
reads only. */
+  static final class BeamTable<T> implements Table, SupportsRead {
+    final UnboundedSource<T, ?> source;
+    final Coder<WindowedValue<T>> coder;
+    final Broadcast<SerializablePipelineOptions> options;
+    final Broadcast<SerializableConfiguration> hadoopConf;
+    final int desiredNumSplits;
+
+    /** Records per micro-batch across all splits, below 1 means unlimited. */
+    final long maxRecordsPerBatch;
+
+    final long maxBatchDurationMillis;
+    final long readerIdleTimeoutMillis;
+    final String transformName;
+
+    BeamTable(
+        UnboundedSource<T, ?> source,
+        Coder<WindowedValue<T>> coder,
+        Broadcast<SerializablePipelineOptions> options,
+        Broadcast<SerializableConfiguration> hadoopConf,
+        int desiredNumSplits,
+        long maxRecordsPerBatch,
+        long maxBatchDurationMillis,
+        long readerIdleTimeoutMillis,
+        String transformName) {
+      this.source = source;
+      this.coder = coder;
+      this.options = options;
+      this.hadoopConf = hadoopConf;
+      this.desiredNumSplits = desiredNumSplits;
+      this.maxRecordsPerBatch = maxRecordsPerBatch;
+      this.maxBatchDurationMillis = maxBatchDurationMillis;
+      this.readerIdleTimeoutMillis = readerIdleTimeoutMillis;
+      this.transformName = transformName;
+    }
+
+    @Override
+    public String name() {
+      return "BeamUnboundedSource[" + transformName + "]";
+    }
+
+    @Override
+    public StructType schema() {
+      return SCHEMA;
+    }
+
+    @Override
+    public Set<TableCapability> capabilities() {
+      return ImmutableSet.of(TableCapability.MICRO_BATCH_READ);
+    }
+
+    @Override
+    public ScanBuilder newScanBuilder(CaseInsensitiveStringMap ignored) {
+      return () -> new BeamScan<>(this);
+    }
+  }
+
+  private static final class BeamScan<T> implements Scan {
+    private final BeamTable<T> table;
+
+    BeamScan(BeamTable<T> table) {
+      this.table = table;
+    }
+
+    @Override
+    public StructType readSchema() {
+      return SCHEMA;
+    }
+
+    @Override
+    public String description() {
+      return table.name();
+    }
+
+    @Override
+    public MicroBatchStream toMicroBatchStream(String checkpointLocation) {
+      return new BeamMicroBatchStream<>(table, checkpointLocation);
+    }
+  }
+
+  /**
+   * Driver side {@link MicroBatchStream} over a Beam {@link UnboundedSource}.
+   *
+   * <p>Offsets are opaque epochs, {@link #latestOffset()} advances by one 
every trigger. The source
+   * is split once and the splits are pinned under the checkpoint location, 
every batch of every run
+   * plans the same partitions. {@link #commit} purges marks below the 
committed epoch on a
+   * background thread, Spark never asks for those again. Partitions carry no 
locality hint, the
+   * reader cache restores a split from its durable mark wherever it lands.
+   */
+  static class BeamMicroBatchStream<T> implements MicroBatchStream {
+
+    private final BeamTable<T> table;
+    private final String checkpointLocation;
+    private final BeamSourceCheckpoint checkpoint;
+
+    private final ExecutorService purger =
+        Executors.newSingleThreadExecutor(
+            runnable -> {
+              Thread thread = new Thread(runnable, "beam-source-mark-purge");
+              thread.setDaemon(true);
+              return thread;
+            });
+    private final AtomicBoolean purgeInFlight = new AtomicBoolean();
+    private final AtomicLong purgeRequested = new AtomicLong();
+
+    private long epoch;
+    private @Nullable List<UnboundedSource<T, ?>> splits;
+
+    BeamMicroBatchStream(BeamTable<T> table, String checkpointLocation) {
+      this.table = table;
+      this.checkpointLocation = checkpointLocation;
+      this.checkpoint =
+          new BeamSourceCheckpoint(checkpointLocation, 
table.hadoopConf.value().value());
+    }
+
+    @Override
+    public Offset initialOffset() {
+      return BeamOffset.ZERO;
+    }
+
+    @Override
+    public synchronized Offset latestOffset() {
+      return new BeamOffset(++epoch);
+    }
+
+    @Override
+    public Offset deserializeOffset(String json) {
+      BeamOffset offset = BeamOffset.fromJson(json);
+      fastForwardEpoch(offset.epoch());
+      return offset;
+    }
+
+    @Override
+    public InputPartition[] planInputPartitions(Offset start, Offset end) {
+      long startEpoch = ((BeamOffset) start).epoch();
+      long endEpoch = ((BeamOffset) end).epoch();
+      fastForwardEpoch(endEpoch);
+      List<UnboundedSource<T, ?>> pinned = splits();
+      try {
+        checkpoint.prepareEpoch(endEpoch);
+      } catch (IOException e) {
+        LOG.warn(
+            "Failed to prepare mark directory of epoch {} at {}.", endEpoch, 
checkpointLocation, e);
+      }
+      long[] quotas = splitQuotas(table.maxRecordsPerBatch, pinned.size(), 
endEpoch);
+      InputPartition[] partitions = new InputPartition[pinned.size()];
+      for (int i = 0; i < pinned.size(); i++) {
+        partitions[i] =
+            new BeamInputPartition<>(
+                pinned.get(i),
+                table.coder,
+                table.options,
+                table.hadoopConf,
+                checkpointLocation,
+                i,
+                startEpoch,
+                endEpoch,
+                quotas[i],
+                table.maxBatchDurationMillis,
+                table.readerIdleTimeoutMillis);
+      }
+      return partitions;
+    }
+
+    @Override
+    public PartitionReaderFactory createReaderFactory() {
+      return new BeamPartitionReaderFactory();
+    }
+
+    /** Purges marks below {@code end} off the stream thread, one purge runs 
at a time. */
+    @Override
+    public void commit(Offset end) {
+      purgeRequested.accumulateAndGet(((BeamOffset) end).epoch(), Math::max);
+      if (purgeInFlight.compareAndSet(false, true)) {
+        purger.execute(this::purgeRequested);
+      }
+    }
+
+    private void purgeRequested() {
+      long epoch;
+      do {
+        epoch = purgeRequested.get();
+        try {
+          checkpoint.purgeMarksBelow(epoch);
+        } catch (IOException | RuntimeException e) {
+          LOG.warn("Failed to purge marks below epoch {} at {}.", epoch, 
checkpointLocation, e);
+        }
+        purgeInFlight.set(false);
+      } while (purgeRequested.get() > epoch && 
purgeInFlight.compareAndSet(false, true));
+    }
+
+    @Override
+    public void stop() {
+      LOG.info(
+          "Stopping Beam micro-batch stream {} at {}.", table.transformName, 
checkpointLocation);
+      purger.shutdown();
+    }
+
+    /** Keeps {@link #latestOffset()} ahead of every epoch Spark logged before 
a restart. */
+    private synchronized void fastForwardEpoch(long seen) {
+      if (seen > epoch) {
+        LOG.info("Fast forwarding epoch of {} from {} to {}.", 
table.transformName, epoch, seen);
+        epoch = seen;
+      }
+    }
+
+    private synchronized List<UnboundedSource<T, ?>> splits() {
+      if (splits != null) {
+        return splits;
+      }
+      List<UnboundedSource<?, ?>> pinned;
+      try {
+        pinned = checkpoint.readSplits();
+      } catch (IOException e) {
+        throw new IllegalStateException("Failed to read pinned splits at " + 
checkpointLocation, e);
+      }
+      if (pinned == null) {
+        pinned = new ArrayList<>(splitSource());
+        try {
+          checkpoint.writeSplits(pinned);
+        } catch (IOException e) {
+          throw new IllegalStateException("Failed to pin splits at " + 
checkpointLocation, e);
+        }
+      } else {
+        LOG.info("Restored {} pinned split(s) from {}.", pinned.size(), 
checkpointLocation);
+      }
+      List<UnboundedSource<T, ?>> typed = new ArrayList<>(pinned.size());
+      for (UnboundedSource<?, ?> split : pinned) {
+        @SuppressWarnings("unchecked") // splits of this source share its 
element type
+        UnboundedSource<T, ?> cast = (UnboundedSource<T, ?>) split;
+        typed.add(cast);
+      }
+      splits = typed;
+      return typed;
+    }
+
+    private List<? extends UnboundedSource<T, ?>> splitSource() {
+      UnboundedSource<T, ?> source = table.source;
+      PipelineOptions options = table.options.value().get();
+      List<? extends UnboundedSource<T, ?>> result;
+      try {
+        result = source.split(table.desiredNumSplits, options);
+      } catch (Exception e) {
+        throw new IllegalStateException(
+            "Failed to split UnboundedSource " + 
source.getClass().getCanonicalName(), e);
+      }
+      if (result.isEmpty()) {
+        result = Collections.singletonList(source);
+      }
+      LOG.info(
+          "Split {} into {} partition(s), desired {}.",
+          table.transformName,
+          result.size(),
+          table.desiredNumSplits);
+      return result;
+    }
+
+    /**
+     * Divides the batch limit over the splits. A limit below 1 means 
unlimited and yields -1 for
+     * every split. Otherwise the remainder rotates with the epoch, so a limit 
below the split count
+     * gives one record to a rotating subset of splits per batch and 0 to the 
others.
+     */
+    static long[] splitQuotas(long maxRecordsPerBatch, int numSplits, long 
epoch) {
+      long[] quotas = new long[numSplits];
+      if (maxRecordsPerBatch < 1) {
+        Arrays.fill(quotas, -1L);
+        return quotas;
+      }
+      long base = maxRecordsPerBatch / numSplits;
+      long remainder = maxRecordsPerBatch % numSplits;
+      for (int i = 0; i < numSplits; i++) {
+        quotas[i] = base + ((i + epoch) % numSplits < remainder ? 1 : 0);
+      }
+      return quotas;
+    }
+  }
+
+  /** One split of a Beam unbounded source for one micro-batch, from epoch 
start to epoch end. */
+  static final class BeamInputPartition<T> implements InputPartition {
+
+    private static final long serialVersionUID = 1L;
+
+    final UnboundedSource<T, ?> split;
+    final Coder<WindowedValue<T>> coder;
+    final Broadcast<SerializablePipelineOptions> options;
+    final Broadcast<SerializableConfiguration> hadoopConf;
+    final String checkpointLocation;
+    final int splitId;
+    final long startEpoch;
+    final long endEpoch;
+
+    /** Records this split may emit in this micro-batch, below 0 means 
unlimited, 0 means none. */
+    final long maxRecords;
+
+    final long maxBatchDurationMillis;
+    final long readerIdleTimeoutMillis;
+
+    BeamInputPartition(
+        UnboundedSource<T, ?> split,
+        Coder<WindowedValue<T>> coder,
+        Broadcast<SerializablePipelineOptions> options,
+        Broadcast<SerializableConfiguration> hadoopConf,
+        String checkpointLocation,
+        int splitId,
+        long startEpoch,
+        long endEpoch,
+        long maxRecords,
+        long maxBatchDurationMillis,
+        long readerIdleTimeoutMillis) {
+      this.split = split;
+      this.coder = coder;
+      this.options = options;
+      this.hadoopConf = hadoopConf;
+      this.checkpointLocation = checkpointLocation;
+      this.splitId = splitId;
+      this.startEpoch = startEpoch;
+      this.endEpoch = endEpoch;
+      this.maxRecords = maxRecords;
+      this.maxBatchDurationMillis = maxBatchDurationMillis;
+      this.readerIdleTimeoutMillis = readerIdleTimeoutMillis;
+    }
+
+    @Override
+    public String toString() {
+      return "BeamInputPartition{checkpointLocation="
+          + checkpointLocation
+          + ", split="
+          + splitId
+          + ", epochs="
+          + startEpoch
+          + ".."
+          + endEpoch
+          + "}";
+    }
+  }
+
+  /** Creates a {@link BeamPartitionReader} for a {@link BeamInputPartition} 
on the executor. */
+  static final class BeamPartitionReaderFactory implements 
PartitionReaderFactory {
+
+    private static final long serialVersionUID = 1L;
+
+    @Override
+    public PartitionReader<InternalRow> createReader(InputPartition partition) 
{
+      try {
+        return new BeamPartitionReader<>((BeamInputPartition<?>) partition);
+      } catch (IOException e) {
+        throw new UncheckedIOException("Failed to open Beam reader for " + 
partition, e);
+      }
+    }
+  }
+
+  /**
+   * Reads one split of a Beam {@link UnboundedSource} for one micro-batch.
+   *
+   * <p>The batch ends at the record quota or at the deadline. The reader then 
writes its checkpoint
+   * mark durably at the end epoch and stays in {@link BeamReaderCache} for 
the next batch. A failed
+   * mark write fails the task. An attempt Spark killed or failed writes 
nothing and its reader is
+   * dropped, the retry restores from the durable mark at the start epoch.
+   *
+   * @param <T> the element type of the split
+   */
+  static final class BeamPartitionReader<T> implements 
PartitionReader<InternalRow> {
+
+    private static final Duration INITIAL_BACKOFF = Duration.millis(10);
+
+    private final String key;
+    private final UnboundedSource<T, ?> split;
+    private final Coder<WindowedValue<T>> coder;
+    private final BeamSourceCheckpoint checkpoint;
+    private final CachedReader<T> cached;
+    private final int splitId;
+    private final long endEpoch;
+    private final long maxRecords;
+    private final long maxBatchDurationMillis;
+
+    private long recordsRead;
+    private long deadlineMillis = -1L;
+    private boolean batchEnded;
+    private @Nullable InternalRow current;
+
+    BeamPartitionReader(BeamInputPartition<T> partition) throws IOException {
+      this.split = partition.split;
+      this.coder = partition.coder;
+      this.splitId = partition.splitId;
+      this.endEpoch = partition.endEpoch;
+      this.maxRecords = partition.maxRecords;
+      this.maxBatchDurationMillis = partition.maxBatchDurationMillis;
+      PipelineOptions options = partition.options.value().get();
+      Configuration conf = partition.hadoopConf.value().value();
+      BeamSourceCheckpoint checkpoint =
+          new BeamSourceCheckpoint(partition.checkpointLocation, conf);
+      this.checkpoint = checkpoint;
+      this.key = BeamReaderCache.key(partition.checkpointLocation, splitId);
+      long startEpoch = partition.startEpoch;
+      int splitId = this.splitId;
+      this.cached =
+          BeamReaderCache.acquire(
+              key,
+              startEpoch,
+              split,
+              options,
+              partition.readerIdleTimeoutMillis,
+              checkpoint::readSparkCommittedEpoch,
+              () -> checkpoint.readMark(splitId, startEpoch));
+    }
+
+    @Override
+    public boolean next() throws IOException {
+      if (deadlineMillis < 0) {
+        deadlineMillis = System.currentTimeMillis() + maxBatchDurationMillis;
+      }
+      BackOff backOff = null;
+      while (true) {
+        if (maxRecords >= 0 && recordsRead >= maxRecords) {
+          return endOfBatch(false);
+        }
+        long remaining = deadlineMillis - System.currentTimeMillis();
+        if (remaining <= 0) {
+          return endOfBatch(false);
+        }
+        if (cached.startOrAdvance()) {
+          recordsRead++;
+          current = toRow();
+          return true;
+        }
+        if (backOff == null) {
+          backOff = backOff(remaining);
+        }
+        try {
+          if (!BackOffUtils.next(Sleeper.DEFAULT, backOff)) {
+            return endOfBatch(false);
+          }
+        } catch (InterruptedException e) {
+          Thread.currentThread().interrupt();
+          return endOfBatch(true);
+        }
+      }
+    }
+
+    @Override
+    public InternalRow get() {
+      InternalRow row = current;
+      if (row == null) {
+        throw new IllegalStateException("No current row, next() did not return 
true.");
+      }
+      return row;
+    }
+
+    @Override
+    public void close() throws IOException {
+      endBatch(attemptDiscarded());
+      current = null;
+    }
+
+    private boolean endOfBatch(boolean discarded) throws IOException {
+      endBatch(discarded);
+      current = null;
+      return false;
+    }
+
+    /**
+     * Ends the batch once. A discarded attempt drops the reader and writes 
nothing. A reader that
+     * was never started has not moved, its start mark is written forward, an 
empty file standing
+     * for a fresh start.
+     */
+    private void endBatch(boolean discarded) throws IOException {
+      if (batchEnded) {
+        return;
+      }
+      batchEnded = true;
+      if (discarded) {
+        LOG.info("Attempt for Beam reader {} was discarded, dropping the 
reader.", key);
+        BeamReaderCache.invalidate(key);
+        return;
+      }
+      if (!cached.started()) {
+        byte[] startMark = cached.positionMark();
+        byte[] codedMark = startMark == null ? new byte[0] : startMark;
+        checkpoint.writeMark(splitId, endEpoch, codedMark);
+        cached.endBatch(endEpoch, null, codedMark);
+        return;
+      }
+      CheckpointMark mark = cached.reader().getCheckpointMark();
+      byte[] codedMark = encodeMark(split, mark);
+      checkpoint.writeMark(splitId, endEpoch, codedMark);
+      cached.endBatch(endEpoch, mark, codedMark);
+      LOG.debug("Beam reader {} read {} record(s) up to epoch {}.", key, 
recordsRead, endEpoch);
+    }
+
+    private static boolean attemptDiscarded() {
+      TaskContext context = TaskContext.get();
+      return context != null && (context.isInterrupted() || 
context.isFailed());
+    }
+
+    private static <MarkT extends CheckpointMark> byte[] encodeMark(
+        UnboundedSource<?, MarkT> source, CheckpointMark mark) {
+      @SuppressWarnings("unchecked") // getCheckpointMark returns the source's 
own mark type
+      MarkT typed = (MarkT) mark;
+      return CoderHelpers.toByteArray(typed, source.getCheckpointMarkCoder());
+    }
+
+    private static BackOff backOff(long remainingMillis) {
+      Duration remaining = Duration.millis(remainingMillis);
+      return FluentBackoff.DEFAULT
+          .withInitialBackoff(INITIAL_BACKOFF)
+          .withMaxBackoff(remaining)
+          .withMaxCumulativeBackoff(remaining)
+          .backoff();
+    }
+
+    private InternalRow toRow() {
+      Instant timestamp = cached.reader().getCurrentTimestamp();
+      WindowedValue<T> value =
+          
WindowedValues.timestampedValueInGlobalWindow(cached.reader().getCurrent(), 
timestamp);
+      byte[] payload = CoderHelpers.toByteArray(value, coder);
+      // Spark stores TimestampType as microseconds.
+      return new GenericInternalRow(new Object[] {payload, 
timestamp.getMillis() * 1000L});
+    }
+  }
+}
diff --git 
a/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamMicroBatchSourceTest.java
 
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamMicroBatchSourceTest.java
new file mode 100644
index 00000000000..4616c7563a2
--- /dev/null
+++ 
b/runners/spark/4/src/test/java/org/apache/beam/runners/spark/structuredstreaming/io/streaming/BeamMicroBatchSourceTest.java
@@ -0,0 +1,1028 @@
+/*
+ * 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.io.streaming;
+
+import static 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.UnboundedSourceDataset.COL_EVENT_TS;
+import static 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.UnboundedSourceDataset.COL_PAYLOAD;
+import static org.junit.Assert.assertArrayEquals;
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertFalse;
+import static org.junit.Assert.assertNotNull;
+import static org.junit.Assert.assertThrows;
+import static org.junit.Assert.assertTrue;
+
+import java.io.File;
+import java.io.IOException;
+import java.io.InputStream;
+import java.io.OutputStream;
+import java.io.Serializable;
+import java.nio.charset.StandardCharsets;
+import java.nio.file.Files;
+import java.sql.Timestamp;
+import java.util.ArrayList;
+import java.util.Arrays;
+import java.util.Collections;
+import java.util.HashSet;
+import java.util.List;
+import java.util.Map;
+import java.util.NoSuchElementException;
+import java.util.Set;
+import java.util.TreeSet;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
+import java.util.concurrent.atomic.AtomicInteger;
+import java.util.function.BooleanSupplier;
+import javax.annotation.Nullable;
+import org.apache.beam.runners.core.construction.SerializablePipelineOptions;
+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.UnboundedSourceDataset.BeamInputPartition;
+import 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.UnboundedSourceDataset.BeamMicroBatchStream;
+import 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.UnboundedSourceDataset.BeamOffset;
+import 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.UnboundedSourceDataset.BeamPartitionReader;
+import 
org.apache.beam.runners.spark.structuredstreaming.io.streaming.UnboundedSourceDataset.BeamTable;
+import org.apache.beam.sdk.coders.Coder;
+import org.apache.beam.sdk.coders.CustomCoder;
+import org.apache.beam.sdk.coders.StringUtf8Coder;
+import org.apache.beam.sdk.coders.VarIntCoder;
+import org.apache.beam.sdk.coders.VarLongCoder;
+import org.apache.beam.sdk.io.CountingSource;
+import org.apache.beam.sdk.io.UnboundedSource;
+import org.apache.beam.sdk.options.PipelineOptions;
+import org.apache.beam.sdk.options.PipelineOptionsFactory;
+import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
+import org.apache.beam.sdk.transforms.windowing.GlobalWindow;
+import org.apache.beam.sdk.util.CoderUtils;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.beam.sdk.values.WindowedValues;
+import org.apache.hadoop.conf.Configuration;
+import org.apache.spark.api.java.function.MapFunction;
+import org.apache.spark.api.java.function.VoidFunction2;
+import org.apache.spark.broadcast.Broadcast;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Encoders;
+import org.apache.spark.sql.Row;
+import org.apache.spark.sql.SparkSession;
+import org.apache.spark.sql.catalyst.InternalRow;
+import org.apache.spark.sql.catalyst.plans.logical.EventTimeWatermark;
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan;
+import org.apache.spark.sql.connector.read.InputPartition;
+import org.apache.spark.sql.streaming.StreamingQuery;
+import org.apache.spark.sql.streaming.StreamingQueryProgress;
+import org.apache.spark.sql.streaming.Trigger;
+import org.apache.spark.sql.util.CaseInsensitiveStringMap;
+import org.apache.spark.util.SerializableConfiguration;
+import org.joda.time.Instant;
+import org.junit.After;
+import org.junit.Before;
+import org.junit.BeforeClass;
+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;
+import scala.reflect.ClassTag;
+
+/**
+ * Tests for the Spark 4 DataSourceV2 micro-batch source wrapping a Beam 
{@link UnboundedSource}.
+ *
+ * <p>The epoch offsets of this source never settle, so {@code 
processAllAvailable()} would block
+ * forever. Every query runs with {@code Trigger.ProcessingTime(100)} and is 
stopped explicitly once
+ * the expected result arrived or the poll deadline expired.
+ */
+@Category(StreamingTest.class)
+@RunWith(JUnit4.class)
+public class BeamMicroBatchSourceTest implements Serializable {
+
+  @ClassRule public static final SparkSessionRule SESSION = new 
SparkSessionRule();
+
+  @Rule public transient TemporaryFolder temp = new TemporaryFolder();
+
+  private static final AtomicInteger TAGS = new AtomicInteger();
+
+  /** Rows per micro-batch per query name, driver side. */
+  private static final Map<String, List<List<Row>>> BATCHES = new 
ConcurrentHashMap<>();
+
+  private static final Coder<WindowedValue<String>> CODER =
+      WindowedValues.getFullCoder(StringUtf8Coder.of(), 
GlobalWindow.Coder.INSTANCE);
+
+  /** 2023-11-14T22:13:20Z, a plain modern timestamp with no rebase or DST 
subtleties. */
+  private static final long BASE_MILLIS = 1_700_000_000_000L;
+
+  private static final long INTERVAL_MILLIS = 1_000L;
+
+  private static final long POLL_TIMEOUT_MILLIS = 120_000L;
+
+  private static Broadcast<SerializablePipelineOptions> optionsBroadcast;
+  private static Broadcast<SerializableConfiguration> hadoopConfBroadcast;
+
+  private String tag;
+
+  @BeforeClass
+  public static void broadcastOnce() {
+    SparkSession session = SESSION.getSession();
+    Configuration conf =
+        ((org.apache.spark.sql.classic.SparkSession) 
session).sessionState().newHadoopConf();
+    optionsBroadcast =
+        session
+            .sparkContext()
+            .broadcast(
+                new 
SerializablePipelineOptions(PipelineOptionsFactory.create()),
+                ClassTag.apply(SerializablePipelineOptions.class));
+    hadoopConfBroadcast =
+        session
+            .sparkContext()
+            .broadcast(
+                new SerializableConfiguration(conf),
+                ClassTag.apply(SerializableConfiguration.class));
+  }
+
+  @Before
+  public void setUp() {
+    tag = "src" + TAGS.incrementAndGet();
+  }
+
+  @After
+  public void tearDown() {
+    BeamReaderCache.invalidateAll();
+    BATCHES.clear();
+    TestSource.forget(tag);
+  }
+
+  /** The {@code EventTimeWatermark} node survives typed maps in the logical 
and analyzed plan. */
+  @Test
+  public void testEventTimeWatermarkSurvivesTypedMap() {
+    Dataset<Row> rows = rows(1, 4, limited(1_000L, 200L));
+    assertTrue("source dataset must be streaming", rows.isStreaming());
+    assertWatermark("directly after withWatermark, logical plan", 
logical(rows));
+    assertWatermark("directly after withWatermark, analyzed plan", 
analyzed(rows));
+
+    Dataset<byte[]> typed =
+        rows.map((MapFunction<Row, byte[]>) row -> row.getAs(COL_PAYLOAD), 
Encoders.BINARY());
+    assertWatermark("after a typed map, logical plan", logical(typed));
+    assertWatermark("after a typed map, analyzed plan", analyzed(typed));
+
+    Dataset<byte[]> chained =
+        typed.map((MapFunction<byte[], byte[]>) bytes -> bytes, 
Encoders.BINARY());
+    assertWatermark("after two chained typed maps, logical plan", 
logical(chained));
+    assertWatermark("after two chained typed maps, analyzed plan", 
analyzed(chained));
+  }
+
+  /** A running query tracks the event time watermark past a typed map. */
+  @Test
+  public void testWatermarkIsTrackedAtRuntimeAfterTypedMap() throws Exception {
+    Dataset<byte[]> typed =
+        rows(1, 8, limited(1_000L, 200L))
+            .map((MapFunction<Row, byte[]>) row -> row.getAs(COL_PAYLOAD), 
Encoders.BINARY());
+    StreamingQuery query =
+        typed
+            .writeStream()
+            .format("noop")
+            .queryName(tag)
+            .outputMode("append")
+            .option("checkpointLocation", 
temp.newFolder(tag).getAbsolutePath())
+            .trigger(Trigger.ProcessingTime(100))
+            .start();
+    try {
+      String watermark = awaitWatermark(query);
+      assertNotNull("query never reported an event time watermark", watermark);
+      assertFalse("watermark stuck at the epoch: " + watermark, 
watermark.startsWith("1970-"));
+    } finally {
+      stopQuietly(query);
+    }
+  }
+
+  /** Payloads decode to the source elements and the timestamp column matches 
the element. */
+  @Test
+  public void testReadsElementsFromUnboundedSource() throws Exception {
+    int count = 8;
+    StreamingQuery query = start(rows(1, count, limited(1_000L, 200L)), tag, 
temp.newFolder(tag));
+    try {
+      await("all rows", () -> values(batches(tag)).size() >= count);
+    } finally {
+      stopQuietly(query);
+    }
+    List<String> values = new ArrayList<>();
+    for (List<Row> batch : batches(tag)) {
+      for (Row row : batch) {
+        WindowedValue<String> value = CoderUtils.decodeFromByteArray(CODER, 
row.getAs(COL_PAYLOAD));
+        values.add(value.getValue());
+        assertEquals(
+            value.getTimestamp().getMillis(), 
row.<Timestamp>getAs(COL_EVENT_TS).getTime());
+        assertEquals(
+            Collections.singletonList(GlobalWindow.INSTANCE), new 
ArrayList<>(value.getWindows()));
+        assertEquals(
+            BASE_MILLIS + TestSource.indexOf(value.getValue()) * 
INTERVAL_MILLIS,
+            value.getTimestamp().getMillis());
+      }
+    }
+    assertEquals(count, values.size());
+    assertEquals(TestSource.elements(tag, 1, count), new HashSet<>(values));
+  }
+
+  /** The default record limit is unlimited, an available source drains in one 
micro-batch. */
+  @Test
+  public void testUnlimitedRecordsPerBatchByDefault() throws Exception {
+    int count = 2500;
+    StreamingQuery query = start(rows(1, count, options(5_000L)), tag, 
temp.newFolder(tag));
+    try {
+      await("a non empty batch", () -> !nonEmptySizes(batches(tag)).isEmpty());
+    } finally {
+      stopQuietly(query);
+    }
+    assertEquals(Collections.singletonList(count), 
nonEmptySizes(batches(tag)));
+  }
+
+  /** The offset is an opaque epoch counter whose JSON is the bare number. */
+  @Test
+  public void testEpochOffsetRoundTrip() {
+    BeamOffset offset = new BeamOffset(42L);
+    assertEquals("42", offset.json());
+    assertEquals(42L, BeamOffset.fromJson("42").epoch());
+    assertEquals(0L, BeamOffset.ZERO.epoch());
+    assertEquals(new BeamOffset(7L), new BeamOffset(7L));
+    assertThrows(IllegalArgumentException.class, () -> 
BeamOffset.fromJson("x"));
+  }
+
+  /** A deserialized offset moves the epoch counter past itself. */
+  @Test
+  public void testEpochFastForwardsPastDeserializedOffset() throws Exception {
+    BeamMicroBatchStream<?> stream = 
newStream(temp.newFolder("ff-offset").getAbsolutePath());
+    stream.deserializeOffset("7");
+    BeamOffset next = (BeamOffset) stream.latestOffset();
+    assertTrue("latestOffset must move past the replayed epoch 7, got " + 
next, next.epoch() > 7L);
+  }
+
+  /** A planned end offset moves the epoch counter past itself. */
+  @Test
+  public void testEpochFastForwardsPastPlannedOffsets() throws Exception {
+    BeamMicroBatchStream<?> stream = 
newStream(temp.newFolder("ff-plan").getAbsolutePath());
+    InputPartition[] partitions =
+        stream.planInputPartitions(new BeamOffset(3L), new BeamOffset(9L));
+    assertTrue("at least one partition expected", partitions.length > 0);
+    BeamOffset next = (BeamOffset) stream.latestOffset();
+    assertTrue("latestOffset must move past the planned epoch 9, got " + next, 
next.epoch() > 9L);
+  }
+
+  /** The batch quota is divided over the splits and the remainder rotates 
with the epoch. */
+  @Test
+  public void testSplitQuotas() {
+    assertArrayEquals(
+        new long[] {1, 1, 1, 0, 0, 0, 0, 0}, 
BeamMicroBatchStream.splitQuotas(3, 8, 0));
+    assertArrayEquals(
+        new long[] {0, 0, 0, 0, 0, 1, 1, 1}, 
BeamMicroBatchStream.splitQuotas(3, 8, 3));
+    assertArrayEquals(new long[] {4, 3, 3}, 
BeamMicroBatchStream.splitQuotas(10, 3, 0));
+    assertArrayEquals(new long[] {3, 4, 3}, 
BeamMicroBatchStream.splitQuotas(10, 3, 2));
+    assertArrayEquals(new long[] {5, 5}, BeamMicroBatchStream.splitQuotas(10, 
2, 0));
+    assertArrayEquals(new long[] {-1, -1, -1}, 
BeamMicroBatchStream.splitQuotas(0, 3, 0));
+    assertArrayEquals(new long[] {-1, -1, -1}, 
BeamMicroBatchStream.splitQuotas(-1, 3, 0));
+    long[] many = BeamMicroBatchStream.splitQuotas(1, 200, 0);
+    assertEquals(200, many.length);
+    assertEquals(1L, many[0]);
+    assertEquals(1L, Arrays.stream(many).sum());
+  }
+
+  /** The record limit of a micro-batch is a total over all splits. */
+  @Test
+  public void testMaxRecordsPerBatchIsSharedAcrossSplits() throws Exception {
+    int shards = 2;
+    int count = 30;
+    StreamingQuery query =
+        start(rows(shards, count, limited(10L, 1_000L)), tag, 
temp.newFolder(tag));
+    try {
+      await("all rows", () -> values(batches(tag)).size() >= count);
+    } finally {
+      stopQuietly(query);
+    }
+    List<Integer> sizes = nonEmptySizes(batches(tag));
+    assertFalse("no rows arrived", sizes.isEmpty());
+    assertTrue("batch exceeds the shared limit: " + sizes, 
Collections.max(sizes) <= 10);
+    assertEquals(TestSource.elements(tag, shards, count), new 
HashSet<>(values(batches(tag))));
+  }
+
+  /**
+   * A limit below the split count emits at most the limit per batch and 
rotates over the splits.
+   */
+  @Test
+  public void testQuotaBelowSplitCountRotates() throws Exception {
+    int shards = 4;
+    StreamingQuery query = start(rows(shards, 40, limited(1L, 1_000L)), tag, 
temp.newFolder(tag));
+    try {
+      await("every shard", () -> shardsOf(values(batches(tag))).size() == 
shards);
+    } finally {
+      stopQuietly(query);
+    }
+    List<Integer> sizes = nonEmptySizes(batches(tag));
+    assertTrue("batch exceeds the limit of 1: " + sizes, 
Collections.max(sizes) <= 1);
+  }
+
+  /** A restart resumes every split from the last committed mark, replaying at 
most one batch. */
+  @Test
+  public void testRestartResumesFromCommittedMark() throws Exception {
+    int shards = 2;
+    int count = 80;
+    long limit = 4L;
+    File checkpointDir = temp.newFolder("restart");
+    String first = tag + "_a";
+    String second = tag + "_b";
+
+    StreamingQuery query = start(rows(shards, count, limited(limit, 1_000L)), 
first, checkpointDir);
+    try {
+      await("two commits", () -> committedBatchIds(checkpointDir).size() >= 2);
+    } finally {
+      stopQuietly(query);
+    }
+    BeamReaderCache.invalidateAll();
+    List<String> firstValues = values(batches(first));
+
+    Set<String> all = TestSource.elements(tag, shards, count);
+    query = start(rows(shards, count, limited(limit, 1_000L)), second, 
checkpointDir);
+    try {
+      await(
+          "union of both runs",
+          () -> {
+            Set<String> union = new HashSet<>(firstValues);
+            union.addAll(values(batches(second)));
+            return union.containsAll(all);
+          });
+    } finally {
+      stopQuietly(query);
+    }
+    List<String> secondValues = values(batches(second));
+
+    Set<String> union = new HashSet<>(firstValues);
+    union.addAll(secondValues);
+    assertEquals(all, union);
+    assertTrue(
+        "more than one batch replayed: " + firstValues.size() + " + " + 
secondValues.size(),
+        firstValues.size() + secondValues.size() <= count + limit);
+    for (int shard = 0; shard < shards; shard++) {
+      int min = Integer.MAX_VALUE;
+      for (String value : secondValues) {
+        if (TestSource.shardOf(value) == shard) {
+          min = Math.min(min, TestSource.indexOf(value));
+        }
+      }
+      assertTrue("run 2 delivered nothing for shard " + shard, min < 
Integer.MAX_VALUE);
+      assertTrue("run 2 restarted shard " + shard + " from element 0", min > 
0);
+    }
+  }
+
+  /** No split finalizes a position beyond its mark at the end epoch of the 
last committed batch. */
+  @Test
+  public void testMarksAreFinalizedOnlyAfterSparkCommit() throws Exception {
+    File checkpointDir = temp.newFolder("finalize");
+    runUntilCommits(checkpointDir, 3);
+
+    int finalizations = 0;
+    for (int shard = 0; shard < 2; shard++) {
+      int committed = committedPosition(checkpointDir, shard);
+      List<Integer> finalized = TestSource.finalized(tag, shard);
+      assertTrue(
+          "shard " + shard + " finalized " + finalized + " beyond committed " 
+ committed,
+          finalized.isEmpty() || Collections.max(finalized) <= committed);
+      finalizations += finalized.size();
+    }
+    assertTrue("no mark was finalized", finalizations > 0);
+  }
+
+  /**
+   * After a run the mark at the last committed epoch exists, every surviving 
mark is at or above
+   * the end epoch of the batch before the last constructed one, and the mark 
of batch 0 is gone.
+   */
+  @Test
+  public void testMarksBelowCommittedOffsetArePurged() throws Exception {
+    File checkpointDir = temp.newFolder("purge");
+    runUntilCommits(checkpointDir, 3);
+
+    long lastConstructed = Collections.max(batchIds(new File(checkpointDir, 
"offsets")));
+    long purgeFloor = endEpoch(checkpointDir, lastConstructed - 1);
+    long committedEpoch = committedEpoch(checkpointDir);
+    long firstEpoch = endEpoch(checkpointDir, 0);
+    assertTrue(committedEpoch >= purgeFloor);
+    assertTrue(purgeFloor > firstEpoch);
+    File sourceDir = sourceDir(checkpointDir);
+    awaitQuietly(
+        10_000L,
+        () -> {
+          for (int shard = 0; shard < 2; shard++) {
+            TreeSet<Long> epochs = markEpochs(sourceDir, shard);
+            if (epochs.isEmpty() || epochs.first() < purgeFloor) {
+              return false;
+            }
+          }
+          return true;
+        });
+    for (int shard = 0; shard < 2; shard++) {
+      TreeSet<Long> remaining = markEpochs(sourceDir, shard);
+      assertTrue(
+          "shard "
+              + shard
+              + " lost the mark at committed epoch "
+              + committedEpoch
+              + ": "
+              + remaining,
+          remaining.contains(committedEpoch));
+      assertTrue(
+          "shard " + shard + " kept marks below purge floor " + purgeFloor + 
": " + remaining,
+          remaining.first() >= purgeFloor);
+      assertFalse("shard " + shard + " kept the mark of batch 0", 
remaining.contains(firstEpoch));
+    }
+  }
+
+  /**
+   * After a stop an idle sweep finalizes exactly the marks of the last 
committed epoch, {@link
+   * BeamReaderCache#closeIdle(long)} is the one white box hook these tests 
use.
+   */
+  @Test
+  public void testStoppedQueryFinalizesLastCommittedMarks() throws Exception {
+    File checkpointDir = temp.newFolder("stopped");
+    runUntilCommits(checkpointDir, 3);
+    BeamReaderCache.closeIdle(Long.MAX_VALUE);
+
+    for (int shard = 0; shard < 2; shard++) {
+      int committed = committedPosition(checkpointDir, shard);
+      List<Integer> finalized = TestSource.finalized(tag, shard);
+      assertTrue(
+          "shard " + shard + " finalized " + finalized + ", committed " + 
committed,
+          finalized.contains(committed) && Collections.max(finalized) == 
committed);
+    }
+  }
+
+  /** A retried batch restarts from the durable mark at its start and 
finalizes nothing. */
+  @Test
+  public void testRetriedBatchRestartsFromDurableMark() throws Exception {
+    String location = sourceDir(temp.newFolder("protocol")).getAbsolutePath();
+    assertEquals(shardZero(0, 1, 2), readBatch(partition(location, 0, 1)));
+    assertEquals(shardZero(0, 1, 2), readBatch(partition(location, 0, 1)));
+    assertEquals(Collections.emptyList(), TestSource.finalized(tag, 0));
+    assertEquals(2, TestSource.created(tag));
+  }
+
+  /** A start epoch above zero without a durable mark is an invariant 
violation. */
+  @Test
+  public void testMissingMarkThrows() throws Exception {
+    String location = sourceDir(temp.newFolder("protocol")).getAbsolutePath();
+    assertThrows(
+        IllegalStateException.class, () -> new 
BeamPartitionReader<>(partition(location, 5, 6)));
+    assertEquals(0, TestSource.created(tag));
+  }
+
+  /** A failed mark write fails the batch after its rows, the retry recreates 
the reader. */
+  @Test
+  public void testRetryAfterFailedMarkWriteRecreatesReader() throws Exception {
+    File location = sourceDir(temp.newFolder("protocol"));
+    assertTrue(location.getParentFile().mkdirs() && location.createNewFile());
+    String file = location.getAbsolutePath();
+    assertEquals(shardZero(0, 1, 2), drainUntilFailure(partition(file, 0, 1), 
IOException.class));
+    assertEquals(shardZero(0, 1, 2), drainUntilFailure(partition(file, 0, 1), 
IOException.class));
+    assertEquals(Collections.emptyList(), TestSource.finalized(tag, 0));
+    assertEquals(2, TestSource.created(tag));
+  }
+
+  // 
---------------------------------------------------------------------------------------------
+  // query helpers
+  // 
---------------------------------------------------------------------------------------------
+
+  private static SparkStructuredStreamingPipelineOptions options(long 
maxBatchDurationMillis) {
+    SparkStructuredStreamingPipelineOptions options =
+        
PipelineOptionsFactory.create().as(SparkStructuredStreamingPipelineOptions.class);
+    options.setWatermarkDelayMillis(0L);
+    options.setMaxBatchDurationMillis(maxBatchDurationMillis);
+    return options;
+  }
+
+  private static SparkStructuredStreamingPipelineOptions limited(
+      long maxRecordsPerBatch, long maxBatchDurationMillis) {
+    SparkStructuredStreamingPipelineOptions options = 
options(maxBatchDurationMillis);
+    options.setMaxRecordsPerBatch(maxRecordsPerBatch);
+    return options;
+  }
+
+  private Dataset<Row> rows(
+      int shards, int count, SparkStructuredStreamingPipelineOptions options) {
+    return UnboundedSourceDataset.of(
+        SESSION.getSession(),
+        new TestSource(tag, shards, count),
+        CODER,
+        options,
+        "Read(TestSource)");
+  }
+
+  /** Builds the driver side stream through the table, with the session's 
broadcasts. */
+  private static BeamMicroBatchStream<?> newStream(String checkpointLocation) {
+    BeamTable<Long> table =
+        new BeamTable<>(
+            CountingSource.unbounded(),
+            WindowedValues.getFullCoder(VarLongCoder.of(), 
GlobalWindow.Coder.INSTANCE),
+            optionsBroadcast,
+            hadoopConfBroadcast,
+            2,
+            -1L,
+            200L,
+            600_000L,
+            "Read(CountingSource)");
+    return (BeamMicroBatchStream<?>)
+        table
+            .newScanBuilder(CaseInsensitiveStringMap.empty())
+            .build()
+            .toMicroBatchStream(checkpointLocation);
+  }
+
+  /** Starts a query collecting every micro-batch as one list into {@link 
#BATCHES}. */
+  private static StreamingQuery start(Dataset<Row> dataset, String queryName, 
File checkpointDir)
+      throws Exception {
+    BATCHES.put(queryName, Collections.synchronizedList(new ArrayList<>()));
+    return dataset
+        .writeStream()
+        .foreachBatch(
+            (VoidFunction2<Dataset<Row>, Long>)
+                (batch, batchId) -> {
+                  List<List<Row>> target = BATCHES.get(queryName);
+                  if (target != null) {
+                    target.add(batch.collectAsList());
+                  }
+                })
+        .queryName(queryName)
+        .outputMode("append")
+        .option("checkpointLocation", checkpointDir.getAbsolutePath())
+        .trigger(Trigger.ProcessingTime(100))
+        .start();
+  }
+
+  /** Runs two shards with a limit of 4 over a source that never drains until 
Spark committed. */
+  private void runUntilCommits(File checkpointDir, int commits) throws 
Exception {
+    StreamingQuery query = start(rows(2, 4_000, limited(4L, 1_000L)), tag, 
checkpointDir);
+    try {
+      await(commits + " commits", () -> 
committedBatchIds(checkpointDir).size() >= commits);
+    } finally {
+      stopQuietly(query);
+    }
+  }
+
+  private static void stopQuietly(StreamingQuery query) {
+    try {
+      query.stop();
+    } catch (Exception e) {
+      // Nothing useful to do while tearing a test query down.
+    }
+  }
+
+  private static List<List<Row>> batches(String queryName) {
+    List<List<Row>> batches = BATCHES.getOrDefault(queryName, 
Collections.emptyList());
+    synchronized (batches) {
+      return new ArrayList<>(batches);
+    }
+  }
+
+  private static List<String> values(List<List<Row>> batches) {
+    List<String> values = new ArrayList<>();
+    for (List<Row> batch : batches) {
+      for (Row row : batch) {
+        values.add(decode(row.getAs(COL_PAYLOAD)));
+      }
+    }
+    return values;
+  }
+
+  private static String decode(byte[] payload) {
+    try {
+      return CoderUtils.decodeFromByteArray(CODER, payload).getValue();
+    } catch (IOException e) {
+      throw new IllegalStateException(e);
+    }
+  }
+
+  private static List<Integer> nonEmptySizes(List<List<Row>> batches) {
+    List<Integer> sizes = new ArrayList<>();
+    for (List<Row> batch : batches) {
+      if (!batch.isEmpty()) {
+        sizes.add(batch.size());
+      }
+    }
+    return sizes;
+  }
+
+  private static Set<Integer> shardsOf(List<String> values) {
+    Set<Integer> shards = new HashSet<>();
+    for (String value : values) {
+      shards.add(TestSource.shardOf(value));
+    }
+    return shards;
+  }
+
+  private static void await(String what, BooleanSupplier condition) throws 
Exception {
+    if (!awaitQuietly(POLL_TIMEOUT_MILLIS, condition)) {
+      throw new AssertionError("timed out waiting for " + what);
+    }
+  }
+
+  private static boolean awaitQuietly(long timeoutMillis, BooleanSupplier 
condition)
+      throws Exception {
+    long deadline = System.currentTimeMillis() + timeoutMillis;
+    while (System.currentTimeMillis() < deadline) {
+      if (condition.getAsBoolean()) {
+        return true;
+      }
+      Thread.sleep(50L);
+    }
+    return condition.getAsBoolean();
+  }
+
+  /** Polls the query progress until it reports an event time watermark past 
the epoch. */
+  private static @Nullable String awaitWatermark(StreamingQuery query) throws 
Exception {
+    long deadline = System.currentTimeMillis() + POLL_TIMEOUT_MILLIS;
+    String last = null;
+    while (System.currentTimeMillis() < deadline) {
+      for (StreamingQueryProgress progress : query.recentProgress()) {
+        String watermark = progress.eventTime().get("watermark");
+        if (watermark != null) {
+          last = watermark;
+          if (!watermark.startsWith("1970-")) {
+            return watermark;
+          }
+        }
+      }
+      Thread.sleep(100L);
+    }
+    return last;
+  }
+
+  private static LogicalPlan logical(Dataset<?> dataset) {
+    return ((org.apache.spark.sql.classic.Dataset<?>) 
dataset).queryExecution().logical();
+  }
+
+  private static LogicalPlan analyzed(Dataset<?> dataset) {
+    return ((org.apache.spark.sql.classic.Dataset<?>) 
dataset).queryExecution().analyzed();
+  }
+
+  private static void assertWatermark(String what, LogicalPlan plan) {
+    assertTrue(
+        "no EventTimeWatermark node found " + what + ":\n" + plan.treeString(),
+        containsWatermark(plan));
+  }
+
+  private static boolean containsWatermark(LogicalPlan plan) {
+    if (plan instanceof EventTimeWatermark) {
+      return true;
+    }
+    scala.collection.Iterator<LogicalPlan> children = 
plan.children().iterator();
+    while (children.hasNext()) {
+      if (containsWatermark(children.next())) {
+        return true;
+      }
+    }
+    return false;
+  }
+
+  // 
---------------------------------------------------------------------------------------------
+  // checkpoint helpers
+  // 
---------------------------------------------------------------------------------------------
+
+  private static File sourceDir(File checkpointDir) {
+    return new File(checkpointDir, "sources/0");
+  }
+
+  /** Numeric file names in a Spark log directory, temp and hidden files 
excluded. */
+  private static TreeSet<Long> batchIds(File dir) {
+    TreeSet<Long> ids = new TreeSet<>();
+    String[] names = dir.list();
+    if (names == null) {
+      return ids;
+    }
+    for (String name : names) {
+      if (!name.startsWith(".") && !name.endsWith(".tmp")) {
+        try {
+          ids.add(Long.parseLong(name));
+        } catch (NumberFormatException e) {
+          // not a log entry
+        }
+      }
+    }
+    return ids;
+  }
+
+  /** Epochs under {@code marks/<epoch>/} holding a mark file of {@code 
shard}. */
+  private static TreeSet<Long> markEpochs(File sourceDir, int shard) {
+    TreeSet<Long> epochs = new TreeSet<>();
+    for (long epoch : batchIds(new File(sourceDir, "marks"))) {
+      if (new File(sourceDir, "marks/" + epoch + "/" + shard).exists()) {
+        epochs.add(epoch);
+      }
+    }
+    return epochs;
+  }
+
+  private static TreeSet<Long> committedBatchIds(File checkpointDir) {
+    return batchIds(new File(checkpointDir, "commits"));
+  }
+
+  /** The end epoch of a batch, the offset line of the single source in {@code 
offsets/<id>}. */
+  private static long endEpoch(File checkpointDir, long batchId) throws 
IOException {
+    File file = new File(new File(checkpointDir, "offsets"), 
Long.toString(batchId));
+    List<String> lines = new ArrayList<>();
+    for (String line : Files.readAllLines(file.toPath(), 
StandardCharsets.UTF_8)) {
+      if (!line.trim().isEmpty()) {
+        lines.add(line.trim());
+      }
+    }
+    assertEquals("one source expected in " + lines, 3, lines.size());
+    return BeamOffset.fromJson(lines.get(2)).epoch();
+  }
+
+  private static long committedEpoch(File checkpointDir) throws IOException {
+    return endEpoch(checkpointDir, committedBatchIds(checkpointDir).last());
+  }
+
+  /** The position in the mark of {@code shard} at the end epoch of the last 
committed batch. */
+  private static int committedPosition(File checkpointDir, int shard) throws 
IOException {
+    long epoch = committedEpoch(checkpointDir);
+    BeamSourceCheckpoint checkpoint =
+        new BeamSourceCheckpoint(sourceDir(checkpointDir).getAbsolutePath(), 
new Configuration());
+    byte[] coded = checkpoint.readMark(shard, epoch);
+    assertNotNull("no mark at committed epoch " + epoch + " for shard " + 
shard, coded);
+    return CoderUtils.decodeFromByteArray(TestSource.MARK_CODER, coded).next;
+  }
+
+  // 
---------------------------------------------------------------------------------------------
+  // hand built partition helpers
+  // 
---------------------------------------------------------------------------------------------
+
+  /** Split 0 of a single shard source of 100 elements from epoch {@code 
start} to {@code end}. */
+  private BeamInputPartition<String> partition(String location, long start, 
long end) {
+    TestSource split = new TestSource(tag, 1, 100).split(1, 
PipelineOptionsFactory.create()).get(0);
+    return new BeamInputPartition<>(
+        split,
+        CODER,
+        optionsBroadcast,
+        hadoopConfBroadcast,
+        location,
+        0,
+        start,
+        end,
+        3L,
+        30_000L,
+        600_000L);
+  }
+
+  private static List<String> readBatch(BeamInputPartition<String> partition) 
throws IOException {
+    List<String> values = new ArrayList<>();
+    drainInto(new BeamPartitionReader<>(partition), values);
+    return values;
+  }
+
+  private static void drainInto(BeamPartitionReader<String> reader, 
List<String> values)
+      throws IOException {
+    while (reader.next()) {
+      InternalRow row = reader.get();
+      values.add(decode(row.getBinary(0)));
+    }
+    reader.close();
+  }
+
+  /** Opens and drains a batch expected to fail, returns what it delivered 
before failing. */
+  private static List<String> drainUntilFailure(
+      BeamInputPartition<String> partition, Class<? extends Exception> 
failure) throws IOException {
+    BeamPartitionReader<String> reader = new BeamPartitionReader<>(partition);
+    List<String> values = new ArrayList<>();
+    assertThrows(failure, () -> drainInto(reader, values));
+    return values;
+  }
+
+  private List<String> shardZero(int... indexes) {
+    List<String> elements = new ArrayList<>();
+    for (int index : indexes) {
+      elements.add(TestSource.element(tag, 0, index));
+    }
+    return elements;
+  }
+
+  // 
---------------------------------------------------------------------------------------------
+  // the shared in memory UnboundedSource
+  // 
---------------------------------------------------------------------------------------------
+
+  /**
+   * Splits into one sub source per shard, each over {@code count / shards} 
elements named {@code
+   * <tag>-<shard>-<index>} with evenly spaced timestamps. Marks are not Java 
serializable, they
+   * record the position they finalize under {@code <tag>/<shard>}, readers 
are counted per tag.
+   */
+  static final class TestSource extends UnboundedSource<String, 
TestSource.Mark> {
+    private static final long serialVersionUID = 1L;
+
+    static final Coder<Mark> MARK_CODER = new MarkCoder();
+
+    private static final ConcurrentMap<String, List<Integer>> FINALIZED = new 
ConcurrentHashMap<>();
+    private static final ConcurrentMap<String, AtomicInteger> CREATED = new 
ConcurrentHashMap<>();
+
+    private final String tag;
+    private final int shard;
+    private final int shards;
+    private final int perShard;
+
+    TestSource(String tag, int shards, int count) {
+      this(tag, -1, shards, count / shards);
+    }
+
+    private TestSource(String tag, int shard, int shards, int perShard) {
+      this.tag = tag;
+      this.shard = shard;
+      this.shards = shards;
+      this.perShard = perShard;
+    }
+
+    static Set<String> elements(String tag, int shards, int count) {
+      Set<String> elements = new HashSet<>();
+      for (int shard = 0; shard < shards; shard++) {
+        for (int index = 0; index < count / shards; index++) {
+          elements.add(element(tag, shard, index));
+        }
+      }
+      return elements;
+    }
+
+    static String element(String tag, int shard, int index) {
+      return tag + "-" + shard + "-" + index;
+    }
+
+    static int shardOf(String element) {
+      String head = element.substring(0, element.lastIndexOf('-'));
+      return Integer.parseInt(head.substring(head.lastIndexOf('-') + 1));
+    }
+
+    static int indexOf(String element) {
+      return Integer.parseInt(element.substring(element.lastIndexOf('-') + 1));
+    }
+
+    static List<Integer> finalized(String tag, int shard) {
+      List<Integer> positions = FINALIZED.get(key(tag, shard));
+      if (positions == null) {
+        return Collections.emptyList();
+      }
+      synchronized (positions) {
+        return new ArrayList<>(positions);
+      }
+    }
+
+    static int created(String tag) {
+      AtomicInteger created = CREATED.get(tag);
+      return created == null ? 0 : created.get();
+    }
+
+    static void forget(String tag) {
+      FINALIZED.keySet().removeIf(key -> key.startsWith(tag + "/"));
+      CREATED.remove(tag);
+    }
+
+    private static String key(String tag, int shard) {
+      return tag + "/" + shard;
+    }
+
+    @Override
+    public List<TestSource> split(int desiredNumSplits, PipelineOptions 
options) {
+      if (shard >= 0) {
+        return Collections.singletonList(this);
+      }
+      List<TestSource> splits = new ArrayList<>();
+      for (int i = 0; i < shards; i++) {
+        splits.add(new TestSource(tag, i, shards, perShard));
+      }
+      return splits;
+    }
+
+    @Override
+    public UnboundedReader<String> createReader(PipelineOptions options, 
@Nullable Mark mark) {
+      if (shard < 0) {
+        throw new IllegalStateException("split before reading");
+      }
+      CREATED.computeIfAbsent(tag, t -> new AtomicInteger()).incrementAndGet();
+      return new Reader(this, mark == null ? 0 : mark.next);
+    }
+
+    @Override
+    public Coder<Mark> getCheckpointMarkCoder() {
+      return MARK_CODER;
+    }
+
+    @Override
+    public Coder<String> getOutputCoder() {
+      return StringUtf8Coder.of();
+    }
+
+    /** Position of the next element of a shard, deliberately not {@link 
Serializable}. */
+    static final class Mark implements UnboundedSource.CheckpointMark {
+      private final String tag;
+      private final int shard;
+      final int next;
+
+      Mark(String tag, int shard, int next) {
+        this.tag = tag;
+        this.shard = shard;
+        this.next = next;
+      }
+
+      @Override
+      public void finalizeCheckpoint() {
+        FINALIZED
+            .computeIfAbsent(key(tag, shard), k -> 
Collections.synchronizedList(new ArrayList<>()))
+            .add(next);
+      }
+    }
+
+    private static final class MarkCoder extends CustomCoder<Mark> {
+      private static final long serialVersionUID = 1L;
+
+      @Override
+      public void encode(Mark mark, OutputStream out) throws IOException {
+        StringUtf8Coder.of().encode(mark.tag, out);
+        VarIntCoder.of().encode(mark.shard, out);
+        VarIntCoder.of().encode(mark.next, out);
+      }
+
+      @Override
+      public Mark decode(InputStream in) throws IOException {
+        return new Mark(
+            StringUtf8Coder.of().decode(in),
+            VarIntCoder.of().decode(in),
+            VarIntCoder.of().decode(in));
+      }
+    }
+
+    private static final class Reader extends UnboundedReader<String> {
+      private final TestSource source;
+      private int next;
+      private int current = -1;
+
+      Reader(TestSource source, int next) {
+        this.source = source;
+        this.next = next;
+      }
+
+      @Override
+      public boolean start() {
+        return advance();
+      }
+
+      @Override
+      public boolean advance() {
+        if (next < source.perShard) {
+          current = next++;
+          return true;
+        }
+        return false;
+      }
+
+      @Override
+      public String getCurrent() throws NoSuchElementException {
+        if (current < 0) {
+          throw new NoSuchElementException();
+        }
+        return element(source.tag, source.shard, current);
+      }
+
+      @Override
+      public Instant getCurrentTimestamp() throws NoSuchElementException {
+        if (current < 0) {
+          throw new NoSuchElementException();
+        }
+        return new Instant(
+            BASE_MILLIS + (source.shard * source.perShard + current) * 
INTERVAL_MILLIS);
+      }
+
+      @Override
+      public Instant getWatermark() {
+        return current < 0 ? BoundedWindow.TIMESTAMP_MIN_VALUE : 
getCurrentTimestamp();
+      }
+
+      @Override
+      public CheckpointMark getCheckpointMark() {
+        return new Mark(source.tag, source.shard, next);
+      }
+
+      @Override
+      public UnboundedSource<String, ?> getCurrentSource() {
+        return source;
+      }
+
+      @Override
+      public void close() {}
+    }
+  }
+}
diff --git 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineOptions.java
 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineOptions.java
index 391350fd348..fb0192dba86 100644
--- 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineOptions.java
+++ 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineOptions.java
@@ -63,6 +63,15 @@ public interface SparkStructuredStreamingPipelineOptions 
extends SparkCommonPipe
 
   void setMaxBatchDurationMillis(long value);
 
+  @Description(
+      "Idle time in milliseconds after which an executor closes a cached 
unbounded reader. Must "
+          + "exceed the longest gap between two micro-batches, a closed 
reader's last checkpoint "
+          + "mark is not finalized and the source redelivers (streaming mode 
only).")
+  @Default.Long(600_000)
+  long getReaderIdleTimeoutMillis();
+
+  void setReaderIdleTimeoutMillis(long value);
+
   @Description(
       "Test-oriented: gracefully stop streaming queries after this many 
consecutive empty "
           + "micro-batches. Disabled if negative (streaming mode only).")

Reply via email to