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..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,14 @@ 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 + /** Create a LogicalRelation with version-appropriate constructor */ def createLogicalRelation( relation: HadoopFsRelation, @@ -44,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 2d7b89b2e..baf315321 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,11 +49,12 @@ import io.substrait.relation.AbstractWriteRel.{CreateMode, OutputMode, 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 -import java.net.URI +import java.net.{URI, URISyntaxException} import java.util.Optional import scala.annotation.nowarn @@ -312,11 +313,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 @@ -386,6 +385,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 { @@ -393,8 +408,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( @@ -439,7 +572,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( @@ -460,6 +593,17 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes remap(plan, localFiles.getRemap) } + private def toFilePath(path: String): Path = { + try { + 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) + } + } + 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 6ab8b9524..dc0d09c4f 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala @@ -17,7 +17,7 @@ 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 @@ -38,7 +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.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} @@ -56,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} @@ -64,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) @@ -311,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 @@ -488,29 +554,112 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { } private def buildLocalFileScan(fsRelation: HadoopFsRelation): relation.AbstractReadRel = { + 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( + 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 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 outputMapping = partitionOutputMapping(fsRelation) + val remap = relation.Rel.Remap.of(outputMapping.map(Int.box).toSeq.asJava) + + 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 + } + ToSubstraitLiteral(Literal(value, field.dataType), Some(field.nullable)) + } + relation.Project + .builder() + .input(read) + .addAllExpressions(values.toSeq.asJava) + .remap(remap) + .build() + } + + val result = 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() + } + result.withHint( + Optional.of( + Hint.builder + .addAllOutputNames(ToSubstraitType.toNamedStruct(fsRelation.schema).names()) + .build())) + } + private def convertFileFormat( fileFormat: DSFileFormat, options: Map[String, String]): FileFormat = fileFormat match { @@ -533,7 +682,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 +707,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..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,14 @@ 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( relation: HadoopFsRelation, output: Seq[AttributeReference], @@ -57,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 new file mode 100644 index 000000000..e727d93f5 --- /dev/null +++ b/spark/src/test/scala/io/substrait/spark/PartitionedFilesSuite.scala @@ -0,0 +1,419 @@ +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.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.{DecimalType, IntegerType, LongType, StringType, StructField, StructType} + +import io.substrait.expression.ExpressionCreator +import io.substrait.plan.{PlanProtoConverter, ProtoPlanConverter} +import io.substrait.relation.{LocalFiles => SubstraitLocalFiles, Project => SubstraitProject, Set => SubstraitSet} +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) + 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 { + 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("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 { + 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 reject only filters on the overlap") { + 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 { + 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(data.filter("p = 10").queryExecution.optimizedPlan) + } + 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) + } + } + } + } + + 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))) + 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) + } + } + } + } +}