Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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],
Expand Down
164 changes: 154 additions & 10 deletions spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -386,15 +385,149 @@ 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 {
case SetOp.UNION_ALL => Union(children, byName = false, allowMissingCol = false)
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(
Expand Down Expand Up @@ -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(
Expand All @@ -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 =>
Expand Down
Loading
Loading