diff --git a/spark/src/main/scala/io/substrait/spark/FileHolder.scala b/spark/src/main/scala/io/substrait/spark/FileHolder.scala index baa4cf5f8..cf774d779 100644 --- a/spark/src/main/scala/io/substrait/spark/FileHolder.scala +++ b/spark/src/main/scala/io/substrait/spark/FileHolder.scala @@ -6,6 +6,7 @@ import io.substrait.relation.{ProtoRelConverter, RelProtoConverter} import io.substrait.relation.Extension.WriteExtensionObject import io.substrait.relation.files.FileOrFiles +/** File target for unpartitioned, unbucketed INSERT writes with append semantics. */ case class FileHolder(fileOrFiles: FileOrFiles) extends WriteExtensionObject { override def toProto(converter: RelProtoConverter): protobuf.Any = { 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..271f944d7 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala @@ -45,7 +45,7 @@ import io.substrait.plan.Plan import io.substrait.relation import io.substrait.relation.{ExtensionWrite, LocalFiles, NamedDdl, NamedWrite} import io.substrait.relation.AbstractDdlRel.{DdlObject, DdlOp} -import io.substrait.relation.AbstractWriteRel.{CreateMode, WriteOp} +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 @@ -531,11 +531,22 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes } override def visit(write: ExtensionWrite, context: EmptyVisitationContext): LogicalPlan = { - val child = write.getInput.accept(this, context) - val mode = write.getOperation match { - case WriteOp.INSERT => SaveMode.Append - case WriteOp.UPDATE => SaveMode.Overwrite - case op => throw new UnsupportedOperationException(s"Write mode $op not supported") + if (write.getOperation != WriteOp.INSERT) { + throw new UnsupportedOperationException(s"Write mode ${write.getOperation} not supported") + } + // The spec defines create_mode for CTAS and is silent on INSERT. Older file writes used + // it for Spark save modes, so reject the modes an append-only file extension cannot honor. + write.getCreateMode match { + case CreateMode.UNSPECIFIED | CreateMode.APPEND_IF_EXISTS => + case createMode => + throw new UnsupportedOperationException( + s"Filesystem INSERT does not support create mode $createMode") + } + write.getOutputMode match { + case OutputMode.UNSPECIFIED | OutputMode.NO_OUTPUT => + case outputMode => + throw new UnsupportedOperationException( + s"Filesystem INSERT does not support output mode $outputMode") } val file = write.getDetail match { @@ -554,6 +565,7 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes val name = file.getPath.get.split('/').reverse.head val table = catalogTable(Seq(name)) + val child = write.getInput.accept(this, context) val plan = withChild(child) { V1Writes.apply( @@ -566,7 +578,7 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes fileFormat = format, options = options, query = child, - mode = mode, + mode = SaveMode.Append, catalogTable = Some(table), fileIndex = None, outputColumnNames = write.getTableSchema.names.asScala.toSeq 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..6ab8b9524 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToSubstraitRel.scala @@ -602,7 +602,13 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { throw new UnsupportedOperationException(s"Unable to convert command: $command") } - private def convertDataWritingCommand(command: V1WriteCommand): relation.AbstractWriteRel = + private def convertDataWritingCommand(command: V1WriteCommand): relation.AbstractWriteRel = { + if (command.staticPartitions.nonEmpty || command.partitionColumns.nonEmpty) { + throw new UnsupportedOperationException("Partitioned writes are not supported") + } + if (command.bucketSpec.nonEmpty) { + throw new UnsupportedOperationException("Bucketed writes are not supported") + } command match { case InsertIntoHadoopFsRelationCommand( outputPath, @@ -617,6 +623,10 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { _, _, outputColumnNames) => + if (mode != SaveMode.Append) { + throw new UnsupportedOperationException( + s"Filesystem writes only support SaveMode.Append, found $mode") + } val file = FileOrFiles .builder() .fileFormat(convertFileFormat(fileFormat, options)) @@ -632,7 +642,7 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { .input(visit(child)) .operation(WriteOp.INSERT) .outputMode(OutputMode.UNSPECIFIED) - .createMode(createMode(mode)) + .createMode(CreateMode.UNSPECIFIED) .tableSchema(outputSchema(child.output, outputColumnNames)) .detail(FileHolder(file)) .build() @@ -649,6 +659,7 @@ class ToSubstraitRel extends AbstractLogicalPlanVisitor with Logging { case _ => throw new UnsupportedOperationException(s"Unable to convert command: ${command.getClass}") } + } private def convertCTAS( table: CatalogTable, diff --git a/spark/src/test/scala/io/substrait/spark/FileWriteSuite.scala b/spark/src/test/scala/io/substrait/spark/FileWriteSuite.scala new file mode 100644 index 000000000..1dee82c33 --- /dev/null +++ b/spark/src/test/scala/io/substrait/spark/FileWriteSuite.scala @@ -0,0 +1,177 @@ +package io.substrait.spark + +import io.substrait.spark.logical.{ToLogicalPlan, ToSubstraitRel} + +import org.apache.spark.SparkFunSuite +import org.apache.spark.sql.{Row, SaveMode} +import org.apache.spark.sql.catalyst.TableIdentifier +import org.apache.spark.sql.catalyst.catalog.BucketSpec +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan +import org.apache.spark.sql.execution.datasources.InsertIntoHadoopFsRelationCommand +import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat +import org.apache.spark.sql.test.SharedSparkSession + +import io.substrait.extension.ExtensionCollector +import io.substrait.relation.{ExtensionWrite, RelProtoConverter} +import io.substrait.relation.AbstractWriteRel.{CreateMode, OutputMode, WriteOp} +import org.apache.hadoop.fs.Path + +class FileWriteSuite extends SparkFunSuite with SharedSparkSession { + + private def withTarget(f: InsertIntoHadoopFsRelationCommand => Unit): Unit = { + withTable("file_write_target") { + spark.sql("CREATE TABLE file_write_target (id INT) USING PARQUET") + spark.sql("INSERT INTO file_write_target VALUES (1), (2)") + val table = spark.sessionState.catalog.getTableMetadata(TableIdentifier("file_write_target")) + val child = spark.sql("SELECT 3 AS id").queryExecution.optimizedPlan + f( + InsertIntoHadoopFsRelationCommand( + outputPath = new Path(table.location), + staticPartitions = Map.empty, + ifPartitionNotExists = false, + partitionColumns = Seq.empty, + bucketSpec = None, + fileFormat = new ParquetFileFormat(), + options = Map.empty, + query = child, + mode = SaveMode.Append, + catalogTable = Some(table), + fileIndex = None, + outputColumnNames = Seq("id") + )) + } + } + + private def targetRows: Seq[Row] = + spark.sql("SELECT id FROM file_write_target ORDER BY id").collect().toSeq + + private def convertWrite(command: InsertIntoHadoopFsRelationCommand): ExtensionWrite = + new ToSubstraitRel().visit(command).asInstanceOf[ExtensionWrite] + + private def importProto(write: ExtensionWrite): LogicalPlan = { + val collector = new ExtensionCollector + val bytes = new RelProtoConverter(collector).toProto(write).toByteArray + val decoded = new FileHolderHandlingProtoRelConverter(collector) + .from(io.substrait.proto.Rel.parseFrom(bytes)) + assertResult(write)(decoded) + new ToLogicalPlan(spark).convert(decoded) + } + + test("append writes preserve existing rows through the file extension protobuf") { + withTarget { + command => + val write = convertWrite(command) + assertResult(CreateMode.UNSPECIFIED)(write.getCreateMode) + val plan = importProto(write) + spark.sessionState.executePlan(plan).executedPlan.execute() + assertResult(Seq(Row(1), Row(2), Row(3)))(targetRows) + } + } + + test("legacy append file extensions remain executable on unpartitioned targets") { + withTarget { + command => + val write = ExtensionWrite + .builder() + .from(convertWrite(command)) + .createMode(CreateMode.APPEND_IF_EXISTS) + .build() + spark.sessionState.executePlan(importProto(write)).executedPlan.execute() + assertResult(Seq(Row(1), Row(2), Row(3)))(targetRows) + } + } + + test("reject filesystem save modes that cannot be represented as INSERT") { + withTarget { + command => + Seq(SaveMode.Overwrite, SaveMode.Ignore, SaveMode.ErrorIfExists).foreach { + mode => + val error = intercept[UnsupportedOperationException] { + convertWrite(command.copy(mode = mode)) + } + assert(error.getMessage.contains(s"SaveMode.Append, found $mode")) + } + } + } + + test("reject partition and bucket metadata that the file extension cannot carry") { + withTarget { + command => + val partitioned = spark.sql("SELECT 3 AS id, 10 AS part").queryExecution.optimizedPlan + val partitionedCommands = Seq( + command.copy(staticPartitions = Map("part" -> "10")), + command.copy( + partitionColumns = Seq(partitioned.output.last), + query = partitioned, + outputColumnNames = Seq("id", "part")) + ) + partitionedCommands.foreach { + unsupported => + val error = intercept[UnsupportedOperationException] { + convertWrite(unsupported) + } + assert(error.getMessage.contains("Partitioned writes are not supported")) + } + val bucketError = intercept[UnsupportedOperationException] { + convertWrite(command.copy(bucketSpec = Some(BucketSpec(2, Seq("id"), Seq.empty)))) + } + assert(bucketError.getMessage.contains("Bucketed writes are not supported")) + } + } + + test("reject legacy file save modes before constructing an executable INSERT") { + withTarget { + command => + Seq(CreateMode.REPLACE_IF_EXISTS, CreateMode.IGNORE_IF_EXISTS, CreateMode.ERROR_IF_EXISTS) + .foreach { + mode => + val write = + ExtensionWrite.builder().from(convertWrite(command)).createMode(mode).build() + val error = intercept[UnsupportedOperationException] { + importProto(write) + } + assert(error.getMessage.contains(s"INSERT does not support create mode $mode")) + } + } + } + + test("reject file UPDATE instead of replacing the entire target") { + withTarget { + command => + val write = + ExtensionWrite.builder().from(convertWrite(command)).operation(WriteOp.UPDATE).build() + val error = intercept[UnsupportedOperationException] { + importProto(write) + } + assert(error.getMessage.contains("Write mode UPDATE not supported")) + } + } + + test("reject file writes that request modified records") { + withTarget { + command => + val write = ExtensionWrite + .builder() + .from(convertWrite(command)) + .outputMode(OutputMode.MODIFIED_RECORDS) + .build() + val error = intercept[UnsupportedOperationException] { + importProto(write) + } + assert(error.getMessage.contains("INSERT does not support output mode MODIFIED_RECORDS")) + } + } + + test("file writes with NO_OUTPUT preserve append semantics") { + withTarget { + command => + val write = ExtensionWrite + .builder() + .from(convertWrite(command)) + .outputMode(OutputMode.NO_OUTPUT) + .build() + spark.sessionState.executePlan(importProto(write)).executedPlan.execute() + assertResult(Seq(Row(1), Row(2), Row(3)))(targetRows) + } + } +} diff --git a/spark/src/test/scala/io/substrait/spark/HiveTableSuite.scala b/spark/src/test/scala/io/substrait/spark/HiveTableSuite.scala index 785ec6669..833b70db1 100644 --- a/spark/src/test/scala/io/substrait/spark/HiveTableSuite.scala +++ b/spark/src/test/scala/io/substrait/spark/HiveTableSuite.scala @@ -4,7 +4,9 @@ import io.substrait.spark.logical.{ToLogicalPlan, ToSubstraitRel} import org.apache.spark.SparkFunSuite import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.catalyst.catalog.BucketSpec import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan +import org.apache.spark.sql.hive.execution.InsertIntoHiveTable import io.substrait.extension.ExtensionLookup import io.substrait.plan.{PlanProtoConverter, ProtoPlanConverter} @@ -112,4 +114,58 @@ class HiveTableSuite extends SparkFunSuite { assertResult(2)(spark.sql("select * from ctas").count()) } + test("reject partitioned Hive writes before losing the partition scope") { + spark.sql("drop table if exists partitioned_write_target") + try { + spark.sql( + "create table partitioned_write_target(id int, p string) using hive partitioned by (p)") + val parsed = spark.sessionState.sqlParser.parsePlan( + "insert overwrite table partitioned_write_target partition (p='a') select 1") + val staticWrite = spark.sessionState + .executePlan(parsed) + .analyzed + .collectFirst { case command: InsertIntoHiveTable => command } + .get + val dynamicInput = spark.sql("select 1 as id, 'a' as p").queryExecution.optimizedPlan + val dynamicWrite = staticWrite.copy( + partition = Map("p" -> None), + partitionColumns = Seq(dynamicInput.output.last), + query = dynamicInput, + outputColumnNames = Seq("id", "p")) + Seq(staticWrite, dynamicWrite) + .foreach { + command => + val error = intercept[UnsupportedOperationException] { + new ToSubstraitRel().convert(command) + } + assert(error.getMessage.contains("Partitioned writes are not supported")) + } + } finally { + spark.sql("drop table if exists partitioned_write_target") + } + } + + test("reject bucketed Hive writes before losing the bucket layout") { + spark.sql("drop table if exists bucketed_write_target") + try { + spark.sql("create table bucketed_write_target(id int) using hive") + val parsed = + spark.sessionState.sqlParser.parsePlan("insert into table bucketed_write_target select 1") + val command = spark.sessionState + .executePlan(parsed) + .analyzed + .collectFirst { case write: InsertIntoHiveTable => write } + .get + val bucketSpec = Some(BucketSpec(2, Seq("id"), Seq.empty)) + val error = intercept[UnsupportedOperationException] { + new ToSubstraitRel().convert( + command + .copy(table = command.table.copy(bucketSpec = bucketSpec), bucketSpec = bucketSpec)) + } + assert(error.getMessage.contains("Bucketed writes are not supported")) + } finally { + spark.sql("drop table if exists bucketed_write_target") + } + } + }