diff --git a/connector/avro/src/test/scala/org/apache/spark/sql/avro/AvroSuite.scala b/connector/avro/src/test/scala/org/apache/spark/sql/avro/AvroSuite.scala index 8134017738346..9767f6c73d644 100644 --- a/connector/avro/src/test/scala/org/apache/spark/sql/avro/AvroSuite.scala +++ b/connector/avro/src/test/scala/org/apache/spark/sql/avro/AvroSuite.scala @@ -40,7 +40,7 @@ import org.apache.spark.sql.TestingUDT.IntervalData import org.apache.spark.sql.avro.AvroCompressionCodec._ import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.catalyst.plans.logical.Filter -import org.apache.spark.sql.catalyst.util.DateTimeTestUtils +import org.apache.spark.sql.catalyst.util.{CharVarcharUtils, DateTimeTestUtils} import org.apache.spark.sql.catalyst.util.DateTimeTestUtils.{withDefaultTimeZone, LA, UTC} import org.apache.spark.sql.execution.{FormattedMode, SparkPlan} import org.apache.spark.sql.execution.datasources.{CommonFileDataSourceSuite, DataSource, FilePartition} @@ -3724,6 +3724,87 @@ abstract class AvroSuite } } + test("SPARK-58814: Avro infers nested CHAR/VARCHAR schema and values") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + withTempPath { dir => + val path = dir.getCanonicalPath + val input = spark.range(1).selectExpr( + "cast('ab' AS CHAR(4)) AS c", + "cast('xy' AS VARCHAR(3)) AS v", + "named_struct('c', cast('z' AS CHAR(2))) AS s", + "array(cast('q' AS VARCHAR(2))) AS a", + "map(cast('k' AS CHAR(2)), cast('v' AS VARCHAR(2))) AS m") + input.write.mode("overwrite").format("avro").save(path) + + val readBack = spark.read.format("avro").load(path) + assert(DataType.equalsIgnoreNullability(readBack.schema, input.schema)) + checkAnswer( + readBack.selectExpr( + "concat('<', c, '>')", + "v", + "concat('<', s.c, '>')", + "a", + "m"), + Row("", "xy", "", Seq("q"), Map("k " -> "v"))) + + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { + assert(DataType.equalsIgnoreNullability( + spark.read.format("avro").load(path).schema, + CharVarcharUtils.replaceCharVarcharWithString(input.schema))) + } + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { + assert(DataType.equalsIgnoreNullability( + spark.read.format("avro").load(path).schema, + input.schema)) + } + } + + withTempPath { dir => + Seq("ab").toDF("c").write.format("avro").save(dir.getCanonicalPath) + val charDf = spark.read.schema("c CHAR(4)").format("avro").load(dir.getCanonicalPath) + checkAnswer( + charDf.selectExpr("concat('<', c, '>')"), + Row("")) + } + withTempPath { dir => + Seq("abcdef").toDF("c").write.format("avro").save(dir.getCanonicalPath) + Seq("CHAR", "VARCHAR").foreach { typ => + checkError( + exception = intercept[SparkRuntimeException] { + spark.read.schema(s"c $typ(4)").format("avro") + .load(dir.getCanonicalPath).collect() + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + } + + withTable("avro_char_varchar_assignment") { + sql( + """CREATE TABLE avro_char_varchar_assignment + |(c CHAR(4), v VARCHAR(4)) USING avro""".stripMargin) + sql("INSERT INTO avro_char_varchar_assignment VALUES ('ab', 'xy')") + assert(spark.table("avro_char_varchar_assignment").schema.map(_.dataType) === + Seq(CharType(4), VarcharType(4))) + checkAnswer( + sql( + """SELECT concat('<', c, '>'), v + |FROM avro_char_varchar_assignment""".stripMargin), + Row("", "xy")) + checkError( + exception = intercept[SparkRuntimeException] { + sql( + """INSERT INTO avro_char_varchar_assignment + |VALUES ('abcde', 'xy')""".stripMargin).collect() + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + } + } + } class AvroV1Suite extends AvroSuite { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/DataSourceScanExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/DataSourceScanExec.scala index a727ccf565063..10a3da7df91a7 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/DataSourceScanExec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/DataSourceScanExec.scala @@ -34,6 +34,7 @@ import org.apache.spark.sql.connector.read.streaming.SparkDataStream import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.sql.execution import org.apache.spark.sql.execution.datasources._ +import org.apache.spark.sql.execution.datasources.orc.OrcFileFormat import org.apache.spark.sql.execution.datasources.parquet.{ParquetFileFormat => ParquetSource} import org.apache.spark.sql.execution.datasources.v2.{PushedDownOperators, TableSampleInfo} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} @@ -752,15 +753,30 @@ case class FileSourceScanExec( lazy val inputRDD: RDD[InternalRow] = { val options = relation.options + (FileFormat.OPTION_RETURNING_BATCH -> supportsColumnar.toString) - val readFile: (PartitionedFile) => Iterator[InternalRow] = - relation.fileFormat.buildReaderWithPartitionValues( - sparkSession = relation.sparkSession, - dataSchema = relation.dataSchema, - partitionSchema = relation.partitionSchema, - requiredSchema = requiredSchema, - filters = pushedDownFilters, - options = options, - hadoopConf = getHadoopConf(relation.sparkSession, relation.options)) + val hadoopConf = getHadoopConf(relation.sparkSession, relation.options) + val readFile: (PartitionedFile) => Iterator[InternalRow] = relation.fileFormat match { + case format: OrcFileFormat + if getTagValue(ApplyCharTypePadding.standardSemanticsTag).isDefined => + format.buildReaderWithPartitionValues( + sparkSession = relation.sparkSession, + dataSchema = relation.dataSchema, + partitionSchema = relation.partitionSchema, + requiredSchema = requiredSchema, + filters = pushedDownFilters, + options = options, + hadoopConf = hadoopConf, + charVarcharStandardSemantics = + getTagValue(ApplyCharTypePadding.standardSemanticsTag).get) + case format => + format.buildReaderWithPartitionValues( + sparkSession = relation.sparkSession, + dataSchema = relation.dataSchema, + partitionSchema = relation.partitionSchema, + requiredSchema = requiredSchema, + filters = pushedDownFilters, + options = options, + hadoopConf = hadoopConf) + } val readRDD = if (bucketedScan) { createBucketedReadRDD(relation.bucketSpec.get, readFile, dynamicallySelectedPartitions) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/ApplyCharTypePadding.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/ApplyCharTypePadding.scala index a1cc576b2b236..d0f277445d3e0 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/ApplyCharTypePadding.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/ApplyCharTypePadding.scala @@ -24,6 +24,7 @@ import org.apache.spark.sql.catalyst.analysis.ApplyCharTypePaddingHelper import org.apache.spark.sql.catalyst.catalog.HiveTableRelation import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.rules.Rule +import org.apache.spark.sql.catalyst.trees.TreeNodeTag import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation import org.apache.spark.sql.internal.SQLConf @@ -39,6 +40,9 @@ import org.apache.spark.sql.internal.SQLConf */ object ApplyCharTypePadding extends Rule[LogicalPlan] { + private[sql] val standardSemanticsTag = + TreeNodeTag[Boolean]("__char_varchar_standard_semantics") + private val readSidePaddingOverrideWarned = new AtomicBoolean(false) private def warnReadSidePaddingOverride(): Unit = { @@ -51,34 +55,59 @@ object ApplyCharTypePadding extends Rule[LogicalPlan] { } override def apply(plan: LogicalPlan): LogicalPlan = { + val standardSemantics = conf.charVarcharStandardSemantics + + def bindStandardSemantics[T <: LogicalPlan](relation: T): T = { + if (relation.getTagValue(standardSemanticsTag).isEmpty) { + relation.setTagValue(standardSemanticsTag, standardSemantics) + } + relation + } + + val boundPlan = if (conf.charVarcharFirstClassTypes) { + plan.resolveOperatorsUp { + case relation: LogicalRelation => bindStandardSemantics(relation) + case relation: DataSourceV2Relation => bindStandardSemantics(relation) + case relation: HiveTableRelation => bindStandardSemantics(relation) + } + } else { + plan + } + // standardSemantics takes precedence over legacy charVarcharAsString. - if (conf.charVarcharAsString && !conf.charVarcharStandardSemantics) { - return plan + if (conf.charVarcharAsString && !standardSemantics) { + return boundPlan } - if (conf.charVarcharStandardSemantics && !conf.readSideCharPadding) { + if (standardSemantics && !conf.readSideCharPadding) { warnReadSidePaddingOverride() } - if (conf.readSideCharPadding || conf.charVarcharStandardSemantics) { - val newPlan = plan.resolveOperatorsUpWithNewOutput { + if (conf.readSideCharPadding || standardSemantics) { + val newPlan = boundPlan.resolveOperatorsUpWithNewOutput { case r: LogicalRelation => + bindStandardSemantics(r) ApplyCharTypePaddingHelper.readSidePadding(r, () => - r.copy(output = r.output.map(CharVarcharUtils.cleanAttrMetadata))) + bindStandardSemantics( + r.copy(output = r.output.map(CharVarcharUtils.cleanAttrMetadata)))) case r: DataSourceV2Relation => + bindStandardSemantics(r) ApplyCharTypePaddingHelper.readSidePadding(r, () => - r.copy(output = r.output.map(CharVarcharUtils.cleanAttrMetadata))) + bindStandardSemantics( + r.copy(output = r.output.map(CharVarcharUtils.cleanAttrMetadata)))) case r: HiveTableRelation => + bindStandardSemantics(r) ApplyCharTypePaddingHelper.readSidePadding(r, () => { val cleanedDataCols = r.dataCols.map(CharVarcharUtils.cleanAttrMetadata) val cleanedPartCols = r.partitionCols.map(CharVarcharUtils.cleanAttrMetadata) - r.copy(dataCols = cleanedDataCols, partitionCols = cleanedPartCols) + bindStandardSemantics( + r.copy(dataCols = cleanedDataCols, partitionCols = cleanedPartCols)) }) } ApplyCharTypePaddingHelper.paddingForStringComparison(newPlan, padCharCol = false) } else { ApplyCharTypePaddingHelper.paddingForStringComparison( - plan, padCharCol = !conf.getConf(SQLConf.LEGACY_NO_CHAR_PADDING_IN_PREDICATE)) + boundPlan, padCharCol = !conf.getConf(SQLConf.LEGACY_NO_CHAR_PADDING_IN_PREDICATE)) } } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/FileSourceStrategy.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/FileSourceStrategy.scala index e2427222d8ebc..f8e7e8d2249ec 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/FileSourceStrategy.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/FileSourceStrategy.scala @@ -340,6 +340,9 @@ object FileSourceStrategy extends Strategy with PredicateHelper with Logging { table.map(_.identifier), markedForSingleTaskExecution = l.getTagValue(MarkSingleTaskExecution.markTag).getOrElse(false)) + l.getTagValue(ApplyCharTypePadding.standardSemanticsTag).foreach { + enabled => scan.setTagValue(ApplyCharTypePadding.standardSemanticsTag, enabled) + } // extra Project node: wrap flat metadata columns to a metadata struct val withMetadataProjections = metadataStructOpt.map { metadataStruct => diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcDeserializer.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcDeserializer.scala index 42fcb0cea2eb4..73168e5297aaf 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcDeserializer.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcDeserializer.scala @@ -136,7 +136,7 @@ class OrcDeserializer( case DoubleType => (ordinal, value) => updater.setDouble(ordinal, value.asInstanceOf[DoubleWritable].get) - case StringType => (ordinal, value) => + case _: StringType => (ordinal, value) => updater.set(ordinal, UTF8String.fromBytes(value.asInstanceOf[Text].copyBytes)) case BinaryType => (ordinal, value) => diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcFileFormat.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcFileFormat.scala index c069abe3b4380..8cace8669d631 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcFileFormat.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcFileFormat.scala @@ -146,7 +146,26 @@ class OrcFileFormat filters: Seq[Filter], options: Map[String, String], hadoopConf: Configuration): (PartitionedFile) => Iterator[InternalRow] = { + buildReaderWithPartitionValues( + sparkSession, + dataSchema, + partitionSchema, + requiredSchema, + filters, + options, + hadoopConf, + charVarcharStandardSemantics = false) + } + private[sql] def buildReaderWithPartitionValues( + sparkSession: SparkSession, + dataSchema: StructType, + partitionSchema: StructType, + requiredSchema: StructType, + filters: Seq[Filter], + options: Map[String, String], + hadoopConf: Configuration, + charVarcharStandardSemantics: Boolean): (PartitionedFile) => Iterator[InternalRow] = { val resultSchema = StructType(requiredSchema.fields ++ partitionSchema.fields) val sqlConf = getSqlConf(sparkSession) val capacity = sqlConf.orcVectorizedReaderBatchSize @@ -205,7 +224,7 @@ class OrcFileFormat val (requestedColIds, canPruneCols) = resultedColPruneInfo.get val resultSchemaString = OrcUtils.orcResultSchemaString(canPruneCols, - dataSchema, resultSchema, partitionSchema, conf) + dataSchema, resultSchema, partitionSchema, conf, charVarcharStandardSemantics) assert(requestedColIds.length == requiredSchema.length, "[BUG] requested column IDs do not match required schema") val taskConf = new Configuration(conf) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcUtils.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcUtils.scala index f533ded6243c2..ff54eee09ae10 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcUtils.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/orc/OrcUtils.scala @@ -427,16 +427,29 @@ object OrcUtils extends Logging { * Given a `StructType` object, this methods converts it to corresponding string representation * in ORC. */ - def getOrcSchemaString(dt: DataType): String = dt match { + def getOrcSchemaString(dt: DataType): String = { + getOrcSchemaString(dt, SQLConf.get.charVarcharStandardSemantics) + } + + private def getOrcSchemaString( + dt: DataType, + charVarcharStandardSemantics: Boolean): String = dt match { case s: StructType => val fieldTypes = s.fields.map { f => - s"${quoteIdentifier(f.name)}:${getOrcSchemaString(f.dataType)}" + s"${quoteIdentifier(f.name)}:" + + s"${getOrcSchemaString(f.dataType, charVarcharStandardSemantics)}" } s"struct<${fieldTypes.mkString(",")}>" case a: ArrayType => - s"array<${getOrcSchemaString(a.elementType)}>" + s"array<${getOrcSchemaString(a.elementType, charVarcharStandardSemantics)}>" case m: MapType => - s"map<${getOrcSchemaString(m.keyType)},${getOrcSchemaString(m.valueType)}>" + s"map<${getOrcSchemaString(m.keyType, charVarcharStandardSemantics)}," + + s"${getOrcSchemaString(m.valueType, charVarcharStandardSemantics)}>" + // Under standard semantics, keep Spark responsible for CHAR/VARCHAR assignment and scan + // checks. Native ORC would truncate or pad before Spark can validate the original value. + // Preserve-only mode retains the native constrained schema and its legacy enforcement. + case _: CharType | _: VarcharType if charVarcharStandardSemantics => + StringType.catalogString case _: DayTimeIntervalType | _: TimestampNTZType => LongType.catalogString case _: YearMonthIntervalType => IntegerType.catalogString // Framework types (TimeType, nanosecond timestamps) supply their own ORC schema string. @@ -535,11 +548,14 @@ object OrcUtils extends Logging { dataSchema: StructType, resultSchema: StructType, partitionSchema: StructType, - conf: Configuration): String = { + conf: Configuration, + charVarcharStandardSemantics: Boolean): String = { val resultSchemaString = if (canPruneCols) { - OrcUtils.getOrcSchemaString(resultSchema) + OrcUtils.getOrcSchemaString(resultSchema, charVarcharStandardSemantics) } else { - OrcUtils.getOrcSchemaString(StructType(dataSchema.fields ++ partitionSchema.fields)) + OrcUtils.getOrcSchemaString( + StructType(dataSchema.fields ++ partitionSchema.fields), + charVarcharStandardSemantics) } OrcConf.MAPRED_INPUT_SCHEMA.setString(conf, resultSchemaString) resultSchemaString diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileScanBuilder.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileScanBuilder.scala index 46f6702003ff3..76424bd5054c6 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileScanBuilder.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileScanBuilder.scala @@ -27,6 +27,10 @@ import org.apache.spark.sql.internal.connector.SupportsPushDownCatalystFilters import org.apache.spark.sql.sources.Filter import org.apache.spark.sql.types.StructType +private[sql] trait SupportsCharVarcharStandardSemantics { + def bindCharVarcharStandardSemantics(enabled: Boolean): Unit +} + abstract class FileScanBuilder( sparkSession: SparkSession, fileIndex: PartitioningAwareFileIndex, diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanRelationPushDown.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanRelationPushDown.scala index d8817b9aa7582..f8b289995cda6 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanRelationPushDown.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanRelationPushDown.scala @@ -35,7 +35,7 @@ import org.apache.spark.sql.connector.expressions.{SortOrder => V2SortOrder} import org.apache.spark.sql.connector.expressions.aggregate.{Aggregation, Avg, Count, CountStar, Max, Min, Sum} import org.apache.spark.sql.connector.expressions.filter.Predicate import org.apache.spark.sql.connector.read.{Scan, ScanBuilder, Statistics => V2Statistics, SupportsPushDownAggregates, SupportsPushDownFilters, SupportsPushDownJoin, SupportsPushDownRequiredColumns, SupportsPushDownVariantExtractions, SupportsReportStatistics, V1Scan, VariantExtraction} -import org.apache.spark.sql.execution.datasources.{DataSourceStrategy, VariantInRelation, VariantMetadata} +import org.apache.spark.sql.execution.datasources.{ApplyCharTypePadding, DataSourceStrategy, VariantInRelation, VariantMetadata} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.internal.connector.VariantExtractionImpl import org.apache.spark.sql.sources @@ -98,7 +98,13 @@ object V2ScanRelationPushDown extends Rule[LogicalPlan] with PredicateHelper { private def createScanBuilder(plan: LogicalPlan) = plan.transform { case r: DataSourceV2Relation => - ScanBuilderHolder(r.output, r, r.table.asReadable.newScanBuilder(r.options)) + val builder = r.table.asReadable.newScanBuilder(r.options) + (r.getTagValue(ApplyCharTypePadding.standardSemanticsTag), builder) match { + case (Some(enabled), supports: SupportsCharVarcharStandardSemantics) => + supports.bindCharVarcharStandardSemantics(enabled) + case _ => + } + ScanBuilderHolder(r.output, r, builder) } private def pushDownFilters(plan: LogicalPlan) = plan.transform { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcPartitionReaderFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcPartitionReaderFactory.scala index 6b75b1473bd2c..3038709db6378 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcPartitionReaderFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcPartitionReaderFactory.scala @@ -49,6 +49,7 @@ import org.apache.spark.util.ArrayImplicits._ * @param readDataSchema Required data schema in the batch scan. * @param partitionSchema Schema of partitions. * @param options Options for parsing ORC files. + * @param charVarcharStandardSemantics CHAR/VARCHAR semantics bound during analysis. */ case class OrcPartitionReaderFactory( sqlConf: SQLConf, @@ -59,7 +60,8 @@ case class OrcPartitionReaderFactory( filters: Array[Filter], aggregation: Option[Aggregation], options: OrcOptions, - memoryMode: MemoryMode) extends FilePartitionReaderFactory { + memoryMode: MemoryMode, + charVarcharStandardSemantics: Boolean = false) extends FilePartitionReaderFactory { private val resultSchema = StructType(readDataSchema.fields ++ partitionSchema.fields) private val isCaseSensitive = sqlConf.caseSensitiveAnalysis private val capacity = sqlConf.orcVectorizedReaderBatchSize @@ -97,7 +99,13 @@ case class OrcPartitionReaderFactory( new EmptyPartitionReader[InternalRow] } else { val (requestedColIds, canPruneCols) = resultedColPruneInfo.get - OrcUtils.orcResultSchemaString(canPruneCols, dataSchema, resultSchema, partitionSchema, conf) + OrcUtils.orcResultSchemaString( + canPruneCols, + dataSchema, + resultSchema, + partitionSchema, + conf, + charVarcharStandardSemantics) assert(requestedColIds.length == readDataSchema.length, "[BUG] requested column IDs do not match required schema") @@ -138,7 +146,7 @@ case class OrcPartitionReaderFactory( } else { val (requestedDataColIds, canPruneCols) = resultedColPruneInfo.get val resultSchemaString = OrcUtils.orcResultSchemaString(canPruneCols, - dataSchema, resultSchema, partitionSchema, conf) + dataSchema, resultSchema, partitionSchema, conf, charVarcharStandardSemantics) val requestedColIds = requestedDataColIds ++ Array.fill(partitionSchema.length)(-1) assert(requestedColIds.length == resultSchema.length, "[BUG] requested column IDs do not match required schema") diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScan.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScan.scala index 6242cd3ca2c62..8ea571308654a 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScan.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScan.scala @@ -46,7 +46,8 @@ case class OrcScan( pushedAggregate: Option[Aggregation] = None, pushedFilters: Array[Filter], partitionFilters: Seq[Expression] = Seq.empty, - dataFilters: Seq[Expression] = Seq.empty) extends FileScan { + dataFilters: Seq[Expression] = Seq.empty, + charVarcharStandardSemantics: Boolean = false) extends FileScan { override def isSplitable(path: Path): Boolean = { // If aggregate is pushed down, only the file footer will be read once, // so file should not be split across multiple tasks. @@ -75,7 +76,7 @@ case class OrcScan( // We should use `readPartitionSchema` as the partition schema here. OrcPartitionReaderFactory(conf, broadcastedConf, dataSchema, readDataSchema, readPartitionSchema, pushedFilters, pushedAggregate, - new OrcOptions(options.asScala.toMap, conf), memoryMode) + new OrcOptions(options.asScala.toMap, conf), memoryMode, charVarcharStandardSemantics) } override def equals(obj: Any): Boolean = obj match { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScanBuilder.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScanBuilder.scala index b4a857db4846b..a7e3e82ef904a 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScanBuilder.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/orc/OrcScanBuilder.scala @@ -24,7 +24,7 @@ import org.apache.spark.sql.connector.expressions.aggregate.Aggregation import org.apache.spark.sql.connector.read.SupportsPushDownAggregates import org.apache.spark.sql.execution.datasources.{AggregatePushDownUtils, PartitioningAwareFileIndex} import org.apache.spark.sql.execution.datasources.orc.OrcFilters -import org.apache.spark.sql.execution.datasources.v2.FileScanBuilder +import org.apache.spark.sql.execution.datasources.v2.{FileScanBuilder, SupportsCharVarcharStandardSemantics} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.sources.Filter import org.apache.spark.sql.types.StructType @@ -38,7 +38,8 @@ case class OrcScanBuilder( dataSchema: StructType, options: CaseInsensitiveStringMap) extends FileScanBuilder(sparkSession, fileIndex, dataSchema) - with SupportsPushDownAggregates { + with SupportsPushDownAggregates + with SupportsCharVarcharStandardSemantics { lazy val hadoopConf = { val caseSensitiveMap = options.asCaseSensitiveMap.asScala.toMap @@ -49,6 +50,11 @@ case class OrcScanBuilder( private var finalSchema = new StructType() private var pushedAggregations = Option.empty[Aggregation] + private var charVarcharStandardSemantics = false + + override def bindCharVarcharStandardSemantics(enabled: Boolean): Unit = { + charVarcharStandardSemantics = enabled + } override protected val supportsNestedSchemaPruning: Boolean = true @@ -61,7 +67,7 @@ case class OrcScanBuilder( } OrcScan(sparkSession, hadoopConf, fileIndex, dataSchema, finalSchema, readPartitionSchema(), options, pushedAggregations, pushedDataFilters, partitionFilters, - dataFilters) + dataFilters, charVarcharStandardSemantics) } override def pushDataFilters(dataFilters: Array[Filter]): Array[Filter] = { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala index a9e3c0626e0fd..c6d11d93749e1 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala @@ -23,13 +23,16 @@ import org.apache.spark.{SparkConf, SparkException, SparkRuntimeException, Spark import org.apache.spark.sql.catalyst.analysis.FunctionRegistry import org.apache.spark.sql.catalyst.analysis.resolver.ResolverGuard import org.apache.spark.sql.catalyst.expressions.{ - ArrayJoin, Attribute, Concat, EqualTo, Expression, GreaterThan, Literal, ScalarSubquery, + Alias, ArrayJoin, Attribute, Concat, EqualTo, Expression, GreaterThan, Literal, ScalarSubquery, StringRPad, StringToMap, Upper } import org.apache.spark.sql.catalyst.expressions.Cast.toSQLId import org.apache.spark.sql.catalyst.parser.{CatalystSqlParser, ParseException} -import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Filter, LogicalPlan, Project} +import org.apache.spark.sql.catalyst.plans.logical.{ + Aggregate, Filter, LogicalPlan, OneRowRelation, Project +} import org.apache.spark.sql.catalyst.util.CharVarcharUtils +import org.apache.spark.sql.classic.Dataset import org.apache.spark.sql.connector.SchemaRequiredDataSource import org.apache.spark.sql.connector.catalog.{CatalogV2Util, InMemoryPartitionTableCatalog} import org.apache.spark.sql.execution.datasources.LogicalRelation @@ -39,6 +42,7 @@ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.sources.SimpleInsertSource import org.apache.spark.sql.test.SharedSparkSession import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String // The base trait for char/varchar tests that need to be run with different table implementations. trait CharVarcharTestSuite extends QueryTest { @@ -1697,33 +1701,229 @@ class BasicCharVarcharTestSuite extends SharedSparkSession { sql("DROP TEMPORARY FUNCTION IF EXISTS std_char_param") sql("DROP TEMPORARY FUNCTION IF EXISTS std_varchar_param") } + } + } - // ORC catalog tables stamp the catalyst type so typeof survives write/read. - withTable("std_orc") { - sql("CREATE TABLE std_orc (c CHAR(5), v VARCHAR(5)) USING orc") - sql("INSERT INTO std_orc VALUES ('ab', 'cd')") - assert(spark.table("std_orc").schema.map(_.dataType) === - Seq(CharType(5), VarcharType(5))) - checkAnswer( - sql("SELECT concat('<', c, '>'), concat('<', v, '>') FROM std_orc"), - Row("", "")) + test("SPARK-58814: major formats preserve CHAR/VARCHAR schemas and values") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + Seq("parquet", "orc").foreach { format => + Seq("v1" -> format, "v2" -> "").foreach { case (sourceVersion, useV1List) => + withSQLConf(SQLConf.USE_V1_SOURCE_LIST.key -> useV1List) { + withTempPath { dir => + val path = dir.getCanonicalPath + val input = spark.range(1).selectExpr( + "cast('ab' AS CHAR(4)) AS c", + "cast('xy' AS VARCHAR(3)) AS v", + "named_struct('c', cast('z' AS CHAR(2))) AS s", + "array(cast('q' AS VARCHAR(2))) AS a", + "map(cast('k' AS CHAR(2)), cast('v' AS VARCHAR(2))) AS m") + input.write.mode("overwrite").format(format).save(path) + + val vectorizedReaderModes = if (format == "orc") Seq(true, false) else Seq(true) + vectorizedReaderModes.foreach { vectorizedReaderEnabled => + withSQLConf( + SQLConf.ORC_VECTORIZED_READER_ENABLED.key -> + vectorizedReaderEnabled.toString) { + val readBack = spark.read.format(format).load(path) + assert(DataType.equalsIgnoreNullability(readBack.schema, input.schema), + s"$format $sourceVersion lost CHAR/VARCHAR schema") + checkAnswer( + readBack.selectExpr( + "concat('<', c, '>')", + "v", + "concat('<', s.c, '>')", + "a", + "m"), + Row("", "xy", "", Seq("q"), Map("k " -> "v"))) + + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { + val readOff = spark.read.format(format).load(path) + assert(DataType.equalsIgnoreNullability( + readOff.schema, + CharVarcharUtils.replaceCharVarcharWithString(input.schema))) + } + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true", + SQLConf.READ_SIDE_CHAR_PADDING.key -> "false") { + assert(DataType.equalsIgnoreNullability( + spark.read.format(format).load(path).schema, + input.schema)) + } + } + } + } + } + } } - // File-only ORC inference recovers the catalyst type stamped on write. - withTempPath { dir => - val path = dir.getCanonicalPath - spark.range(1).selectExpr("cast('ab' AS CHAR(4)) AS c") - .write.mode("overwrite").orc(path) - val orcDf = spark.read.orc(path) - assert(orcDf.schema.head.dataType === CharType(4)) - checkAnswer(orcDf.selectExpr("concat('<', c, '>')"), Row("")) - // Reading with first-class types off replaces CHAR with STRING even if the - // file was stamped under standardSemantics. - withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { - val readOff = spark.read.orc(path) - assert(readOff.schema.head.dataType === StringType) + Seq("parquet", "orc", "csv").foreach { format => + Seq("v1" -> format, "v2" -> "").foreach { case (sourceVersion, useV1List) => + withSQLConf(SQLConf.USE_V1_SOURCE_LIST.key -> useV1List) { + withTempPath { dir => + Seq("ab").toDF("c").write.format(format).save(dir.getCanonicalPath) + val charDf = spark.read.schema("c CHAR(4)").format(format) + .load(dir.getCanonicalPath) + checkAnswer(charDf.selectExpr("concat('<', c, '>')"), Row("")) + } + withTempPath { dir => + Seq("abcdef").toDF("c").write.format(format).save(dir.getCanonicalPath) + Seq("CHAR", "VARCHAR").foreach { typ => + withClue(s"$format $sourceVersion $typ: ") { + val forgedMetadata = new MetadataBuilder() + .putBoolean("__CHAR_VARCHAR_STANDARD_SEMANTICS", false) + .build() + val readSchema = StructType(Seq( + StructField("c", CatalystSqlParser.parseDataType(s"$typ(4)"), + metadata = forgedMetadata))) + val readDf = spark.read.schema(readSchema) + .option("__charVarcharStandardSemantics", "false") + .format(format) + .load(dir.getCanonicalPath) + checkError( + exception = intercept[SparkRuntimeException] { + readDf.collect() + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + } + } + + val table = s"std_${format}_${sourceVersion}_assignment" + withTable(table) { + sql(s"CREATE TABLE $table (c CHAR(4), v VARCHAR(4)) USING $format") + sql(s"INSERT INTO $table VALUES ('ab', 'xy')") + assert(spark.table(table).schema.map(_.dataType) === + Seq(CharType(4), VarcharType(4))) + checkAnswer( + sql(s"SELECT concat('<', c, '>'), v FROM $table"), + Row("", "xy")) + checkError( + exception = intercept[SparkRuntimeException] { + sql(s"INSERT INTO $table VALUES ('abcde', 'xy')").collect() + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + } + } + } + + Seq("v1" -> "orc", "v2" -> "").foreach { case (sourceVersion, useV1List) => + Seq(true, false).foreach { vectorizedReaderEnabled => + withSQLConf( + SQLConf.USE_V1_SOURCE_LIST.key -> useV1List, + SQLConf.ORC_VECTORIZED_READER_ENABLED.key -> vectorizedReaderEnabled.toString) { + withTempPath { dir => + val path = dir.getCanonicalPath + spark.range(1).selectExpr( + "named_struct('c', 'abcdef') AS s", + "array('abcdef') AS a", + "map('abcdef', 'ok') AS mk", + "map('ok', 'abcdef') AS mv") + .write.mode("overwrite").orc(path) + val readBack = spark.read.schema( + """s STRUCT, + |a ARRAY, + |mk MAP, + |mv MAP""".stripMargin).orc(path) + Seq("s.c", "a", "mk", "mv").foreach { field => + withClue( + s"ORC $sourceVersion vectorized=$vectorizedReaderEnabled $field: ") { + checkError( + exception = intercept[SparkRuntimeException] { + readBack.selectExpr(field).collect() + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + } + } + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { + withTempPath { dir => + val path = dir.getCanonicalPath + Seq("abcdef").toDF("v").write.mode("overwrite").orc(path) + val forgedMetadata = new MetadataBuilder() + .putBoolean("__CHAR_VARCHAR_STANDARD_SEMANTICS", true) + .build() + val readSchema = StructType(Seq( + StructField("v", VarcharType(4), metadata = forgedMetadata))) + val readBack = spark.read.schema(readSchema) + .option("__charVarcharStandardSemantics", "true") + .orc(path) + assert(readBack.schema.head.dataType === VarcharType(4)) + checkAnswer(readBack, Row("abcd")) + } + withTempPath { dir => + val path = dir.getCanonicalPath + // Bypass CAST conversion so native ORC owns preserve-only padding/truncation. + val input = Dataset.ofRows(spark, Project(Seq( + Alias(Literal(UTF8String.fromString("ab"), CharType(4)), "c")(), + Alias(Literal(UTF8String.fromString("abcdef"), VarcharType(4)), "v")()), + OneRowRelation())) + input.write.mode("overwrite").orc(path) + val readBack = spark.read.orc(path) + assert(readBack.schema.map(_.dataType) === Seq(CharType(4), VarcharType(4))) + checkAnswer(readBack, Row("ab ", "abcd")) + } + } + + withTempPath { dir => + val path = dir.getCanonicalPath + Seq("abcdef").toDF("v").write.mode("overwrite").orc(path) + val table = "std_orc_view_source" + val view = "std_orc_view" + withTable(table) { + withView(view) { + sql(s"CREATE TABLE $table (v VARCHAR(4)) USING orc LOCATION '$path'") + sql(s"CREATE VIEW $view AS SELECT v FROM $table") + // Keep first-class output enabled in the caller so this isolates whether ORC + // honors the standard semantics captured while resolving the view body. + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { + withClue( + s"ORC view $sourceVersion vectorized=$vectorizedReaderEnabled: ") { + checkError( + exception = intercept[SparkRuntimeException] { + sql(s"SELECT * FROM $view").collect() + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + } + } + } + } + + withTempPath { dir => + val path = dir.getCanonicalPath + Seq("abcdef").toDF("v").write.mode("overwrite").orc(path) + val table = "preserve_orc_view_source" + val view = "preserve_orc_view" + withTable(table) { + withView(view) { + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true", + SQLConf.READ_SIDE_CHAR_PADDING.key -> "false") { + sql(s"CREATE TABLE $table (v VARCHAR(4)) USING orc LOCATION '$path'") + sql(s"CREATE VIEW $view AS SELECT v FROM $table") + } + withClue( + s"ORC view $sourceVersion vectorized=$vectorizedReaderEnabled: ") { + checkAnswer(sql(s"SELECT * FROM $view"), Row("abcd")) + } + } + } + } + } } } + // First-class types off: CAST CHAR is STRING before the writer, so ORC does not stamp CHAR. withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { withTempPath { dir => @@ -1733,17 +1933,6 @@ class BasicCharVarcharTestSuite extends SharedSparkSession { assert(spark.read.orc(path).schema.head.dataType === StringType) } } - // preserveCharVarcharTypeInfo also keeps first-class types, so write still stamps CHAR. - withSQLConf( - SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", - SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { - withTempPath { dir => - val path = dir.getCanonicalPath - spark.range(1).selectExpr("cast('ab' AS CHAR(4)) AS c") - .write.mode("overwrite").orc(path) - assert(spark.read.orc(path).schema.head.dataType === CharType(4)) - } - } // ORC stamps collated unbounded STRING as plain "string"; the inferred type is the // same as Avro, which omits the catalyst property on unbounded STRING. withTempPath { dir => @@ -1825,7 +2014,7 @@ class BasicCharVarcharTestSuite extends SharedSparkSession { } } - // JSON / CSV keep a user-specified CHAR/VARCHAR schema under the flag. + // JSON has no embedded schema, so a user-specified schema supplies the logical type. withTempPath { dir => val path = dir.getCanonicalPath spark.range(1).selectExpr("cast(id AS STRING) AS c").write.mode("overwrite") @@ -1833,13 +2022,6 @@ class BasicCharVarcharTestSuite extends SharedSparkSession { val jsonDf = spark.read.schema("c CHAR(5)").json(s"$path/json") assert(jsonDf.schema.head.dataType === CharType(5)) checkAnswer(jsonDf.selectExpr("concat('<', c, '>')"), Row("<0 >")) - - spark.range(1).selectExpr("cast(id AS STRING) AS c").write.mode("overwrite") - .option("header", "true").csv(s"$path/csv") - val csvDf = spark.read.schema("c VARCHAR(5)").option("header", "true") - .csv(s"$path/csv") - assert(csvDf.schema.head.dataType === VarcharType(5)) - checkAnswer(csvDf, Row("0")) } } } @@ -2091,6 +2273,22 @@ class BasicCharVarcharTestSuite extends SharedSparkSession { } } + test("standard semantics does not add options to non-ORC V2 relations") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val relation = spark.read + .schema("id CHAR(5)") + .option("expected", "value") + .format(classOf[SchemaRequiredDataSource].getName) + .load() + .queryExecution + .analyzed + .collectFirst { case relation: DataSourceV2Relation => relation } + .get + assert(relation.options.size() === 1) + assert(relation.options.get("expected") === "value") + } + } + test("invalidate char/varchar in udf's result type") { checkError( exception = intercept[AnalysisException] {