This is an automated email from the ASF dual-hosted git repository.

JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git


The following commit(s) were added to refs/heads/master by this push:
     new f580fe0c8a [spark] Expose written columns for streaming micro-batches 
(#9023)
f580fe0c8a is described below

commit f580fe0c8a8264ad9ffb6a9396c64c14e768674c
Author: LsomeYeah <[email protected]>
AuthorDate: Mon Aug 17 13:44:46 2026 +0800

    [spark] Expose written columns for streaming micro-batches (#9023)
---
 docs/docs/spark/structured-streaming.md            |  29 +++
 .../apache/paimon/utils/DataEvolutionUtils.java    |  67 +++++-
 .../paimon/utils/DataEvolutionUtilsTest.java       | 167 ++++++++++++++
 .../apache/paimon/spark/PaimonInputPartition.scala |  46 ++++
 .../spark/PaimonSparkMicroBatchMetadata.scala      | 202 ++++++++++++++++
 .../spark/sources/PaimonMicroBatchStream.scala     |  42 +++-
 .../org/apache/paimon/spark/PaimonSourceTest.scala | 253 ++++++++++++++++++++-
 7 files changed, 796 insertions(+), 10 deletions(-)

diff --git a/docs/docs/spark/structured-streaming.md 
b/docs/docs/spark/structured-streaming.md
index 91801f3a4d..41c0a19dbf 100644
--- a/docs/docs/spark/structured-streaming.md
+++ b/docs/docs/spark/structured-streaming.md
@@ -198,6 +198,35 @@ val query = spark.readStream
   .start()
 ```
 
+### Written Columns of a Micro-Batch
+
+`foreachBatch` consumers can inspect which Paimon field IDs were written by 
the data files admitted to the current micro-batch. Call 
`PaimonSparkMicroBatchMetadata.writtenColumnIds` with the raw `Dataset` passed 
to `foreachBatch`. Paimon resolves the file metadata lazily when this method is 
called.
+
+```scala
+import org.apache.paimon.spark.PaimonSparkMicroBatchMetadata
+import org.apache.spark.sql.{Dataset, Row}
+
+val query = spark.readStream
+  .format("paimon")
+  .table("table_name")
+  .writeStream
+  .option("checkpointLocation", "/path/to/checkpoint")
+  .foreachBatch { (batch: Dataset[Row], _: Long) =>
+    val writtenColumnIds = 
PaimonSparkMicroBatchMetadata.writtenColumnIds(batch)
+    if (!writtenColumnIds.isPresent) {
+      // Metadata is unavailable; conservatively process all columns.
+    } else {
+      val fieldIds = writtenColumnIds.get()
+      // Process the exact set of written Paimon field IDs.
+    }
+  }
+  .start()
+```
+
+A present `Optional` contains the complete, immutable list of written field 
IDs in ascending order. The list may be empty; that is a known empty set, not 
unknown metadata.
+
+An empty `Optional` means that metadata is unavailable, for example because a 
file or schema cannot be resolved, the micro-batch is empty, the `Dataset` is 
not the raw batch from a query with exactly one distinct Paimon streaming 
source, or its lineage is incomplete or ambiguous. An empty `Optional` does not 
mean that no columns were written; callers must fall back to processing all 
columns.
+
 Paimon Structured Streaming supports read row in the form of changelog (add 
rowkind column in row to represent its
 change type) in two ways:
 
diff --git 
a/paimon-core/src/main/java/org/apache/paimon/utils/DataEvolutionUtils.java 
b/paimon-core/src/main/java/org/apache/paimon/utils/DataEvolutionUtils.java
index c1294f3946..262da197ae 100644
--- a/paimon-core/src/main/java/org/apache/paimon/utils/DataEvolutionUtils.java
+++ b/paimon-core/src/main/java/org/apache/paimon/utils/DataEvolutionUtils.java
@@ -20,13 +20,21 @@ package org.apache.paimon.utils;
 
 import org.apache.paimon.io.DataFileMeta;
 import org.apache.paimon.schema.TableSchema;
+import org.apache.paimon.table.SpecialFields;
+import org.apache.paimon.table.source.DataSplit;
 import org.apache.paimon.types.DataField;
 
+import java.util.ArrayList;
 import java.util.Collection;
+import java.util.Collections;
 import java.util.Comparator;
+import java.util.HashMap;
 import java.util.HashSet;
 import java.util.List;
+import java.util.Map;
+import java.util.Optional;
 import java.util.Set;
+import java.util.TreeSet;
 import java.util.function.Function;
 import java.util.stream.Collectors;
 
@@ -38,21 +46,76 @@ import static 
org.apache.paimon.utils.Preconditions.checkState;
 /** Util class for data evolution. */
 public class DataEvolutionUtils {
 
+    /**
+     * Collect exact written field ids; an empty list is exact and an empty 
optional is unresolved.
+     */
+    public static Optional<List<Integer>> collectWrittenColumnIds(
+            Collection<DataSplit> splits, Function<Long, TableSchema> 
schemaLoader) {
+        Set<Integer> fieldIds = new TreeSet<>();
+        Map<Long, List<DataField>> schemaFieldsCache = new HashMap<>();
+        Map<Pair<Long, List<String>>, Set<Integer>> fieldIdsCache = new 
HashMap<>();
+        try {
+            for (DataSplit split : splits) {
+                for (DataFileMeta file : split.dataFiles()) {
+                    Pair<Long, List<String>> cacheKey = 
Pair.of(file.schemaId(), file.writeCols());
+                    Set<Integer> fileFieldIds = fieldIdsCache.get(cacheKey);
+                    if (fileFieldIds == null) {
+                        List<DataField> schemaFields =
+                                schemaFieldsCache.computeIfAbsent(
+                                        file.schemaId(),
+                                        schemaId -> {
+                                            TableSchema schema = 
schemaLoader.apply(schemaId);
+                                            checkArgument(
+                                                    schema != null,
+                                                    "Cannot find schema %s.",
+                                                    schemaId);
+                                            return schema.fields();
+                                        });
+                        fileFieldIds = resolveFileFieldIds(schemaFields, file, 
true);
+                        fieldIdsCache.put(cacheKey, fileFieldIds);
+                    }
+                    fieldIds.addAll(fileFieldIds);
+                }
+            }
+        } catch (RuntimeException e) {
+            return Optional.empty();
+        }
+        return Optional.of(Collections.unmodifiableList(new 
ArrayList<>(fieldIds)));
+    }
+
     /**
      * Table field ids physically present in a file, resolved through the 
schema used to write it.
      */
     public static Set<Integer> fileFieldIds(
             Function<Long, TableSchema> scanTableSchema, DataFileMeta file) {
-        TableSchema schema = scanTableSchema.apply(file.schemaId());
+        return 
resolveFileFieldIds(scanTableSchema.apply(file.schemaId()).fields(), file, 
false);
+    }
+
+    private static Set<Integer> resolveFileFieldIds(
+            List<DataField> schemaFields, DataFileMeta file, boolean strict) {
         List<String> writeCols = file.writeCols();
         Set<String> writeColNames = writeCols == null ? null : new 
HashSet<>(writeCols);
+        Set<String> unresolved =
+                strict && writeColNames != null ? new HashSet<>(writeColNames) 
: null;
         Set<Integer> ids = new HashSet<>();
-        for (DataField field : schema.fields()) {
+        for (DataField field : schemaFields) {
             // writeCols may also contain physical row-tracking fields outside 
the table schema.
             if (writeColNames == null || writeColNames.contains(field.name())) 
{
                 ids.add(field.id());
+                if (unresolved != null) {
+                    unresolved.remove(field.name());
+                }
             }
         }
+
+        if (unresolved != null) {
+            unresolved.removeIf(SpecialFields::isSystemField);
+            checkArgument(
+                    unresolved.isEmpty(),
+                    "Cannot find write columns %s in schema %s.",
+                    unresolved,
+                    file.schemaId());
+        }
         return ids;
     }
 
diff --git 
a/paimon-core/src/test/java/org/apache/paimon/utils/DataEvolutionUtilsTest.java 
b/paimon-core/src/test/java/org/apache/paimon/utils/DataEvolutionUtilsTest.java
index 33feb9d850..c499c5f276 100644
--- 
a/paimon-core/src/test/java/org/apache/paimon/utils/DataEvolutionUtilsTest.java
+++ 
b/paimon-core/src/test/java/org/apache/paimon/utils/DataEvolutionUtilsTest.java
@@ -18,11 +18,15 @@
 
 package org.apache.paimon.utils;
 
+import org.apache.paimon.data.BinaryRow;
 import org.apache.paimon.io.DataFileMeta;
+import org.apache.paimon.schema.Schema;
 import org.apache.paimon.schema.TableSchema;
 import org.apache.paimon.stats.SimpleStats;
 import org.apache.paimon.table.SpecialFields;
+import org.apache.paimon.table.source.DataSplit;
 import org.apache.paimon.types.DataField;
+import org.apache.paimon.types.DataTypes;
 import org.apache.paimon.types.IntType;
 
 import org.junit.jupiter.api.Test;
@@ -31,10 +35,17 @@ import java.util.Arrays;
 import java.util.Collections;
 import java.util.HashMap;
 import java.util.List;
+import java.util.Map;
+import java.util.Optional;
+import java.util.concurrent.atomic.AtomicInteger;
 import java.util.function.Function;
 
 import static org.assertj.core.api.Assertions.assertThat;
 import static org.assertj.core.api.Assertions.assertThatThrownBy;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.spy;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
 
 /** Test for {@link DataEvolutionUtils}. */
 public class DataEvolutionUtilsTest {
@@ -107,6 +118,128 @@ public class DataEvolutionUtilsTest {
                                         1,
                                         Collections.singletonList("other"))))
                 .containsExactly(2);
+        assertThat(
+                        DataEvolutionUtils.fileFieldIds(
+                                ignored -> schema,
+                                dataFile(
+                                        "unknown.parquet",
+                                        1,
+                                        Collections.singletonList("unknown"))))
+                .isEmpty();
+    }
+
+    @Test
+    public void testCollectWrittenColumnIdsAcrossSchemas() {
+        Map<Long, TableSchema> schemas = new HashMap<>();
+        schemas.put(
+                0L,
+                tableSchema(
+                        0L,
+                        new DataField(1, "a", DataTypes.INT()),
+                        new DataField(2, "old_name", DataTypes.STRING())));
+        schemas.put(
+                1L,
+                tableSchema(
+                        1L,
+                        new DataField(2, "new_name", DataTypes.STRING()),
+                        new DataField(3, "c", DataTypes.BIGINT())));
+
+        DataFileMeta oldSchemaFile = dataFile(0L, Arrays.asList("a", 
"old_name"));
+        DataFileMeta newSchemaFile = dataFile(1L, Arrays.asList("new_name", 
"c"));
+
+        assertThat(collectWrittenColumnIds(schemas::get, oldSchemaFile, 
newSchemaFile))
+                .hasValue(Arrays.asList(1, 2, 3));
+    }
+
+    @Test
+    public void testCollectWrittenColumnIdsFallsBackWhenResolutionFails() {
+        DataFileMeta unknownSchemaFile = dataFile(99L, 
Collections.singletonList("a"));
+        assertThat(collectWrittenColumnIds(ignored -> null, unknownSchemaFile))
+                .as("unknown schema")
+                .isEmpty();
+
+        DataFileMeta unresolvedSchemaFile = dataFile(1L, 
Collections.singletonList("missing"));
+        assertThat(
+                        collectWrittenColumnIds(
+                                ignored -> {
+                                    throw new IllegalArgumentException("schema 
cannot be resolved");
+                                },
+                                unresolvedSchemaFile))
+                .as("schema loader failure")
+                .isEmpty();
+
+        TableSchema schema = tableSchema(1L, new DataField(1, "a", 
DataTypes.INT()));
+        DataFileMeta unknownColumnFile = dataFile(1L, 
Collections.singletonList("missing"));
+        assertThat(collectWrittenColumnIds(ignored -> schema, 
unknownColumnFile))
+                .as("unknown non-system write column")
+                .isEmpty();
+    }
+
+    @Test
+    public void testCollectWrittenColumnIdsIgnoresSystemFields() {
+        TableSchema schema = tableSchema(1L, new DataField(1, "a", 
DataTypes.INT()));
+        DataFileMeta file =
+                dataFile(
+                        1L,
+                        Arrays.asList(
+                                SpecialFields.ROW_ID.name(),
+                                "a",
+                                SpecialFields.SEQUENCE_NUMBER.name()));
+
+        assertThat(collectWrittenColumnIds(ignored -> schema, file))
+                .hasValue(Collections.singletonList(1));
+
+        assertThat(
+                        collectWrittenColumnIds(
+                                ignored -> schema,
+                                dataFile(
+                                        1L,
+                                        Arrays.asList(
+                                                SpecialFields.ROW_ID.name(),
+                                                
SpecialFields.SEQUENCE_NUMBER.name()))))
+                .hasValue(Collections.emptyList());
+    }
+
+    @Test
+    public void testCollectWrittenColumnIdsCachesSchemaAcrossProjections() {
+        TableSchema schema =
+                spy(
+                        tableSchema(
+                                1L,
+                                new DataField(1, "a", DataTypes.INT()),
+                                new DataField(2, "b", DataTypes.STRING())));
+        DataFileMeta first = dataFile(1L, Collections.singletonList("a"));
+        DataFileMeta second = dataFile(1L, Collections.singletonList("b"));
+        DataFileMeta repeated = dataFile(1L, Collections.singletonList("a"));
+        AtomicInteger schemaLoads = new AtomicInteger();
+
+        Optional<List<Integer>> result =
+                collectWrittenColumnIds(
+                        ignored -> {
+                            schemaLoads.incrementAndGet();
+                            return schema;
+                        },
+                        first,
+                        second,
+                        repeated);
+
+        assertThat(result.get()).containsExactly(1, 2);
+        assertThat(schemaLoads).hasValue(1);
+        verify(schema).fields();
+        verify(repeated).writeCols();
+    }
+
+    @Test
+    public void testCollectWrittenColumnIdsExpandsLegacyFileSchema() {
+        TableSchema schema =
+                tableSchema(
+                        1L,
+                        new DataField(1, "a", DataTypes.INT()),
+                        new DataField(2, "b", DataTypes.STRING()));
+        DataFileMeta legacyFile = dataFile(1L, null);
+
+        assertThat(collectWrittenColumnIds(ignored -> schema, legacyFile))
+                .hasValue(Arrays.asList(1, 2));
     }
 
     @Test
@@ -171,4 +304,38 @@ public class DataEvolutionUtilsTest {
                 0L,
                 writeCols);
     }
+
+    private static DataFileMeta dataFile(long schemaId, java.util.List<String> 
writeCols) {
+        DataFileMeta file = mock(DataFileMeta.class);
+        when(file.schemaId()).thenReturn(schemaId);
+        when(file.writeCols()).thenReturn(writeCols);
+        return file;
+    }
+
+    private static DataSplit dataSplit(DataFileMeta... files) {
+        return DataSplit.builder()
+                .withSnapshot(1L)
+                .withPartition(BinaryRow.EMPTY_ROW)
+                .withBucket(0)
+                .withBucketPath("bucket-0")
+                .withDataFiles(Arrays.asList(files))
+                .build();
+    }
+
+    private static Optional<List<Integer>> collectWrittenColumnIds(
+            Function<Long, TableSchema> schemaLoader, DataFileMeta... files) {
+        return DataEvolutionUtils.collectWrittenColumnIds(
+                Collections.singletonList(dataSplit(files)), schemaLoader);
+    }
+
+    private static TableSchema tableSchema(long id, DataField... fields) {
+        return TableSchema.create(
+                id,
+                new Schema(
+                        Arrays.asList(fields),
+                        Collections.emptyList(),
+                        Collections.emptyList(),
+                        Collections.emptyMap(),
+                        null));
+    }
 }
diff --git 
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonInputPartition.scala
 
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonInputPartition.scala
index 7e3dbf893b..a280b8e3d4 100644
--- 
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonInputPartition.scala
+++ 
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonInputPartition.scala
@@ -24,6 +24,10 @@ import org.apache.spark.sql.catalyst.InternalRow
 import org.apache.spark.sql.catalyst.expressions.GenericInternalRow
 import org.apache.spark.sql.connector.read.{HasPartitionKey, InputPartition, 
SupportsReportPartitioning}
 
+import java.util.{List => JList, Objects, Optional}
+
+import scala.util.control.NonFatal
+
 trait PaimonInputPartition extends InputPartition {
   def splits: Seq[Split]
 
@@ -36,6 +40,48 @@ trait PaimonInputPartition extends InputPartition {
 }
 
 case class SimplePaimonInputPartition(splits: Seq[Split]) extends 
PaimonInputPartition
+
+final private[spark] class PaimonMicroBatchMetadata private[spark] (
+    val sourceId: String,
+    val startOffset: String,
+    val endOffset: String,
+    val splitCount: Int,
+    @transient private var writtenColumnIdsThunk: () => 
Optional[JList[Integer]])
+  extends Serializable {
+
+  @transient private lazy val cachedWrittenColumnIds: Optional[JList[Integer]] 
= {
+    try {
+      val thunk = writtenColumnIdsThunk
+      writtenColumnIdsThunk = null
+      val supplied = if (thunk == null) null else thunk()
+      if (supplied == null) Optional.empty() else supplied
+    } catch {
+      case NonFatal(_) => Optional.empty()
+      case _: LinkageError => Optional.empty()
+    }
+  }
+
+  def writtenColumnIds: Optional[JList[Integer]] = cachedWrittenColumnIds
+
+  override def equals(other: Any): Boolean =
+    other match {
+      case that: PaimonMicroBatchMetadata =>
+        sourceId == that.sourceId &&
+        startOffset == that.startOffset &&
+        endOffset == that.endOffset &&
+        splitCount == that.splitCount
+      case _ => false
+    }
+
+  override def hashCode(): Int =
+    Objects.hash(sourceId, startOffset, endOffset, Integer.valueOf(splitCount))
+}
+
+private[spark] case class PaimonMicroBatchInputPartition(
+    splits: Seq[Split],
+    @transient metadata: PaimonMicroBatchMetadata)
+  extends PaimonInputPartition
+
 object PaimonInputPartition {
   def apply(split: Split): PaimonInputPartition = {
     SimplePaimonInputPartition(Seq(split))
diff --git 
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonSparkMicroBatchMetadata.scala
 
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonSparkMicroBatchMetadata.scala
new file mode 100644
index 0000000000..cc5b94fe6b
--- /dev/null
+++ 
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/PaimonSparkMicroBatchMetadata.scala
@@ -0,0 +1,202 @@
+/*
+ * 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.paimon.spark
+
+import org.apache.paimon.annotation.Experimental
+import org.apache.paimon.spark.sources.PaimonMicroBatchStream
+
+import org.apache.spark.Partition
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.Dataset
+import org.apache.spark.sql.connector.read.InputPartition
+import org.apache.spark.sql.execution.datasources.v2.DataSourceRDD
+
+import java.util.{IdentityHashMap, List => JList, Map => JMap, Optional, UUID}
+
+import scala.util.control.NonFatal
+
+/** Driver-side access to metadata planned for a Paimon streaming micro-batch. 
*/
+@Experimental
+final class PaimonSparkMicroBatchMetadata private ()
+
+object PaimonSparkMicroBatchMetadata {
+
+  private val StreamingQueryIdKey = "sql.streaming.queryId"
+
+  /**
+   * Returns written columns for a raw foreachBatch Dataset with exactly one 
Paimon streaming
+   * source. This method only inspects driver-side RDD planning metadata and 
does not run a Spark
+   * job. The result is empty when the Dataset is not backed by a Paimon 
source, the lineage is
+   * incomplete, or multiple Paimon sources make the result ambiguous.
+   */
+  def writtenColumnIds(batch: Dataset[_]): Optional[JList[Integer]] = {
+    try {
+      extractWrittenColumnIds(batch)
+    } catch {
+      case NonFatal(_) => Optional.empty()
+      case _: LinkageError => Optional.empty()
+    }
+  }
+
+  private def extractWrittenColumnIds(batch: Dataset[_]): 
Optional[JList[Integer]] = {
+    if (!hasExactlyOnePaimonSource(batch)) {
+      return Optional.empty()
+    }
+
+    val visited = new IdentityHashMap[RDD[_], java.lang.Boolean]()
+    var only: PaimonMicroBatchMetadata = null
+
+    def inspectOccurrence(dataSourceRDD: DataSourceRDD): Boolean = {
+      var occurrenceOnly: PaimonMicroBatchMetadata = null
+      var inputCount = 0
+      var valid = true
+      val partitions = dataSourceRDD.partitions
+      var partitionIndex = 0
+
+      while (valid && partitionIndex < partitions.length) {
+        val inputs = 
dataSourceInputPartitions(partitions(partitionIndex)).iterator
+        while (valid && inputs.hasNext) {
+          inputs.next() match {
+            case input: PaimonMicroBatchInputPartition =>
+              val current = input.metadata
+              if (current eq null) {
+                valid = false
+              } else if (occurrenceOnly eq null) {
+                occurrenceOnly = current
+                inputCount += 1
+              } else if ((occurrenceOnly eq current) || occurrenceOnly == 
current) {
+                inputCount += 1
+              } else {
+                valid = false
+              }
+            case _: PaimonInputPartition => valid = false
+            case _ =>
+          }
+        }
+        partitionIndex += 1
+      }
+
+      if (!valid || ((occurrenceOnly ne null) && inputCount != 
occurrenceOnly.splitCount)) {
+        false
+      } else if (occurrenceOnly eq null) {
+        true
+      } else if (only eq null) {
+        only = occurrenceOnly
+        true
+      } else {
+        (only eq occurrenceOnly) || only == occurrenceOnly
+      }
+    }
+
+    def visit(rdd: RDD[_]): Boolean = {
+      if (visited.containsKey(rdd)) {
+        true
+      } else {
+        visited.put(rdd, java.lang.Boolean.TRUE)
+        val valid =
+          rdd match {
+            case dataSourceRDD: DataSourceRDD => 
inspectOccurrence(dataSourceRDD)
+            case _ => true
+          }
+        if (!valid) {
+          false
+        } else {
+          val dependencies = rdd.dependencies.iterator
+          var complete = true
+          while (complete && dependencies.hasNext) {
+            complete = visit(dependencies.next().rdd)
+          }
+          complete
+        }
+      }
+    }
+
+    if (!visit(batch.queryExecution.toRdd) || (only eq null)) {
+      Optional.empty()
+    } else {
+      only.writtenColumnIds
+    }
+  }
+
+  private def dataSourceInputPartitions(partition: Partition): 
Seq[InputPartition] = {
+    if (partition == null) {
+      throw new IllegalArgumentException("Data source RDD partition must not 
be null.")
+    }
+
+    val pluralMethod =
+      try {
+        Some(partition.getClass.getMethod("inputPartitions"))
+      } catch {
+        case _: NoSuchMethodException => None
+      }
+
+    pluralMethod match {
+      case Some(method) => requireInputPartitions(method.invoke(partition))
+      case None =>
+        
Seq(requireInputPartition(partition.getClass.getMethod("inputPartition").invoke(partition)))
+    }
+  }
+
+  private def requireInputPartitions(value: Any): Seq[InputPartition] =
+    value match {
+      case null => throw new IllegalArgumentException("Input partitions must 
not be null.")
+      case values: scala.collection.Seq[_] =>
+        values.iterator.map(requireInputPartition).toVector
+      case other =>
+        throw new IllegalArgumentException(
+          s"Unexpected input partitions type ${other.getClass.getName}.")
+    }
+
+  private def requireInputPartition(value: Any): InputPartition =
+    value match {
+      case input: InputPartition => input
+      case null => throw new IllegalArgumentException("Input partition must 
not be null.")
+      case other =>
+        throw new IllegalArgumentException(
+          s"Unexpected input partition type ${other.getClass.getName}.")
+    }
+
+  private def hasExactlyOnePaimonSource(batch: Dataset[_]): Boolean = {
+    val queryId = 
batch.sparkSession.sparkContext.getLocalProperty(StreamingQueryIdKey)
+    if (queryId == null) {
+      return false
+    }
+
+    val sharedState =
+      
batch.sparkSession.getClass.getMethod("sharedState").invoke(batch.sparkSession)
+    val activeQueries =
+      sharedState.getClass
+        .getMethod("activeStreamingQueries")
+        .invoke(sharedState)
+        .asInstanceOf[JMap[UUID, AnyRef]]
+    val execution = activeQueries.get(UUID.fromString(queryId))
+    if (execution == null) {
+      return false
+    }
+
+    // Spark replaces sources without new offsets with LocalRelation before 
foreachBatch. Their
+    // RDD lineage therefore contains no InputPartition to inspect. The active 
StreamExecution is
+    // the only per-query structure which still retains every source. Keep 
this Spark-internal
+    // access isolated here and fail closed if a Spark version changes it.
+    val sources =
+      
execution.getClass.getMethod("sources").invoke(execution).asInstanceOf[Seq[AnyRef]]
+    sources.headOption.exists(
+      first => first.isInstanceOf[PaimonMicroBatchStream] && sources.forall(_ 
eq first))
+  }
+}
diff --git 
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/sources/PaimonMicroBatchStream.scala
 
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/sources/PaimonMicroBatchStream.scala
index c3d2dfc881..96883c04c6 100644
--- 
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/sources/PaimonMicroBatchStream.scala
+++ 
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/sources/PaimonMicroBatchStream.scala
@@ -20,14 +20,21 @@ package org.apache.paimon.spark.sources
 
 import org.apache.paimon.CoreOptions
 import org.apache.paimon.options.Options
-import org.apache.paimon.spark.{PaimonImplicits, PaimonInputPartition, 
PaimonPartitionReaderFactory, SparkConnectorOptions}
+import org.apache.paimon.schema.TableSchema
+import org.apache.paimon.spark.{PaimonImplicits, 
PaimonMicroBatchInputPartition, PaimonMicroBatchMetadata, 
PaimonPartitionReaderFactory, SparkConnectorOptions}
 import org.apache.paimon.table.DataTable
-import org.apache.paimon.table.source.ReadBuilder
+import org.apache.paimon.table.source.{DataSplit, ReadBuilder}
+import org.apache.paimon.utils.DataEvolutionUtils
 
 import org.apache.spark.internal.Logging
 import org.apache.spark.sql.connector.read.{InputPartition, 
PartitionReaderFactory}
 import org.apache.spark.sql.connector.read.streaming.{MicroBatchStream, 
Offset, ReadLimit, SupportsTriggerAvailableNow}
 
+import java.lang.{Long => JLong}
+import java.util.{ArrayList, Collections}
+import java.util.concurrent.ConcurrentHashMap
+import java.util.function.Function
+
 import scala.collection.mutable
 
 class PaimonMicroBatchStream(
@@ -93,6 +100,14 @@ class PaimonMicroBatchStream(
 
   private lazy val blobAsDescriptor: Boolean = 
options.get(CoreOptions.BLOB_AS_DESCRIPTOR)
 
+  private[spark] lazy val schemaLoader: Function[JLong, TableSchema] = {
+    val schemaManager = table.schemaManager()
+    val schemaCache = new ConcurrentHashMap[JLong, TableSchema]()
+    val uncachedSchemaLoader: Function[JLong, TableSchema] =
+      schemaId => schemaManager.schema(schemaId.longValue())
+    schemaId => schemaCache.computeIfAbsent(schemaId, uncachedSchemaLoader)
+  }
+
   override def getDefaultReadLimit: ReadLimit = defaultReadLimit
 
   override def prepareForTriggerAvailableNow(): Unit = {
@@ -134,11 +149,30 @@ class PaimonMicroBatchStream(
     }
     val endOffset = PaimonSourceOffset(end)
 
-    getBatch(startOffset, Some(endOffset), None)
-      .map(ids => PaimonInputPartition(ids.entry))
+    val admittedSplits = getBatch(startOffset, Some(endOffset), None)
+    val metadata = createMicroBatchMetadata(startOffset, endOffset, 
admittedSplits)
+    admittedSplits
+      .map(ids => PaimonMicroBatchInputPartition(Seq(ids.entry), metadata))
       .toArray[InputPartition]
   }
 
+  private def createMicroBatchMetadata(
+      startOffset: PaimonSourceOffset,
+      endOffset: PaimonSourceOffset,
+      admittedSplits: Array[IndexedDataSplit]): PaimonMicroBatchMetadata = {
+    val splits = new ArrayList[DataSplit](admittedSplits.length)
+    admittedSplits.foreach(split => splits.add(split.entry))
+    val admittedSplitSnapshot = Collections.unmodifiableList(splits)
+
+    new PaimonMicroBatchMetadata(
+      checkpointLocation,
+      startOffset.json(),
+      endOffset.json(),
+      admittedSplits.length,
+      () => DataEvolutionUtils.collectWrittenColumnIds(admittedSplitSnapshot, 
schemaLoader)
+    )
+  }
+
   override def createReaderFactory(): PartitionReaderFactory = {
     PaimonPartitionReaderFactory(readBuilder, blobAsDescriptor = 
blobAsDescriptor)
   }
diff --git 
a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/PaimonSourceTest.scala
 
b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/PaimonSourceTest.scala
index e8b685664c..624c2c9dc7 100644
--- 
a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/PaimonSourceTest.scala
+++ 
b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/PaimonSourceTest.scala
@@ -18,18 +18,36 @@
 
 package org.apache.paimon.spark
 
-import org.apache.paimon.spark.sources.PaimonSourceOffset
+import org.apache.paimon.schema.{SchemaManager, TableSchema}
+import org.apache.paimon.spark.sources.{PaimonMicroBatchStream, 
PaimonSourceOffset}
+import org.apache.paimon.table.DataTable
+import org.apache.paimon.utils.InstantiationUtil
 
-import org.apache.spark.sql.Row
-import org.apache.spark.sql.streaming.{StreamingQueryException, StreamTest, 
Trigger}
+import org.apache.spark.sql.{Dataset, Row}
+import org.apache.spark.sql.streaming.{StreamingQuery, 
StreamingQueryException, StreamTest, Trigger}
 import org.junit.jupiter.api.Assertions
+import org.mockito.Mockito.{mock, times, verify, when}
 
-import java.util.concurrent.TimeUnit
+import java.lang.{Long => JLong}
+import java.util.{Collections, List => JList, Optional}
+import java.util.concurrent.{atomic, TimeUnit}
+
+import scala.collection.JavaConverters._
 
 class PaimonSourceTest extends PaimonSparkTestBase with StreamTest {
 
   import testImplicits._
 
+  private def testMetadata(
+      writtenColumnIds: => Optional[JList[Integer]]): PaimonMicroBatchMetadata 
=
+    new PaimonMicroBatchMetadata("source", "start", "end", 0, () => 
writtenColumnIds)
+
+  private def withStartedQuery(query: => StreamingQuery)(body: StreamingQuery 
=> Unit): Unit = {
+    val started = query
+    try body(started)
+    finally started.stop()
+  }
+
   test("Paimon Source: EQUAL_NULL_SAFE") {
     withTempDir {
       _ =>
@@ -48,6 +66,233 @@ class PaimonSourceTest extends PaimonSparkTestBase with 
StreamTest {
     }
   }
 
+  test("Paimon Source: keep micro-batch metadata on the driver") {
+    val metadata = 
testMetadata(Optional.of(Collections.singletonList(Integer.valueOf(1))))
+    val partition = PaimonMicroBatchInputPartition(Seq.empty, metadata)
+
+    val restored = InstantiationUtil.clone(partition)
+
+    assert(restored.splits.isEmpty)
+    assert(restored.metadata == null)
+  }
+
+  test("Paimon Source: lazily memoize micro-batch written columns") {
+    val evaluations = new atomic.AtomicInteger()
+    val expected = Collections.singletonList(Integer.valueOf(1))
+    val metadata = testMetadata {
+      evaluations.incrementAndGet()
+      Optional.of(expected)
+    }
+
+    assert(evaluations.get() == 0)
+    assert(metadata.writtenColumnIds == Optional.of(expected))
+    assert(metadata.writtenColumnIds == Optional.of(expected))
+    assert(evaluations.get() == 1)
+
+    val equivalent = testMetadata {
+      throw new AssertionError("Equality must not evaluate the thunk.")
+    }
+    assert(metadata == equivalent)
+    assert(metadata.hashCode() == equivalent.hashCode())
+  }
+
+  test("Paimon Source: cache schemas for the stream lifetime") {
+    val table = mock(classOf[DataTable])
+    val schemaManager = mock(classOf[SchemaManager])
+    val initialSchema = mock(classOf[TableSchema])
+    val evolvedSchema = mock(classOf[TableSchema])
+    when(table.options()).thenReturn(Collections.emptyMap[String, String]())
+    when(table.schemaManager()).thenReturn(schemaManager)
+    when(schemaManager.schema(1L)).thenReturn(initialSchema)
+    when(schemaManager.schema(2L)).thenReturn(evolvedSchema)
+
+    val stream = new PaimonMicroBatchStream(table, null, "checkpoint")
+
+    assert(stream.schemaLoader.apply(JLong.valueOf(1L)) eq initialSchema)
+    assert(stream.schemaLoader.apply(JLong.valueOf(1L)) eq initialSchema)
+    assert(stream.schemaLoader.apply(JLong.valueOf(2L)) eq evolvedSchema)
+    assert(stream.schemaLoader.apply(JLong.valueOf(2L)) eq evolvedSchema)
+    verify(schemaManager, times(1)).schema(1L)
+    verify(schemaManager, times(1)).schema(2L)
+  }
+
+  test("Paimon Source: expose written columns to raw foreachBatch") {
+    withTempDir {
+      checkpointDir =>
+        val TableSnapshotState(_, location, snapshotData, _, _) =
+          prepareTableAndGetLocation(1, hasPk = true)
+        val expectedFieldIds =
+          loadTable("T").schema().fields().asScala.map(field => 
Integer.valueOf(field.id())).sorted
+        @volatile var writtenColumnIds: JList[Integer] = null
+        @volatile var metadataLookupStartedNoSparkJob = false
+        @volatile var rowCount = 0L
+
+        withStartedQuery(
+          spark.readStream
+            .format("paimon")
+            .load(location)
+            .select("a")
+            .writeStream
+            .option("checkpointLocation", checkpointDir.getCanonicalPath)
+            .foreachBatch {
+              (batch: Dataset[Row], _: Long) =>
+                val jobGroup = s"written-columns-metadata-${System.nanoTime()}"
+                val previousJobGroup = 
spark.sparkContext.getLocalProperty("spark.jobGroup.id")
+                spark.sparkContext.setLocalProperty("spark.jobGroup.id", 
jobGroup)
+                val metadata =
+                  try {
+                    PaimonSparkMicroBatchMetadata.writtenColumnIds(batch)
+                  } finally {
+                    metadataLookupStartedNoSparkJob =
+                      
spark.sparkContext.statusTracker.getJobIdsForGroup(jobGroup).isEmpty
+                    spark.sparkContext.setLocalProperty("spark.jobGroup.id", 
previousJobGroup)
+                  }
+                if (metadata.isPresent) {
+                  writtenColumnIds = metadata.get()
+                }
+                rowCount += batch.count()
+                ()
+            }
+            .start()) {
+          query =>
+            query.processAllAvailable()
+            assert(writtenColumnIds == expectedFieldIds.asJava)
+            assert(metadataLookupStartedNoSparkJob)
+            assert(rowCount == snapshotData.size)
+        }
+    }
+  }
+
+  test("Paimon Source: expose written columns for a self-union") {
+    withTempDir {
+      checkpointDir =>
+        val TableSnapshotState(_, location, _, _, _) =
+          prepareTableAndGetLocation(1, hasPk = true)
+        val expectedFieldIds =
+          loadTable("T").schema().fields().asScala.map(field => 
Integer.valueOf(field.id())).sorted
+        @volatile var writtenColumnIds: JList[Integer] = null
+
+        val source = spark.readStream
+          .format("paimon")
+          .load(location)
+        withStartedQuery(
+          source
+            .union(source)
+            .writeStream
+            .option("checkpointLocation", checkpointDir.getCanonicalPath)
+            .foreachBatch {
+              (batch: Dataset[Row], _: Long) =>
+                val metadata = 
PaimonSparkMicroBatchMetadata.writtenColumnIds(batch)
+                if (metadata.isPresent) {
+                  writtenColumnIds = metadata.get()
+                }
+                batch.count()
+                ()
+            }
+            .start()) {
+          query =>
+            query.processAllAvailable()
+            assert(writtenColumnIds == expectedFieldIds.asJava)
+        }
+    }
+  }
+
+  test("Paimon Source: written columns metadata is ambiguous with an empty 
second source") {
+    withTable("written_columns_source_1", "written_columns_source_2") {
+      withTempDir {
+        checkpointDir =>
+          spark.sql("CREATE TABLE written_columns_source_1 (id INT)")
+          spark.sql("CREATE TABLE written_columns_source_2 (id INT)")
+          spark.sql("INSERT INTO written_columns_source_1 VALUES (1)")
+          spark.sql("INSERT INTO written_columns_source_2 VALUES (2)")
+
+          val source1 = spark.readStream
+            .table("written_columns_source_1")
+          val source2 = spark.readStream
+            .table("written_columns_source_2")
+          @volatile var nonEmptyBatchMetadataPresent = Seq.empty[Boolean]
+
+          withStartedQuery(
+            source1
+              .union(source2)
+              .writeStream
+              .option("checkpointLocation", checkpointDir.getCanonicalPath)
+              .foreachBatch {
+                (batch: Dataset[Row], _: Long) =>
+                  val metadataPresent =
+                    
PaimonSparkMicroBatchMetadata.writtenColumnIds(batch).isPresent
+                  if (batch.count() > 0) {
+                    nonEmptyBatchMetadataPresent = 
nonEmptyBatchMetadataPresent :+ metadataPresent
+                  }
+                  ()
+              }
+              .start()) {
+            query =>
+              query.processAllAvailable()
+              nonEmptyBatchMetadataPresent = Seq.empty
+
+              spark.sql("INSERT INTO written_columns_source_1 VALUES (3)")
+              query.processAllAvailable()
+
+              assert(nonEmptyBatchMetadataPresent == Seq(false))
+          }
+      }
+    }
+  }
+
+  test("Paimon Source: expose partial data evolution written columns") {
+    withSparkSQLConf("spark.paimon.write.use-v2-write" -> "false") {
+      withTable("T") {
+        withTempDir {
+          checkpointDir =>
+            spark.sql(
+              "CREATE TABLE T (id INT, b INT, c INT) " +
+                "TBLPROPERTIES ('row-tracking.enabled' = 'true', " +
+                "'data-evolution.enabled' = 'true')")
+            spark.sql("INSERT INTO T VALUES (1, 10, 100), (2, 20, 200)")
+            val fieldIds =
+              loadTable("T")
+                .schema()
+                .fields()
+                .asScala
+                .map(field => field.name() -> field.id())
+                .toMap
+            @volatile var nonEmptyBatchColumns = Seq.empty[JList[Integer]]
+
+            withStartedQuery(
+              spark.readStream
+                .option(SparkConnectorOptions.MAX_FILES_PER_TRIGGER.key(), 1)
+                .option("scan.mode", "latest")
+                .table("`T$row_tracking`")
+                .writeStream
+                .option("checkpointLocation", checkpointDir.getCanonicalPath)
+                .foreachBatch {
+                  (batch: Dataset[Row], _: Long) =>
+                    val metadata = 
PaimonSparkMicroBatchMetadata.writtenColumnIds(batch)
+                    if (batch.count() > 0 && metadata.isPresent) {
+                      nonEmptyBatchColumns = nonEmptyBatchColumns :+ 
metadata.get()
+                    }
+                    ()
+                }
+                .start()) {
+              query =>
+                query.processAllAvailable()
+                spark.sql("UPDATE T SET b = 22 WHERE id = 2")
+                spark.sql("UPDATE T SET c = NULL WHERE id = 1")
+                query.processAllAvailable()
+
+                assert(nonEmptyBatchColumns.size >= 2)
+                val partialBatchColumns = nonEmptyBatchColumns.takeRight(2)
+                assert(
+                  partialBatchColumns == Seq(
+                    Seq(Integer.valueOf(fieldIds("b"))).asJava,
+                    Seq(Integer.valueOf(fieldIds("c"))).asJava))
+            }
+        }
+      }
+    }
+  }
+
   test("Paimon Source: default scan mode") {
     withTempDir {
       checkpointDir =>

Reply via email to