From ec2bec49be4e7f9e4f813d8a605a599a365a6ad1 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Thu, 3 Sep 2026 19:22:39 +0000 Subject: [PATCH 1/3] fix(spark)!: preserve partition values in file scans A partitioned Parquet read containing rows (1,10) and (2,20) round-trips as (1,NULL) and (2,NULL). The exporter lists leaf files with the full table schema, while the importer treats virtual partition columns as physical file columns. Export each selected partition as a scan of its data schema with typed partition literals, preserve the merged column order through emit, and combine partitions with UNION_ALL. This uses Spark's resolved values rather than reconstructing them from paths, including nulls, explicit types, overlapping columns, pruned indexes, and multiple roots. Serialize canonical file URIs and decode them as URIs so escaped partition values address the original files. Local paths with unescaped spaces remain accepted. The representation uses standard Substrait v0.102.0 relations and adds one scan/project branch per nonempty partition. BREAKING CHANGE: Spark 3.4 conversion now rejects partitioned file scans whose file and partition column names differ only by case under case-insensitive resolution. Resolve the schema collision before conversion; this prevents changing the source plan's values to the different partition values. --- .../substrait/spark/compat/SparkCompat.scala | 3 + .../spark/logical/ToLogicalPlan.scala | 13 +- .../spark/logical/ToSubstraitRel.scala | 95 +++++++- .../spark/compat/SparkCompatImpl.scala | 2 + .../spark/PartitionedFilesSuite.scala | 226 ++++++++++++++++++ 5 files changed, 334 insertions(+), 5 deletions(-) create mode 100644 spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala diff --git a/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala b/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala index 418234671..49f028dc9 100644 --- a/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala +++ b/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala @@ -10,6 +10,9 @@ import org.apache.spark.sql.execution.datasources.{HadoopFsRelation, LogicalRela */ trait SparkCompat { + /** Whether partition values override differently cased file columns in case-insensitive mode. */ + def supportsCaseInsensitivePartitionOverlap: Boolean = true + /** Create a LogicalRelation with version-appropriate constructor */ def createLogicalRelation( relation: HadoopFsRelation, diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala index c3fa30dba..1c8734983 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala @@ -53,7 +53,7 @@ import io.substrait.relation.physical.{BroadcastExchange, MultiBucketExchange, R import io.substrait.util.EmptyVisitationContext import org.apache.hadoop.fs.Path -import java.net.URI +import java.net.{URI, URISyntaxException} import java.util.Optional import scala.annotation.nowarn @@ -440,7 +440,7 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes val (format, options) = convertFileFormat(formats.head) val location = SparkCompat.instance.createInMemoryFileIndex( spark, - localFiles.getItems.asScala.map(i => new Path(i.getPath.get())).toSeq, + localFiles.getItems.asScala.map(i => toFilePath(i.getPath.get())).toSeq, Map(), Some(schema)) val hadoopFsRelation = SparkCompat.instance.createHadoopFsRelation( @@ -461,6 +461,15 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes remap(plan, localFiles.getRemap) } + private def toFilePath(path: String): Path = { + try { + new Path(new URI(path)) + } catch { + // Preserve support for unescaped local paths, such as filenames containing spaces. + case _: URISyntaxException => new Path(path) + } + } + def convertFileFormat(fileFormat: FileFormat): (SparkFileFormat, Map[String, String]) = { fileFormat match { case csv: FileFormat.DelimiterSeparatedTextReadOptions => diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala index 90f4a3538..6e9543c73 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala @@ -17,14 +17,14 @@ package io.substrait.spark.logical import io.substrait.spark.{FileHolder, SparkExtension, ToSubstraitType} -import io.substrait.spark.compat.WindowGroupLimitCase +import io.substrait.spark.compat.{SparkCompat, WindowGroupLimitCase} import io.substrait.spark.expression._ import io.substrait.spark.utils.Util import org.apache.spark.internal.Logging import org.apache.spark.sql.SaveMode import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.analysis.ResolvedIdentifier +import org.apache.spark.sql.catalyst.analysis.{caseInsensitiveResolution, caseSensitiveResolution, ResolvedIdentifier} import org.apache.spark.sql.catalyst.catalog.{CatalogTable, HiveTableRelation} import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, Average, Sum} @@ -38,6 +38,7 @@ import org.apache.spark.sql.execution.datasources.orc.OrcFileFormat import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat import org.apache.spark.sql.execution.datasources.v2.{DataSourceV2Relation, DataSourceV2ScanRelation, V2SessionCatalog} import org.apache.spark.sql.hive.execution.{CreateHiveTableAsSelectCommand, InsertIntoHiveTable} +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{NullType, StructField, StructType} import io.substrait.`type`.{NamedStruct, Type} @@ -511,6 +512,92 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { .build() } + private def buildPartitionedFileScan(fsRelation: HadoopFsRelation): relation.Rel = { + // LocalFiles does not carry directory partition values. Encode Spark's resolved values as + // literals so consumers do not have to infer their types or base paths again. + val fileFormat = convertFileFormat(fsRelation.fileFormat, fsRelation.options) + val dataSchema = ToSubstraitType.toNamedStruct(fsRelation.dataSchema) + val resolver = + if ( + SparkCompat.instance.getConf(fsRelation.sparkSession, SQLConf.CASE_SENSITIVE.key).toBoolean + ) { + caseSensitiveResolution + } else { + caseInsensitiveResolution + } + + if ( + !SparkCompat.instance.supportsCaseInsensitivePartitionOverlap && + fsRelation.dataSchema.exists( + data => + fsRelation.partitionSchema.exists( + partition => data.name != partition.name && resolver(data.name, partition.name))) + ) { + throw new UnsupportedOperationException( + "This Spark version cannot reliably read file and partition columns that differ only in case") + } + + // Partition values override overlapping file columns, preserving the merged schema's order. + val outputMapping = fsRelation.schema.map { + field => + val partitionIndex = + fsRelation.partitionSchema.indexWhere(p => resolver(p.name, field.name)) + if (partitionIndex >= 0) { + fsRelation.dataSchema.size + partitionIndex + } else { + fsRelation.dataSchema.indexWhere(d => resolver(d.name, field.name)) + } + } + val remap = relation.Rel.Remap.of(outputMapping.map(Int.box).toSeq.asJava) + + val partitions = fsRelation.location.listFiles(Nil, Nil).filter(_.files.nonEmpty).map { + partition => + val read = relation.LocalFiles + .builder() + .initialSchema(dataSchema) + .addAllItems( + partition.files + .map { + file => + FileOrFiles + .builder() + .fileFormat(fileFormat) + .partitionIndex(0) + .start(0) + .length(file.getLen) + .path(file.getPath.toUri.toString) + .pathType(PathType.URI_FILE) + .build() + } + .toSeq + .asJava) + .build() + val values = fsRelation.partitionSchema.zipWithIndex.map { + case (field, index) => + ToSubstraitLiteral( + Literal(partition.values.get(index, field.dataType), field.dataType), + Some(field.nullable)) + } + relation.Project + .builder() + .input(read) + .addAllExpressions(values.toSeq.asJava) + .remap(remap) + .build() + } + + partitions.size match { + case 0 => + relation.VirtualTableScan + .builder() + .initialSchema(ToSubstraitType.toNamedStruct(fsRelation.schema)) + .build() + case 1 => partitions.head + case _ => + relation.Set.builder().setOp(SetOp.UNION_ALL).addAllInputs(partitions.toSeq.asJava).build() + } + } + private def convertFileFormat( fileFormat: DSFileFormat, options: Map[String, String]): FileFormat = fileFormat match { @@ -533,7 +620,7 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { } /** Read Operator: https://substrait.io/relations/logical_relations/#read-operator */ - private def convertReadOperator(plan: LeafNode): relation.AbstractReadRel = { + private def convertReadOperator(plan: LeafNode): relation.Rel = { var tableNames: List[String] = null plan match { case logicalRelation: LogicalRelation if logicalRelation.catalogTable.isDefined => @@ -558,6 +645,8 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { buildVirtualTableScan(rdd.schema, rdd.rdd.take(_rddLimit).toIndexedSeq) case logicalRelation: LogicalRelation => logicalRelation.relation match { + case fsRelation: HadoopFsRelation if fsRelation.partitionSchema.nonEmpty => + buildPartitionedFileScan(fsRelation) case fsRelation: HadoopFsRelation => buildLocalFileScan(fsRelation) case _ => diff --git a/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala b/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala index 022882eee..9c41738f1 100644 --- a/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala +++ b/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala @@ -8,6 +8,8 @@ import io.substrait.relation class SparkCompatImpl extends SparkCompat { + override def supportsCaseInsensitivePartitionOverlap: Boolean = false + override def createLogicalRelation( relation: HadoopFsRelation, output: Seq[AttributeReference], diff --git a/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala new file mode 100644 index 000000000..8a916fff3 --- /dev/null +++ b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala @@ -0,0 +1,226 @@ +package io.substrait.spark + +import io.substrait.spark.compat.SparkCompat +import io.substrait.spark.logical.{ToLogicalPlan, ToSubstraitRel} + +import org.apache.spark.sql.Row +import org.apache.spark.sql.catalyst.analysis.caseSensitiveResolution +import org.apache.spark.sql.catalyst.expressions.Expression +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan +import org.apache.spark.sql.classic.DatasetUtil +import org.apache.spark.sql.execution.datasources.{FileIndex, HadoopFsRelation, LogicalRelation, PartitionDirectory} +import org.apache.spark.sql.test.SharedSparkSession +import org.apache.spark.sql.types.{DataType, IntegerType, LongType, StringType, StructField, StructType} + +import io.substrait.plan.{PlanProtoConverter, ProtoPlanConverter} +import io.substrait.relation.{LocalFiles => SubstraitLocalFiles} +import io.substrait.relation.files.FileOrFiles +import org.apache.hadoop.fs.Path + +import java.net.URI +import java.time.LocalDate + +import scala.jdk.CollectionConverters._ + +class PartitionedFilesSuite extends SharedSparkSession { + + private def assertRoundTrip(plan: LogicalPlan, expected: Seq[Row]): Unit = { + val original = DatasetUtil.fromLogicalPlan(spark, plan).collect().toSeq + assertResult(expected.sortBy(_.toString))(original.sortBy(_.toString)) + + val substrait = new ToSubstraitRel().convert(plan) + val bytes = new PlanProtoConverter().toProto(substrait).toByteArray + val decoded = new ProtoPlanConverter().from(io.substrait.proto.Plan.parseFrom(bytes)) + assertResult(substrait)(decoded) + + val converted = new ToLogicalPlan(spark).convert(decoded) + assert( + DataType.equalsStructurallyByName(plan.schema, converted.schema, caseSensitiveResolution)) + val actual = DatasetUtil.fromLogicalPlan(spark, converted).collect().toSeq + assertResult(expected.sortBy(_.toString))(actual.sortBy(_.toString)) + } + + Seq("parquet", "orc", "csv").foreach { + format => + test(s"partition values survive $format reads and file options") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark + .sql("select 1 id, 'left|right' value, 10 part union all select 2, 'other', 20") + .write + .format(format) + .option("header", true) + .option("delimiter", "|") + .partitionBy("part") + .save(path) + val schema = StructType( + Seq( + StructField("id", IntegerType), + StructField("value", StringType), + StructField("part", IntegerType))) + val data = spark.read + .format(format) + .schema(schema) + .option("header", true) + .option("delimiter", "|") + .load(path) + assertRoundTrip( + data.queryExecution.optimizedPlan, + Seq(Row(1, "left|right", 10), Row(2, "other", 20))) + } + } + } + + test("multiple roots and basePath preserve the selected partition values") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark + .sql("select 1 id, 10 part union all select 2, 20 union all select 3, 30") + .write + .partitionBy("part") + .parquet(path) + val selected = spark.read + .option("basePath", path) + .parquet(s"$path/part=10", s"$path/part=20") + assertRoundTrip(selected.queryExecution.optimizedPlan, Seq(Row(1, 10), Row(2, 20))) + assertRoundTrip(selected.filter("part = 10").queryExecution.optimizedPlan, Seq(Row(1, 10))) + } + } + + test("date null and escaped string partition values retain their types and values") { + withSQLConf("spark.sql.datetime.java8API.enabled" -> "true") { + withTempPath { + directory => + val path = directory.getAbsolutePath + "/root with spaces" + spark + .sql("select 1 id, date '2024-01-02' day, 'a/b% c' label " + + "union all select 2, cast(null as date), cast(null as string)") + .write + .partitionBy("day", "label") + .parquet(path) + val data = spark.read.parquet(path) + assertRoundTrip( + data.queryExecution.optimizedPlan, + Seq(Row(1, LocalDate.of(2024, 1, 2), "a/b% c"), Row(2, null, null))) + } + } + } + + test("local file reads still accept unescaped paths containing spaces") { + withTempPath { + directory => + val path = directory.getAbsolutePath + "/root with spaces" + spark.sql("select 1 id").write.parquet(path) + val original = spark.read.parquet(path).queryExecution.optimizedPlan + val scan = new ToSubstraitRel().visit(original).asInstanceOf[SubstraitLocalFiles] + val files = scan.getItems.asScala.map { + file => FileOrFiles.builder().from(file).path(new URI(file.getPath.get()).getPath).build() + } + val rawPaths = SubstraitLocalFiles.builder().from(scan).items(files.toSeq.asJava).build() + val converted = new ToLogicalPlan(spark).convert(rawPaths) + assertResult(Seq(Row(1)))(DatasetUtil.fromLogicalPlan(spark, converted).collect().toSeq) + } + } + + test("explicit partition types are retained") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id, 10 part").write.partitionBy("part").parquet(path) + val schema = StructType(Seq(StructField("id", IntegerType), StructField("part", LongType))) + val data = spark.read.schema(schema).option("basePath", path).parquet(s"$path/part=10") + assertRoundTrip(data.queryExecution.optimizedPlan, Seq(Row(1, 10L))) + } + } + + test("partition values override overlapping file columns in merged schema order") { + withSQLConf("spark.sql.caseSensitive" -> "false") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id, '999' p, 'physical' value").write.parquet(s"$path/p=10") + val data = spark.read.parquet(path) + assertResult(Seq("id", "p", "value"))(data.columns.toSeq) + assertRoundTrip(data.queryExecution.optimizedPlan, Seq(Row(1, 10, "physical"))) + } + } + } + + test("mixed-case overlapping columns are rejected when Spark does not override them") { + withSQLConf("spark.sql.caseSensitive" -> "false") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id, 999 P, 'physical' value").write.parquet(s"$path/p=10") + val data = spark.read.parquet(path) + val plan = data.queryExecution.optimizedPlan + if (SparkCompat.instance.supportsCaseInsensitivePartitionOverlap) { + assertRoundTrip(plan, Seq(Row(1, 10, "physical"))) + } else { + assertResult(Seq(Row(1, 999, "physical")))(data.collect().toSeq) + val error = intercept[UnsupportedOperationException] { + new ToSubstraitRel().convert(plan) + } + assert(error.getMessage.contains("differ only in case")) + } + } + } + } + + test("case-sensitive file and partition column names remain distinct") { + withSQLConf("spark.sql.caseSensitive" -> "true") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id, 999 P, 'physical' value").write.parquet(s"$path/p=10") + val data = spark.read.parquet(path) + assertRoundTrip(data.queryExecution.optimizedPlan, Seq(Row(1, 999, "physical", 10))) + } + } + } + + private def withPartitions( + original: HadoopFsRelation, + partitions: Seq[PartitionDirectory]): LogicalPlan = { + val index = new FileIndex { + override def rootPaths: Seq[Path] = original.location.rootPaths + override def listFiles( + partitionFilters: Seq[Expression], + dataFilters: Seq[Expression]): Seq[PartitionDirectory] = partitions + override def inputFiles: Array[String] = + partitions.flatMap(_.files.map(_.getPath.toString)).toArray + override def refresh(): Unit = () + override def sizeInBytes: Long = partitions.flatMap(_.files.map(_.getLen)).sum + override def partitionSchema: StructType = original.partitionSchema + } + val relation = original.copy(location = index)(spark) + SparkCompat.instance.createLogicalRelation( + relation, + ToSparkType.toAttributeSeq(ToSubstraitType.toNamedStruct(relation.schema)), + None, + false) + } + + test("pruned and empty file indexes do not restore excluded partitions") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark + .sql("select 1 id, 10 part union all select 2, 20") + .write + .partitionBy("part") + .parquet(path) + val logical = spark.read + .parquet(path) + .queryExecution + .optimizedPlan + .asInstanceOf[LogicalRelation] + val original = logical.relation.asInstanceOf[HadoopFsRelation] + val selected = original.location.listFiles(Nil, Nil).filter(_.values.getInt(0) == 10) + assertRoundTrip(withPartitions(original, selected), Seq(Row(1, 10))) + assertRoundTrip(withPartitions(original, Seq.empty), Seq.empty) + } + } +} From df5025037040919a7773b511a5717cafa050f37b Mon Sep 17 00:00:00 2001 From: bvolpato Date: Sat, 3 Oct 2026 17:57:47 -0400 Subject: [PATCH 2/3] fix(spark): refine partition scan conversion Coalesce compatible imported scans and prune direct partition predicates during export. Preserve captured schema positions, path semantics, names, and decimal fitting while rejecting unsafe Spark 3.4 overlap filters. --- .../substrait/spark/compat/SparkCompat.scala | 7 +- .../spark/logical/ToLogicalPlan.scala | 160 +++++++++++- .../spark/logical/ToSubstraitRel.scala | 232 +++++++++++------- .../spark/compat/SparkCompatImpl.scala | 8 +- .../spark/compat/SparkCompatImpl.scala | 8 +- .../spark/compat/SparkCompatImpl.scala | 8 +- .../spark/PartitionedFilesSuite.scala | 195 ++++++++++++++- 7 files changed, 505 insertions(+), 113 deletions(-) diff --git a/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala b/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala index 49f028dc9..36db657b0 100644 --- a/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala +++ b/spark/src/main/scala/io/substrait/spark/compat/SparkCompat.scala @@ -10,6 +10,11 @@ import org.apache.spark.sql.execution.datasources.{HadoopFsRelation, LogicalRela */ trait SparkCompat { + def createPartitionDirectory( + values: org.apache.spark.sql.catalyst.InternalRow, + files: Seq[org.apache.hadoop.fs.FileStatus] + ): org.apache.spark.sql.execution.datasources.PartitionDirectory + /** Whether partition values override differently cased file columns in case-insensitive mode. */ def supportsCaseInsensitivePartitionOverlap: Boolean = true @@ -47,7 +52,7 @@ trait SparkCompat { def createHadoopFsRelation( spark: AnyRef, - location: org.apache.spark.sql.execution.datasources.InMemoryFileIndex, + location: org.apache.spark.sql.execution.datasources.FileIndex, partitionSchema: org.apache.spark.sql.types.StructType, dataSchema: org.apache.spark.sql.types.StructType, bucketSpec: Option[org.apache.spark.sql.catalyst.catalog.BucketSpec], diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala index 1c8734983..fd4e41545 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala @@ -30,13 +30,13 @@ import org.apache.spark.sql.catalyst.plans.{FullOuter, Inner, LeftAnti, LeftOute import org.apache.spark.sql.catalyst.plans.logical._ import org.apache.spark.sql.catalyst.util.toPrettySQL import org.apache.spark.sql.execution.command.{CreateDataSourceTableAsSelectCommand, CreateTableCommand, DataWritingCommand, DropTableCommand, LeafRunnableCommand} -import org.apache.spark.sql.execution.datasources.{FileFormat => SparkFileFormat, InsertIntoHadoopFsRelationCommand, V1Writes} +import org.apache.spark.sql.execution.datasources.{FileFormat => SparkFileFormat, FileIndex, InsertIntoHadoopFsRelationCommand, PartitionDirectory, V1Writes} import org.apache.spark.sql.execution.datasources.csv.CSVFileFormat import org.apache.spark.sql.execution.datasources.orc.OrcFileFormat import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat import org.apache.spark.sql.hive.execution.{CreateHiveTableAsSelectCommand, InsertIntoHiveTable} import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} -import org.apache.spark.sql.types.{DataType, IntegerType, StructField, StructType} +import org.apache.spark.sql.types.{ArrayType, DataType, IntegerType, MapType, StructField, StructType} import io.substrait.`type`.{NamedStruct, StringTypeVisitor, Type} import io.substrait.{expression => exp} @@ -49,6 +49,7 @@ import io.substrait.relation.AbstractWriteRel.{CreateMode, WriteOp} import io.substrait.relation.Expand.{ConsistentField, SwitchingField} import io.substrait.relation.Set.SetOp import io.substrait.relation.files.FileFormat +import io.substrait.relation.files.FileOrFiles.PathType import io.substrait.relation.physical.{BroadcastExchange, MultiBucketExchange, RoundRobinExchange, ScatterExchange, SingleBucketExchange} import io.substrait.util.EmptyVisitationContext import org.apache.hadoop.fs.Path @@ -303,11 +304,9 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes private def fieldNames(rel: relation.Rel): Option[Seq[String]] = { if (rel.getHint.isPresent && !rel.getHint.get().getOutputNames.isEmpty) { Some( - ToSubstraitType - .toNamedStruct(ToSparkType.toStructType( - NamedStruct.of(rel.getHint.get.getOutputNames, rel.getRecordType))) - .names - .asScala + ToSparkType + .toStructType(NamedStruct.of(rel.getHint.get.getOutputNames, rel.getRecordType)) + .fieldNames .toSeq) } else { None @@ -340,7 +339,12 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes } else { allExpressions } - Project(remapped, child) + val named = if (names.size == remapped.size) { + remapped.zip(names).map { case (expr, name) => Alias(expr, name)() } + } else { + remapped + } + Project(named, child) } else { val aggregate: Aggregate = child.asInstanceOf[Aggregate] aggregate.copy(aggregateExpressions = projectList) @@ -387,6 +391,22 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes } override def visit(set: relation.Set, context: EmptyVisitationContext): LogicalPlan = { + def finish(plan: LogicalPlan): LogicalPlan = { + val remapped = remap(plan, set.getRemap) + fieldNames(set) match { + case Some(names) if names.size == remapped.output.size => + Project( + remapped.output.zip(names).map { case (attribute, name) => Alias(attribute, name)() }, + remapped) + case _ => remapped + } + } + if (set.getSetOp == SetOp.UNION_ALL) { + combinePartitionScans(set.getInputs.asScala.toSeq, context) match { + case Some(plan) => return finish(plan) + case None => + } + } val children = set.getInputs.asScala.map(_.accept(this, context)).toSeq withOutput(children.flatMap(_.output)) { val plan = set.getSetOp match { @@ -394,8 +414,126 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes case op => throw new UnsupportedOperationException(s"Operation not currently supported: $op") } - remap(plan, set.getRemap) + finish(plan) + } + } + + private def combinePartitionScans( + inputs: Seq[relation.Rel], + context: EmptyVisitationContext): Option[LogicalPlan] = { + val branches = inputs.map { + case project: relation.Project + if project.getExpressions.asScala.forall(_.isInstanceOf[SExpression.Literal]) => + project.getInput match { + case read: LocalFiles + if read.getRemap.isEmpty && read.getFilter.isEmpty && + read.getProjection.isEmpty && read.getItems.asScala.forall( + item => + item.pathType.orElse( + null) == PathType.URI_FILE && item.getPath.isPresent && item.getStart == 0) => + Some((project, read)) + case _ => None + } + case _ => None + } + if (branches.isEmpty || branches.exists(_.isEmpty)) return None + val scans = branches.flatten + val (firstProject, firstRead) = scans.head + val types = firstProject.getExpressions.asScala.map(_.getType).toSeq + if ( + !scans.forall { + case (project, read) => + read.getInitialSchema == firstRead.getInitialSchema && + project.getRemap == firstProject.getRemap && fieldNames(project) == fieldNames( + firstProject) && + project.getExpressions.asScala.map(_.getType).toSeq == types + } + ) return None + val items = scans.flatMap(_._2.getItems.asScala) + val formats = items.map(_.getFileFormat).distinct + if (items.isEmpty || formats.size != 1 || formats.head.isEmpty) return None + val dataSchema = ToSparkType.toStructType(firstRead.getInitialSchema) + // Internal names must not overlap file columns: the project's emit supplies the final order. + var prefix = "__substrait_partition_" + while (dataSchema.fieldNames.exists(_.toLowerCase(java.util.Locale.ROOT).startsWith(prefix))) { + prefix = "_" + prefix } + val literals = scans.map { + case (project, _) => + project.getExpressions.asScala + .map(_.accept(expressionConverter, context).asInstanceOf[Literal]) + .toSeq + } + if ( + literals.head.exists(_.dataType match { + case _: ArrayType | _: MapType | _: StructType => true + case _ => false + }) + ) return None + val partitionType = StructType(literals.head.zipWithIndex.map { + case (literal, index) => + StructField(s"$prefix$index", literal.dataType, types(index).nullable()) + }) + val paths = items.map(item => toFilePath(item.getPath.get())) + val listing = + SparkCompat.instance.createInMemoryFileIndex(spark, paths, Map(), Some(dataSchema)) + val files = listing.allFiles().map(file => file.getPath -> file).toMap + if (!paths.forall(files.contains)) return None + val partitions = scans.zip(literals).map { + case ((_, read), values) => + SparkCompat.instance.createPartitionDirectory( + InternalRow.fromSeq(values.map(_.value)), + read.getItems.asScala.map(item => files(toFilePath(item.getPath.get()))).toSeq) + } + val index = new FileIndex { + override def rootPaths: Seq[Path] = paths + override def inputFiles: Array[String] = paths.map(_.toUri.toString).toArray + override def refresh(): Unit = listing.refresh() + override def sizeInBytes: Long = partitions.flatMap(_.files.map(_.getLen)).sum + override def partitionSchema: StructType = partitionType + override def listFiles( + partitionFilters: Seq[Expression], + dataFilters: Seq[Expression]): Seq[PartitionDirectory] = { + if (partitionFilters.isEmpty) partitions + else { + val bound = partitionFilters.reduceLeft(And).transform { + case attribute: AttributeReference => + BoundReference( + partitionSchema.fieldIndex(attribute.name), + attribute.dataType, + attribute.nullable) + } + val predicate = Predicate.createInterpreted(bound) + partitions.filter(partition => predicate.eval(partition.values)) + } + } + } + val (format, options) = convertFileFormat(formats.head.get()) + val fsRelation = SparkCompat.instance.createHadoopFsRelation( + spark, + index, + partitionType, + dataSchema, + None, + format, + options) + val output = fsRelation.schema.map( + field => AttributeReference(field.name, field.dataType, field.nullable, field.metadata)()) + val scan = SparkCompat.instance.createLogicalRelation(fsRelation, output, None, false) + val projected = remap(scan, firstProject.getRemap) + val hintNames = fieldNames(firstProject) + val expressionNames = hintNames + .filter(_.size == literals.head.size) + .getOrElse(literals.head.map(toPrettySQL)) + val appendedNames = dataSchema.fieldNames.toSeq ++ expressionNames + val remappedNames = if (firstProject.getRemap.isPresent) { + firstProject.getRemap.get().indices().asScala.map(appendedNames(_)).toSeq + } else appendedNames + val names = hintNames.filter(_.size == projected.output.size).getOrElse(remappedNames) + Some( + Project( + projected.output.zip(names).map { case (attribute, name) => Alias(attribute, name)() }, + projected)) } override def visit( @@ -463,7 +601,9 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes private def toFilePath(path: String): Path = { try { - new Path(new URI(path)) + val uri = new URI(path) + if (uri.getScheme == null) new Path(path) + else new Path(uri.getScheme, uri.getAuthority, uri.getPath) } catch { // Preserve support for unescaped local paths, such as filenames containing spaces. case _: URISyntaxException => new Path(path) diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala index 6e9543c73..ae1fc26e3 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala @@ -24,7 +24,7 @@ import io.substrait.spark.utils.Util import org.apache.spark.internal.Logging import org.apache.spark.sql.SaveMode import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.analysis.{caseInsensitiveResolution, caseSensitiveResolution, ResolvedIdentifier} +import org.apache.spark.sql.catalyst.analysis.ResolvedIdentifier import org.apache.spark.sql.catalyst.catalog.{CatalogTable, HiveTableRelation} import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, Average, Sum} @@ -38,8 +38,7 @@ import org.apache.spark.sql.execution.datasources.orc.OrcFileFormat import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat import org.apache.spark.sql.execution.datasources.v2.{DataSourceV2Relation, DataSourceV2ScanRelation, V2SessionCatalog} import org.apache.spark.sql.hive.execution.{CreateHiveTableAsSelectCommand, InsertIntoHiveTable} -import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{NullType, StructField, StructType} +import org.apache.spark.sql.types.{Decimal, DecimalType, NullType, StructField, StructType} import io.substrait.`type`.{NamedStruct, Type} import io.substrait.{proto, relation} @@ -57,6 +56,7 @@ import io.substrait.relation.Set.SetOp import io.substrait.relation.files.{FileFormat, FileOrFiles} import io.substrait.relation.files.FileOrFiles.PathType import io.substrait.util.EmptyVisitationContext +import org.apache.hadoop.fs.Path import java.util import java.util.{Collections, Optional} @@ -65,7 +65,7 @@ import scala.collection.mutable import scala.collection.mutable.ArrayBuffer import scala.jdk.CollectionConverters._ -class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { +class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging with PredicateHelper { private val toSubstraitExp = new WithLogicalSubQuery(this) @@ -312,11 +312,76 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { } override def visitFilter(p: Filter): relation.Rel = { - val input = visit(p.child) + checkPartitionFilter(p) + val input = p.child match { + case logical: LogicalRelation if logical.catalogTable.isEmpty => + logical.relation match { + case fs: HadoopFsRelation if fs.partitionSchema.nonEmpty => + val mapping = partitionOutputMapping(fs) + val partitionColumns = AttributeSet(logical.output.zip(mapping).collect { + case (attribute, index) if index >= fs.dataSchema.size => attribute + }) + val filters = splitConjunctivePredicates(p.condition).filter { + filter => + filter.deterministic && !SubqueryExpression.hasSubquery(filter) && + filter.references.nonEmpty && filter.references.subsetOf(partitionColumns) + } + buildPartitionedFileScan(fs, filters) + case _ => visit(p.child) + } + case _ => visit(p.child) + } val condition = toExpression(p.child.output)(p.condition) relation.Filter.builder().condition(condition).input(input).build() } + private def checkPartitionFilter(filter: Filter): Unit = { + if (SparkCompat.instance.supportsCaseInsensitivePartitionOverlap) return + val source = filter + .collect { + case logical: LogicalRelation if logical.catalogTable.isEmpty => + logical.relation + } + .collectFirst { + case fs: HadoopFsRelation if fs.dataSchema.zip(fs.schema).exists { + case (data, merged) => data.name != merged.name + } => + fs + } + if (source.isEmpty) return + // Inspect Spark's optimized plan so aliases, inferred predicates, and pushdown barriers + // follow the source engine's rules. The original plan is still used for conversion. + val optimized = + SparkCompat.instance.createQueryExecution(source.get.sparkSession, filter).optimizedPlan + optimized.foreach { + case Filter(condition, logical: LogicalRelation) if logical.catalogTable.isEmpty => + logical.relation match { + case fs: HadoopFsRelation => + val mapping = partitionOutputMapping(fs) + val overlap = AttributeSet(logical.output.zipWithIndex.collect { + case (attribute, index) + if index < fs.dataSchema.size && + fs.schema(index).name != fs.dataSchema(index).name && mapping(index) == index => + attribute + }) + val partitionColumns = AttributeSet(logical.output.zip(fs.schema).collect { + case (attribute, field) if fs.partitionSchema.contains(field) => attribute + }) + if ( + splitConjunctivePredicates(condition) + .filter(_.deterministic) + .flatMap(extractPredicatesWithinOutputSet(_, partitionColumns)) + .exists(_.references.intersect(overlap).nonEmpty) + ) { + throw new UnsupportedOperationException( + "This Spark version filters differently cased overlapping partition columns using values it does not output") + } + case _ => + } + case _ => + } + } + private def toSubstraitJoin(joinType: JoinType): relation.Join.JoinType = joinType match { case Inner | Cross => relation.Join.JoinType.INNER case LeftOuter => relation.Join.JoinType.LEFT @@ -489,104 +554,90 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { } private def buildLocalFileScan(fsRelation: HadoopFsRelation): relation.AbstractReadRel = { + buildLocalFileScan( + ToSubstraitType.toNamedStruct(fsRelation.schema), + fsRelation.location.listFiles(Nil, Nil).flatMap(_.files.map(f => (f.getPath, f.getLen))), + convertFileFormat(fsRelation.fileFormat, fsRelation.options) + ) + } + + private def buildLocalFileScan( + schema: NamedStruct, + files: Seq[(Path, Long)], + format: FileFormat): relation.LocalFiles = { relation.LocalFiles .builder() - .initialSchema(ToSubstraitType.toNamedStruct(fsRelation.schema)) + .initialSchema(schema) .addAllItems( - fsRelation.location.inputFiles - .map( - file => { - FileOrFiles - .builder() - .fileFormat(convertFileFormat(fsRelation.fileFormat, fsRelation.options)) - .partitionIndex(0) - .start(0) - .length(fsRelation.sizeInBytes) - .path(file) - .pathType(PathType.URI_FILE) - .build() - }) - .toList - .asJava + files.map { + case (path, length) => + FileOrFiles + .builder() + .fileFormat(format) + .partitionIndex(0) + .start(0) + .length(length) + .path(path.toUri.toString) + .pathType(PathType.URI_FILE) + .build() + }.asJava ) .build() } - private def buildPartitionedFileScan(fsRelation: HadoopFsRelation): relation.Rel = { + private def partitionOutputMapping(fsRelation: HadoopFsRelation): Seq[Int] = { + // The merged schema captures Spark's resolution when the relation was constructed. Its first + // fields retain data-schema positions, with overlapping fields replaced by partition fields. + fsRelation.schema.zipWithIndex.map { + case (field, index) => + val partitionIndex = fsRelation.partitionSchema.indexOf(field) + val physicalCaseOverlap = !SparkCompat.instance.supportsCaseInsensitivePartitionOverlap && + index < fsRelation.dataSchema.size && field.name != fsRelation.dataSchema(index).name + if (partitionIndex >= 0 && !physicalCaseOverlap) { + fsRelation.dataSchema.size + partitionIndex + } else { + index + } + }.toSeq + } + + private def buildPartitionedFileScan( + fsRelation: HadoopFsRelation, + partitionFilters: Seq[Expression] = Nil): relation.Rel = { // LocalFiles does not carry directory partition values. Encode Spark's resolved values as // literals so consumers do not have to infer their types or base paths again. val fileFormat = convertFileFormat(fsRelation.fileFormat, fsRelation.options) val dataSchema = ToSubstraitType.toNamedStruct(fsRelation.dataSchema) - val resolver = - if ( - SparkCompat.instance.getConf(fsRelation.sparkSession, SQLConf.CASE_SENSITIVE.key).toBoolean - ) { - caseSensitiveResolution - } else { - caseInsensitiveResolution - } - - if ( - !SparkCompat.instance.supportsCaseInsensitivePartitionOverlap && - fsRelation.dataSchema.exists( - data => - fsRelation.partitionSchema.exists( - partition => data.name != partition.name && resolver(data.name, partition.name))) - ) { - throw new UnsupportedOperationException( - "This Spark version cannot reliably read file and partition columns that differ only in case") - } - - // Partition values override overlapping file columns, preserving the merged schema's order. - val outputMapping = fsRelation.schema.map { - field => - val partitionIndex = - fsRelation.partitionSchema.indexWhere(p => resolver(p.name, field.name)) - if (partitionIndex >= 0) { - fsRelation.dataSchema.size + partitionIndex - } else { - fsRelation.dataSchema.indexWhere(d => resolver(d.name, field.name)) - } - } + val outputMapping = partitionOutputMapping(fsRelation) val remap = relation.Rel.Remap.of(outputMapping.map(Int.box).toSeq.asJava) - val partitions = fsRelation.location.listFiles(Nil, Nil).filter(_.files.nonEmpty).map { - partition => - val read = relation.LocalFiles - .builder() - .initialSchema(dataSchema) - .addAllItems( - partition.files - .map { - file => - FileOrFiles - .builder() - .fileFormat(fileFormat) - .partitionIndex(0) - .start(0) - .length(file.getLen) - .path(file.getPath.toUri.toString) - .pathType(PathType.URI_FILE) - .build() + val partitions = + fsRelation.location.listFiles(partitionFilters, Nil).filter(_.files.nonEmpty).map { + partition => + val read = buildLocalFileScan( + dataSchema, + partition.files.map(f => (f.getPath, f.getLen)).toSeq, + fileFormat) + val values = fsRelation.partitionSchema.zipWithIndex.map { + case (field, index) => + val value = partition.values.get(index, field.dataType) match { + case decimal: Decimal => + val dt = field.dataType.asInstanceOf[DecimalType] + val fitted = decimal.clone() + if (fitted.changePrecision(dt.precision, dt.scale)) fitted else null + case other => other } - .toSeq - .asJava) - .build() - val values = fsRelation.partitionSchema.zipWithIndex.map { - case (field, index) => - ToSubstraitLiteral( - Literal(partition.values.get(index, field.dataType), field.dataType), - Some(field.nullable)) - } - relation.Project - .builder() - .input(read) - .addAllExpressions(values.toSeq.asJava) - .remap(remap) - .build() - } + ToSubstraitLiteral(Literal(value, field.dataType), Some(field.nullable)) + } + relation.Project + .builder() + .input(read) + .addAllExpressions(values.toSeq.asJava) + .remap(remap) + .build() + } - partitions.size match { + val result = partitions.size match { case 0 => relation.VirtualTableScan .builder() @@ -596,6 +647,11 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { case _ => relation.Set.builder().setOp(SetOp.UNION_ALL).addAllInputs(partitions.toSeq.asJava).build() } + result.withHint( + Optional.of( + Hint.builder + .addAllOutputNames(ToSubstraitType.toNamedStruct(fsRelation.schema).names()) + .build())) } private def convertFileFormat( diff --git a/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala b/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala index 9c41738f1..09bb0ca30 100644 --- a/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala +++ b/spark/src/main/spark-3.4/io/substrait/spark/compat/SparkCompatImpl.scala @@ -8,6 +8,12 @@ import io.substrait.relation class SparkCompatImpl extends SparkCompat { + override def createPartitionDirectory( + values: org.apache.spark.sql.catalyst.InternalRow, + files: Seq[org.apache.hadoop.fs.FileStatus] + ): org.apache.spark.sql.execution.datasources.PartitionDirectory = + org.apache.spark.sql.execution.datasources.PartitionDirectory(values, files) + override def supportsCaseInsensitivePartitionOverlap: Boolean = false override def createLogicalRelation( @@ -59,7 +65,7 @@ class SparkCompatImpl extends SparkCompat { override def createHadoopFsRelation( spark: AnyRef, - location: org.apache.spark.sql.execution.datasources.InMemoryFileIndex, + location: org.apache.spark.sql.execution.datasources.FileIndex, partitionSchema: org.apache.spark.sql.types.StructType, dataSchema: org.apache.spark.sql.types.StructType, bucketSpec: Option[org.apache.spark.sql.catalyst.catalog.BucketSpec], diff --git a/spark/src/main/spark-3.5/io/substrait/spark/compat/SparkCompatImpl.scala b/spark/src/main/spark-3.5/io/substrait/spark/compat/SparkCompatImpl.scala index e689e0467..750b93005 100644 --- a/spark/src/main/spark-3.5/io/substrait/spark/compat/SparkCompatImpl.scala +++ b/spark/src/main/spark-3.5/io/substrait/spark/compat/SparkCompatImpl.scala @@ -8,6 +8,12 @@ import io.substrait.relation class SparkCompatImpl extends SparkCompat { + override def createPartitionDirectory( + values: org.apache.spark.sql.catalyst.InternalRow, + files: Seq[org.apache.hadoop.fs.FileStatus] + ): org.apache.spark.sql.execution.datasources.PartitionDirectory = + org.apache.spark.sql.execution.datasources.PartitionDirectory(values, files.toArray) + override def createLogicalRelation( relation: HadoopFsRelation, output: Seq[AttributeReference], @@ -57,7 +63,7 @@ class SparkCompatImpl extends SparkCompat { override def createHadoopFsRelation( spark: AnyRef, - location: org.apache.spark.sql.execution.datasources.InMemoryFileIndex, + location: org.apache.spark.sql.execution.datasources.FileIndex, partitionSchema: org.apache.spark.sql.types.StructType, dataSchema: org.apache.spark.sql.types.StructType, bucketSpec: Option[org.apache.spark.sql.catalyst.catalog.BucketSpec], diff --git a/spark/src/main/spark-4.0/io/substrait/spark/compat/SparkCompatImpl.scala b/spark/src/main/spark-4.0/io/substrait/spark/compat/SparkCompatImpl.scala index 848f4cbe0..8380d0c39 100644 --- a/spark/src/main/spark-4.0/io/substrait/spark/compat/SparkCompatImpl.scala +++ b/spark/src/main/spark-4.0/io/substrait/spark/compat/SparkCompatImpl.scala @@ -6,6 +6,12 @@ import org.apache.spark.sql.execution.datasources.{HadoopFsRelation, LogicalRela class SparkCompatImpl extends SparkCompat { + override def createPartitionDirectory( + values: org.apache.spark.sql.catalyst.InternalRow, + files: Seq[org.apache.hadoop.fs.FileStatus] + ): org.apache.spark.sql.execution.datasources.PartitionDirectory = + org.apache.spark.sql.execution.datasources.PartitionDirectory(values, files.toArray) + override def createLogicalRelation( relation: HadoopFsRelation, output: Seq[AttributeReference], @@ -55,7 +61,7 @@ class SparkCompatImpl extends SparkCompat { override def createHadoopFsRelation( spark: AnyRef, - location: org.apache.spark.sql.execution.datasources.InMemoryFileIndex, + location: org.apache.spark.sql.execution.datasources.FileIndex, partitionSchema: org.apache.spark.sql.types.StructType, dataSchema: org.apache.spark.sql.types.StructType, bucketSpec: Option[org.apache.spark.sql.catalyst.catalog.BucketSpec], diff --git a/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala index 8a916fff3..88d214e07 100644 --- a/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala +++ b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala @@ -4,16 +4,17 @@ import io.substrait.spark.compat.SparkCompat import io.substrait.spark.logical.{ToLogicalPlan, ToSubstraitRel} import org.apache.spark.sql.Row -import org.apache.spark.sql.catalyst.analysis.caseSensitiveResolution -import org.apache.spark.sql.catalyst.expressions.Expression -import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Alias, Ascending, EqualTo, Expression, Literal, SortOrder} +import org.apache.spark.sql.catalyst.plans.logical.{Filter, GlobalLimit, LocalLimit, LogicalPlan, Project, Sort, Union} import org.apache.spark.sql.classic.DatasetUtil import org.apache.spark.sql.execution.datasources.{FileIndex, HadoopFsRelation, LogicalRelation, PartitionDirectory} import org.apache.spark.sql.test.SharedSparkSession -import org.apache.spark.sql.types.{DataType, IntegerType, LongType, StringType, StructField, StructType} +import org.apache.spark.sql.types.{DecimalType, IntegerType, LongType, StringType, StructField, StructType} +import io.substrait.expression.ExpressionCreator import io.substrait.plan.{PlanProtoConverter, ProtoPlanConverter} -import io.substrait.relation.{LocalFiles => SubstraitLocalFiles} +import io.substrait.relation.{LocalFiles => SubstraitLocalFiles, Project => SubstraitProject, Set => SubstraitSet} import io.substrait.relation.files.FileOrFiles import org.apache.hadoop.fs.Path @@ -34,10 +35,16 @@ class PartitionedFilesSuite extends SharedSparkSession { assertResult(substrait)(decoded) val converted = new ToLogicalPlan(spark).convert(decoded) - assert( - DataType.equalsStructurallyByName(plan.schema, converted.schema, caseSensitiveResolution)) + assertResult(plan.schema.map(field => (field.name, field.dataType))) { + converted.schema.map(field => (field.name, field.dataType)) + } val actual = DatasetUtil.fromLogicalPlan(spark, converted).collect().toSeq assertResult(expected.sortBy(_.toString))(actual.sortBy(_.toString)) + + val bare = new ToLogicalPlan(spark).convert(new ToSubstraitRel().visit(plan)) + assertResult(plan.schema.map(field => (field.name, field.dataType))) { + bare.schema.map(field => (field.name, field.dataType)) + } } Seq("parquet", "orc", "csv").foreach { @@ -148,7 +155,7 @@ class PartitionedFilesSuite extends SharedSparkSession { } } - test("mixed-case overlapping columns are rejected when Spark does not override them") { + test("mixed-case overlapping columns reject only filters on the overlap") { withSQLConf("spark.sql.caseSensitive" -> "false") { withTempPath { directory => @@ -159,11 +166,29 @@ class PartitionedFilesSuite extends SharedSparkSession { if (SparkCompat.instance.supportsCaseInsensitivePartitionOverlap) { assertRoundTrip(plan, Seq(Row(1, 10, "physical"))) } else { - assertResult(Seq(Row(1, 999, "physical")))(data.collect().toSeq) + assertRoundTrip(plan, Seq(Row(1, 999, "physical"))) + assertRoundTrip(data.select("id").queryExecution.optimizedPlan, Seq(Row(1))) + assertRoundTrip( + data.select("id", "value").queryExecution.optimizedPlan, + Seq(Row(1, "physical"))) val error = intercept[UnsupportedOperationException] { - new ToSubstraitRel().convert(plan) + new ToSubstraitRel().convert(data.filter("p = 10").queryExecution.optimizedPlan) } - assert(error.getMessage.contains("differ only in case")) + assert(error.getMessage.contains("overlapping partition columns")) + val alias = Alias(plan.output(1), "aliased")() + val projected = Project(Seq(plan.output.head, alias), plan) + val sorted = Sort(Seq(SortOrder(projected.output.head, Ascending)), true, projected) + Seq(projected, sorted).foreach { + child => + val filtered = Filter(EqualTo(alias.toAttribute, Literal(10)), child) + intercept[UnsupportedOperationException](new ToSubstraitRel().convert(filtered)) + } + val union = Union(Seq(plan, plan), byName = false, allowMissingCol = false) + intercept[UnsupportedOperationException] { + new ToSubstraitRel().convert(Filter(EqualTo(union.output(1), Literal(10)), union)) + } + val limited = GlobalLimit(Literal(1), LocalLimit(Literal(1), plan)) + assertRoundTrip(Filter(EqualTo(limited.output(1), Literal(10)), limited), Seq.empty) } } } @@ -220,7 +245,155 @@ class PartitionedFilesSuite extends SharedSparkSession { val original = logical.relation.asInstanceOf[HadoopFsRelation] val selected = original.location.listFiles(Nil, Nil).filter(_.values.getInt(0) == 10) assertRoundTrip(withPartitions(original, selected), Seq(Row(1, 10))) + val emptyDirectory = + PartitionDirectory(InternalRow(30), Nil) + val withEmpty = withPartitions(original, selected :+ emptyDirectory) + assertRoundTrip(withEmpty, Seq(Row(1, 10))) + assert(new ToSubstraitRel().visit(withEmpty).isInstanceOf[SubstraitProject]) assertRoundTrip(withPartitions(original, Seq.empty), Seq.empty) } } + + test("partition projects import as one scan and partition filters prune exported files") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark + .range(8) + .selectExpr("id", "cast(id as int) part") + .write + .partitionBy("part") + .parquet(path) + val data = spark.read.parquet(path) + val exported = new ToSubstraitRel().visit(data.queryExecution.optimizedPlan) + assert(exported.isInstanceOf[SubstraitSet]) + val imported = new ToLogicalPlan(spark).convert(exported) + assertResult(1)(imported.collect { case _: LogicalRelation => 1 }.size) + assertResult(data.collect().toSeq.sortBy(_.toString)) { + DatasetUtil.fromLogicalPlan(spark, imported).collect().toSeq.sortBy(_.toString) + } + assertResult(Seq(Row(7L, 7))) { + DatasetUtil.fromLogicalPlan(spark, imported).filter("part = 7").collect().toSeq + } + val filtered = new ToSubstraitRel() + .visit(data.filter("part = 7").queryExecution.optimizedPlan) + .asInstanceOf[io.substrait.relation.Filter] + assert(filtered.getInput.isInstanceOf[SubstraitProject]) + assertRoundTrip(data.filter("part = 7").queryExecution.optimizedPlan, Seq(Row(7L, 7))) + } + } + + test("raw paths preserve URI punctuation and folder URIs normalize trailing slashes") { + Seq("a%3Ab", "a#b", "a?b").foreach { + name => + withTempPath { + directory => + val path = new java.io.File(directory, name) + spark.sql("select 1 id").write.parquet(path.getAbsolutePath) + val original = spark.read.parquet(path.getAbsolutePath).queryExecution.optimizedPlan + val scan = new ToSubstraitRel().visit(original).asInstanceOf[SubstraitLocalFiles] + val items = scan.getItems.asScala + .map( + file => + FileOrFiles + .builder() + .from(file) + .path(new URI(file.getPath.get()).getPath) + .build()) + .toSeq + val raw = SubstraitLocalFiles.builder().from(scan).items(items.asJava).build() + assertResult(Seq(Row(1)))( + DatasetUtil + .fromLogicalPlan(spark, new ToLogicalPlan(spark).convert(raw)) + .collect() + .toSeq) + val folder = FileOrFiles + .builder() + .from(scan.getItems.get(0)) + .path(path.toURI.toString) + .pathType(FileOrFiles.PathType.URI_FOLDER) + .build() + val folderScan = + SubstraitLocalFiles.builder().from(scan).items(Seq(folder).asJava).build() + assertResult(Seq(Row(1)))( + DatasetUtil + .fromLogicalPlan(spark, new ToLogicalPlan(spark).convert(folderScan)) + .collect() + .toSeq) + val empty = SubstraitLocalFiles + .builder() + .from(scan) + .items(Seq(FileOrFiles.builder().from(folder).path("").build()).asJava) + .build() + intercept[IllegalArgumentException](new ToLogicalPlan(spark).convert(empty)) + } + } + } + + test("complex literal projects remain separate scans on import") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id").write.parquet(path) + val scan = new ToSubstraitRel().visit(spark.read.parquet(path).queryExecution.optimizedPlan) + val literal = ExpressionCreator.list(false, ExpressionCreator.i32(false, 7)) + val project = SubstraitProject.builder().input(scan).addExpressions(literal).build() + val union = SubstraitSet + .builder() + .setOp(SubstraitSet.SetOp.UNION_ALL) + .addInputs(project, project) + .build() + val imported = new ToLogicalPlan(spark).convert(union) + assertResult(2)(imported.collect { case _: LogicalRelation => 1 }.size) + assertResult(Seq(Row(1, Seq(7)), Row(1, Seq(7)))) { + DatasetUtil.fromLogicalPlan(spark, imported).collect().toSeq + } + } + } + + test("partition decimals are rounded and overflow to null using their declared type") { + // Spark's vectorized reader reinterprets the unscaled directory value instead of fitting it. + withSQLConf("spark.sql.parquet.enableVectorizedReader" -> "false") { + Seq( + ("1.25", DecimalType(10, 1), new java.math.BigDecimal("1.3")), + ("123.4", DecimalType(3, 1), null)).foreach { + case (value, decimalType, expected) => + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id").write.parquet(s"$path/part=$value") + val data = spark.read + .schema( + StructType(Seq(StructField("id", IntegerType), StructField("part", decimalType)))) + .parquet(path) + assertRoundTrip(data.queryExecution.optimizedPlan, Seq(Row(1, expected))) + } + } + } + } + + test("partition mapping uses the captured merged schema and Unicode names") { + withSQLConf("spark.sql.caseSensitive" -> "false ") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id, 999 `İl`").write.parquet(s"$path/il=34") + val data = spark.read.parquet(path) + assertRoundTrip(data.queryExecution.optimizedPlan, Seq(Row(1, 999, 34))) + } + } + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id, 999 P").write.parquet(s"$path/p=10") + withSQLConf("spark.sql.caseSensitive" -> "true") { + val plan = spark.read.parquet(path).queryExecution.optimizedPlan + withSQLConf("spark.sql.caseSensitive" -> "false") { + val read = new ToSubstraitRel().visit(plan).asInstanceOf[SubstraitProject] + assertResult(Seq(0, 1, 2))( + read.getRemap.get().indices().asScala.map(_.intValue()).toSeq) + } + } + } + } } From 6889df97429fdb1dccab00c339b5ed72d45a5d7c Mon Sep 17 00:00:00 2001 From: bvolpato Date: Wed, 7 Oct 2026 01:55:50 -0400 Subject: [PATCH 3/3] fix(spark): preserve empty unpartitioned file scans --- .../spark/logical/ToSubstraitRel.scala | 16 ++++++++++----- .../spark/PartitionedFilesSuite.scala | 20 +++++++++++++++++++ 2 files changed, 31 insertions(+), 5 deletions(-) diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala index 247f0b0a8..dc0d09c4f 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala @@ -554,11 +554,17 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging with Predic } private def buildLocalFileScan(fsRelation: HadoopFsRelation): relation.AbstractReadRel = { - buildLocalFileScan( - ToSubstraitType.toNamedStruct(fsRelation.schema), - fsRelation.location.listFiles(Nil, Nil).flatMap(_.files.map(f => (f.getPath, f.getLen))), - convertFileFormat(fsRelation.fileFormat, fsRelation.options) - ) + val files = + fsRelation.location.listFiles(Nil, Nil).flatMap(_.files.map(f => (f.getPath, f.getLen))) + if (files.isEmpty) { + buildVirtualTableScan(fsRelation.schema, Nil) + } else { + buildLocalFileScan( + ToSubstraitType.toNamedStruct(fsRelation.schema), + files, + convertFileFormat(fsRelation.fileFormat, fsRelation.options) + ) + } } private def buildLocalFileScan( diff --git a/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala index 88d214e07..e727d93f5 100644 --- a/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala +++ b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala @@ -96,6 +96,26 @@ class PartitionedFilesSuite extends SharedSparkSession { } } + test("an unpartitioned zero-byte CSV file round-trips as an empty table") { + withTempPath { + directory => + val path = directory.getAbsolutePath + spark.sql("select 1 id where false").coalesce(1).write.csv(path) + val data = spark.read.schema("id INT").csv(path) + assert(data.inputFiles.nonEmpty) + assertRoundTrip(data.queryExecution.optimizedPlan, Seq.empty) + } + } + + test("an unpartitioned directory without data files round-trips as an empty table") { + withTempPath { + directory => + assert(directory.mkdirs()) + val data = spark.read.schema("id INT").csv(directory.getAbsolutePath) + assertRoundTrip(data.queryExecution.optimizedPlan, Seq.empty) + } + } + test("date null and escaped string partition values retain their types and values") { withSQLConf("spark.sql.datetime.java8API.enabled" -> "true") { withTempPath {