Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
15 commits
Select commit Hold shift + click to select a range
813df92
[GH-56132][STREAMING] Call pruneColumns on V2 streaming scan builders…
zikangh May 27, 2026
2e4fd61
[GH-56132][STREAMING] Address review: fix ContinuousExecution, improv…
zikangh May 27, 2026
195409a
[GH-56132][STREAMING] Address review: fix ContinuousExecution, improv…
zikangh May 27, 2026
1e184cf
[GH-56132][STREAMING] Add escape-hatch config spark.sql.streaming.v2.…
zikangh May 27, 2026
a9ec1ed
[GH-56132][STREAMING] Add E2E streaming metadata column test in Metad…
zikangh May 27, 2026
3fa3c33
[GH-56132][STREAMING] Assert AIOOBE in disabled-flag test; add flag=f…
zikangh May 27, 2026
b3e1e35
[GH-56132][STREAMING] Fix line length lint in MicroBatchExecution
zikangh May 27, 2026
11dbc06
[GH-56132][STREAMING] Fix import ordering in ContinuousExecution
zikangh May 28, 2026
8a94ac4
[GH-56132][STREAMING] Fix import ordering in ContinuousExecution (SQL…
zikangh May 28, 2026
0394c86
[GH-56132][STREAMING] Fix import order in MetadataColumnSuite, line l…
zikangh May 28, 2026
dae00fe
[GH-56132][STREAMING] Fix remaining import order in MetadataColumnSuite
zikangh May 28, 2026
f208acd
[GH-56132][STREAMING] Add bindingPolicy to config; simplify AIOOBE te…
zikangh May 28, 2026
7ce9cdf
[GH-56132][STREAMING] Fix bindingPolicy call order (before booleanConf)
zikangh May 28, 2026
cd700c6
[GH-56132][STREAMING] Fix InMemoryMicroBatchReaderFactory: append met…
zikangh May 28, 2026
7faba11
[GH-56132][STREAMING] Remove escape-hatch config; unconditionally cal…
zikangh May 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -545,7 +545,8 @@ abstract class InMemoryBaseTable(
override def json(): String = rowCount.toString
}

class InMemoryMicroBatchStream extends MicroBatchStream {
class InMemoryMicroBatchStream(readSchema: StructType, tableSchema: StructType)
extends MicroBatchStream {
override def initialOffset(): Offset = new InMemoryTableOffset(0)
override def latestOffset(): Offset =
new InMemoryTableOffset(InMemoryBaseTable.this.rows.size.toLong)
Expand All @@ -554,14 +555,13 @@ abstract class InMemoryBaseTable(
val e = end.asInstanceOf[InMemoryTableOffset].rowCount.toInt
Array(InMemoryMicroBatchPartition(InMemoryBaseTable.this.rows.slice(s, e)))
}
override def createReaderFactory(): PartitionReaderFactory = { partition =>
val rows = partition.asInstanceOf[InMemoryMicroBatchPartition].rows
new PartitionReader[InternalRow] {
private var idx = -1
override def next(): Boolean = { idx += 1; idx < rows.size }
override def get(): InternalRow = rows(idx)
override def close(): Unit = {}
override def createReaderFactory(): PartitionReaderFactory = {
val metadataColNames = new mutable.ArrayBuffer[String]()
readSchema.foreach {
case MetadataStructFieldWithLogicalName(_, name) => metadataColNames += name
case _ =>
}
new InMemoryMicroBatchReaderFactory(metadataColNames.toArray)
}
override def deserializeOffset(json: String): Offset = new InMemoryTableOffset(json.toLong)
override def commit(end: Offset): Unit = {}
Expand Down Expand Up @@ -655,7 +655,7 @@ abstract class InMemoryBaseTable(
}

override def toMicroBatchStream(checkpointLocation: String): MicroBatchStream =
new InMemoryMicroBatchStream
new InMemoryMicroBatchStream(readSchema, tableSchema)
}

case class InMemoryBatchScan(
Expand Down Expand Up @@ -954,6 +954,30 @@ class BufferedRows(val key: Seq[Any], val schema: StructType)
def clear(): Unit = rows.clear()
}

private class InMemoryMicroBatchReaderFactory(
metaNames: Array[String]) extends PartitionReaderFactory with Serializable {
override def createReader(partition: InputPartition): PartitionReader[InternalRow] = {
val rows = partition.asInstanceOf[InMemoryMicroBatchPartition].rows
new PartitionReader[InternalRow] {
private var idx = -1
override def next(): Boolean = { idx += 1; idx < rows.size }
override def get(): InternalRow = {
val rawRow = rows(idx)
if (metaNames.isEmpty) rawRow
else {
val metaRow = new GenericInternalRow(metaNames.map {
case "index" => idx.asInstanceOf[Any]
case "_partition" => UTF8String.fromString("").asInstanceOf[Any]
case _ => null
})
new JoinedRow(rawRow, metaRow)
}
}
override def close(): Unit = {}
}
}
}

object BufferedRows {
def apply(key: Seq[Any], schema: Array[Column]): BufferedRows = {
new BufferedRows(key, CatalogV2Util.v2ColumnsToStructType(schema))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ import org.apache.spark.sql.catalyst.trees.TreePattern.CURRENT_LIKE
import org.apache.spark.sql.classic.SparkSession
import org.apache.spark.sql.connector.catalog.{SupportsRead, SupportsWrite, TableCapability}
import org.apache.spark.sql.connector.distributions.UnspecifiedDistribution
import org.apache.spark.sql.connector.read.SupportsPushDownRequiredColumns
import org.apache.spark.sql.connector.read.streaming.{ContinuousStream, PartitionOffset, ReadLimit, SparkDataStream}
import org.apache.spark.sql.connector.write.{RequiresDistributionAndOrdering, Write}
import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors}
Expand Down Expand Up @@ -92,7 +93,15 @@ class ContinuousExecution(
log"from DataSourceV2 named '${MDC(STREAMING_DATA_SOURCE_NAME, sourceName)}' " +
log"${MDC(STREAMING_DATA_SOURCE_DESCRIPTION, dsStr)}")
// TODO: operator pushdown.
val scan = table.newScanBuilder(options).build()
// Passes the full output schema (not a pruned subset) so that connectors
// implementing SupportsMetadataColumns can include metadata columns in readSchema().
val scanBuilder = table.newScanBuilder(options)
scanBuilder match {
case r: SupportsPushDownRequiredColumns =>
r.pruneColumns(output.toStructType)
case _ =>
}
val scan = scanBuilder.build()
val stream = scan.toContinuousStream(metadataPath)
val relation = StreamingDataSourceV2Relation(
table, output, catalog, identifier, options, metadataPath)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ import org.apache.spark.sql.catalyst.util.truncatedString
import org.apache.spark.sql.classic.{Dataset, SparkSession}
import org.apache.spark.sql.classic.ClassicConversions.castToImpl
import org.apache.spark.sql.connector.catalog.{SupportsRead, SupportsWrite, TableCapability, TransactionalCatalogPlugin}
import org.apache.spark.sql.connector.read.SupportsPushDownRequiredColumns
import org.apache.spark.sql.connector.read.streaming.{MicroBatchStream, Offset => OffsetV2, ReadLimit, SparkDataStream, SupportsAdmissionControl, SupportsRealTimeMode, SupportsTriggerAvailableNow}
import org.apache.spark.sql.errors.QueryExecutionErrors
import org.apache.spark.sql.execution.{SparkPlan, SQLExecution}
Expand Down Expand Up @@ -224,7 +225,15 @@ class MicroBatchExecution(
log"from DataSourceV2 named '${MDC(LogKeys.STREAMING_DATA_SOURCE_NAME, srcName)}' " +
log"${MDC(LogKeys.STREAMING_DATA_SOURCE_DESCRIPTION, dsStr)}")
// TODO: operator pushdown.
val scan = table.newScanBuilder(options).build()
// Passes the full output schema (not a pruned subset) so that connectors
// implementing SupportsMetadataColumns can include metadata columns in readSchema().
val scanBuilder = table.newScanBuilder(options)
scanBuilder match {
case r: SupportsPushDownRequiredColumns =>
r.pruneColumns(output.toStructType)
case _ =>
}
val scan = scanBuilder.build()

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The batch analogue in PushDownUtils.pruneColumns rebinds the relation's output from scan.readSchema() after build() (via toOutputAttrs(scan.readSchema(), relation)) so the relation reflects what the scan actually produces. Here we keep the analyzed output and trust the connector to produce a matching readSchema(). If a connector reorders fields or silently drops an unrecognized column, downstream binding will fail with a different cryptic error that looks like the bug this PR fixes — making future regressions confusing to diagnose. Consider either adopting the batch defensive rebind, or adding a short comment explaining why the analyzed output is safe to trust here.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not an issue here because we pass the entire output into pruneColumns.

val stream = scan.toMicroBatchStream(metadataPath)
val relation = StreamingDataSourceV2Relation(
table,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -376,6 +376,32 @@ class MetadataColumnSuite extends DatasourceV2SQLBase {
}
}

test("SPARK-56132: streaming read of metadata columns from V2 source") {
withTable(tbl) {
prepareTable()
withTempDir { checkpointDir =>
// "index" is a metadata column (not in the table schema); "id" and "data" are data columns.
val df = spark.readStream.table(tbl).select("id", "data", "index")
val q = df.writeStream
.format("memory")
.queryName("result_56132")
.option("checkpointLocation", checkpointDir.getCanonicalPath)
.start()
try {
q.processAllAvailable()
val result = spark.table("result_56132")
// Verify data columns arrive correctly and index (metadata) is non-null.
checkAnswer(result.select("id", "data").orderBy("id"),
Seq(Row(1, "a"), Row(2, "b"), Row(3, "c")))
assert(result.select("index").collect().forall(!_.isNullAt(0)),
"index metadata column should be non-null in streaming output")
} finally {
q.stop()
}
}
}
}

test("SPARK-43123: Metadata column related field metadata should not be leaked to catalogs") {
withTable(tbl, "testcat.target") {
prepareTable()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,8 @@ import org.apache.spark.sql.catalyst.streaming.StreamingRelationV2
import org.apache.spark.sql.connector.{FakeV2Provider, FakeV2ProviderWithCustomSchema, InMemoryTableSessionCatalog}
import org.apache.spark.sql.connector.catalog.{Column, Identifier, InMemoryTable, InMemoryTableCatalog, MetadataColumn, SupportsMetadataColumns, SupportsRead, Table, TableCapability, TableInfo, V2TableWithV1Fallback}
import org.apache.spark.sql.connector.expressions.{ClusterByTransform, FieldReference, Transform}
import org.apache.spark.sql.connector.read.ScanBuilder
import org.apache.spark.sql.connector.read.{Scan, ScanBuilder, SupportsPushDownRequiredColumns}
import org.apache.spark.sql.connector.read.streaming.MicroBatchStream
import org.apache.spark.sql.execution.streaming.runtime.{MemoryStream, MemoryStreamScanBuilder, StreamingQueryWrapper}
import org.apache.spark.sql.functions.lit
import org.apache.spark.sql.internal.SQLConf
Expand Down Expand Up @@ -564,6 +565,42 @@ class DataStreamTableAPISuite extends StreamTest with BeforeAndAfter {
}
}

test("SPARK-56132: pruneColumns called on SupportsPushDownRequiredColumns " +
"V2 streaming scan builder") {
val tblName = "teststream.table_name"
withTable(tblName) {
spark.sql(s"CREATE TABLE $tblName (data int) USING foo")
val stream = MemoryStream[Int]
val testCatalog = spark.sessionState.catalogManager.catalog("teststream").asTableCatalog
val table = testCatalog.loadTable(Identifier.of(Array(), "table_name"))
.asInstanceOf[InMemoryStreamTable]
table.setStream(stream)

// Wrap the table's scan builder so we can record pruneColumns calls.
val recorded = new PrunedSchemaRecorder
table.scanBuilderWrapper = Some(inner => new RecordingPruneScanBuilder(inner, recorded))

withTempDir { checkpointDir =>
val q = spark.readStream.table(tblName)
.select("value", "_seq")
.writeStream.format("noop")
.option("checkpointLocation", checkpointDir.getCanonicalPath)
.start()
try {
// logicalPlan is initialized lazily when the query thread starts; wait for it.
eventually(timeout(streamingTimeout)) {
assert(recorded.called,
"pruneColumns should have been called on the streaming scan builder")
}
assert(recorded.schema.fieldNames.toSet === Set("value", "_seq"),
s"Expected pruneColumns to receive {value, _seq}, got ${recorded.schema}")
} finally {
q.stop()
}
}
}
}

private def checkForStreamTable(dir: Option[File], tableName: String): Unit = {
val memory = MemoryStream[Int]
val dsw = memory.toDS().writeStream.format("parquet")
Expand Down Expand Up @@ -683,6 +720,7 @@ class InMemoryStreamTable(override val name: String)
with SupportsRead
with SupportsMetadataColumns {
var stream: MemoryStream[Int] = _
var scanBuilderWrapper: Option[MemoryStreamScanBuilder => ScanBuilder] = None

def setStream(inputData: MemoryStream[Int]): Unit = stream = inputData

Expand All @@ -693,7 +731,8 @@ class InMemoryStreamTable(override val name: String)
}

override def newScanBuilder(options: CaseInsensitiveStringMap): ScanBuilder = {
new MemoryStreamScanBuilder(stream)
val inner = new MemoryStreamScanBuilder(stream)
scanBuilderWrapper.map(_(inner)).getOrElse(inner)
}

private object SeqColumn extends MetadataColumn {
Expand All @@ -705,6 +744,36 @@ class InMemoryStreamTable(override val name: String)
override val metadataColumns: Array[MetadataColumn] = Array(SeqColumn)
}

class PrunedSchemaRecorder {
@volatile var called = false
@volatile var schema: StructType = new StructType()
}

class RecordingPruneScanBuilder(inner: MemoryStreamScanBuilder, recorder: PrunedSchemaRecorder)
extends ScanBuilder
with SupportsPushDownRequiredColumns {

override def pruneColumns(requiredSchema: StructType): Unit = {
recorder.called = true
recorder.schema = requiredSchema
}

override def build(): Scan = {
val innerScan = inner.build()
val prunedSchema = recorder.schema
// Return a scan whose readSchema() reflects the pruned schema so the streaming plan
// and scan agree on output columns. Without the fix, pruneColumns is never called and
// readSchema() defaults to the full table schema, causing ArrayIndexOutOfBoundsException
// when metadata columns are in the plan output but absent from the scan output.
new Scan {
override def readSchema(): StructType =
if (recorder.called) prunedSchema else innerScan.readSchema()
override def toMicroBatchStream(checkpointLocation: String): MicroBatchStream =
innerScan.toMicroBatchStream(checkpointLocation)
}
}
}

class NonStreamV2Table(override val name: String)
extends Table with SupportsRead with V2TableWithV1Fallback {
override def schema(): StructType = StructType(Nil)
Expand Down