Skip to content
Open
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 @@ -41,7 +41,7 @@ import org.apache.spark.internal.Logging
import org.apache.spark.internal.LogKeys.PATH
import org.apache.spark.sql.{SPARK_VERSION_METADATA_KEY, SparkSession}
import org.apache.spark.sql.catalyst.{FileSourceOptions, InternalRow}
import org.apache.spark.sql.catalyst.analysis.caseSensitiveResolution
import org.apache.spark.sql.catalyst.analysis.{caseInsensitiveResolution, caseSensitiveResolution}
import org.apache.spark.sql.catalyst.expressions.JoinedRow
import org.apache.spark.sql.catalyst.parser.CatalystSqlParser
import org.apache.spark.sql.catalyst.util.{quoteIdentifier, CaseInsensitiveMap, CharVarcharUtils}
Expand Down Expand Up @@ -585,7 +585,9 @@ object OrcUtils extends Logging {
partitionSchema: StructType,
aggregation: Aggregation,
aggSchema: StructType,
partitionValues: InternalRow): InternalRow = {
partitionValues: InternalRow,
isCaseSensitive: Boolean,
conf: Configuration): InternalRow = {
var columnsStatistics: OrcColumnStatistics = null
try {
columnsStatistics = OrcFooterReader.readStatistics(reader)
Expand All @@ -595,10 +597,23 @@ object OrcUtils extends Logging {
s"ORC aggregate push down by setting 'spark.sql.orc.aggregatePushdown' to false.", e)
}

// Get column statistics with column name.
def getColumnStatistics(columnName: String): ColumnStatistics = {
val columnIndex = dataSchema.getFieldIndex(columnName).getOrElse(-1)
columnsStatistics.get(columnIndex).getStatistics
// Resolve to the file's own ordinal, mirroring `requestedColumnIds`; -1 if absent.
val orcFieldNames = reader.getSchema.getFieldNames.asScala
val forcePositionalEvolution = OrcConf.FORCE_POSITIONAL_EVOLUTION.getBoolean(conf)
def fileColumnIndex(columnName: String): Int = {
if (forcePositionalEvolution || orcFieldNames.forall(_.startsWith("_col"))) {
val index = dataSchema.getFieldIndex(columnName).getOrElse(-1)
if (index >= 0 && index < orcFieldNames.length) index else -1
} else if (isCaseSensitive) {
orcFieldNames.indexWhere(caseSensitiveResolution(_, columnName))
} else {
orcFieldNames.indexWhere(caseInsensitiveResolution(_, columnName))
}
}

def getColumnStatistics(columnName: String): Option[ColumnStatistics] = {
val columnIndex = fileColumnIndex(columnName)
if (columnIndex >= 0) Some(columnsStatistics.get(columnIndex).getStatistics) else None
}

// Get Min/Max statistics and store as ORC `WritableComparable` format.
Expand Down Expand Up @@ -657,14 +672,14 @@ object OrcUtils extends Logging {
aggregation.aggregateExpressions.zipWithIndex.map {
case (max: Max, index) if V2ColumnUtils.extractV2Column(max.column).isDefined =>
val columnName = V2ColumnUtils.extractV2Column(max.column).get
val statistics = getColumnStatistics(columnName)
val dataType = schemaWithoutGroupBy(index).dataType
getMinMaxFromColumnStatistics(statistics, dataType, isMax = true)
getColumnStatistics(columnName)
.map(getMinMaxFromColumnStatistics(_, dataType, isMax = true)).orNull
case (min: Min, index) if V2ColumnUtils.extractV2Column(min.column).isDefined =>
val columnName = V2ColumnUtils.extractV2Column(min.column).get
val statistics = getColumnStatistics(columnName)
val dataType = schemaWithoutGroupBy.apply(index).dataType
getMinMaxFromColumnStatistics(statistics, dataType, isMax = false)
getColumnStatistics(columnName)
.map(getMinMaxFromColumnStatistics(_, dataType, isMax = false)).orNull
case (count: Count, _) if V2ColumnUtils.extractV2Column(count.column).isDefined =>
val columnName = V2ColumnUtils.extractV2Column(count.column).get
val isPartitionColumn = partitionSchema.fields.map(_.name).contains(columnName)
Expand All @@ -675,7 +690,7 @@ object OrcUtils extends Logging {
val nonNullRowsCount = if (isPartitionColumn) {
columnsStatistics.getStatistics.getNumberOfValues
} else {
getColumnStatistics(columnName).getNumberOfValues
getColumnStatistics(columnName).map(_.getNumberOfValues).getOrElse(0L)
}
new LongWritable(nonNullRowsCount)
case (_: CountStar, _) =>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ import org.apache.spark.internal.Logging
import org.apache.spark.internal.LogKeys.{CLASS_NAME, CONFIG}
import org.apache.spark.sql.{Row, SparkSession}
import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.catalyst.analysis.{caseInsensitiveResolution, caseSensitiveResolution}
import org.apache.spark.sql.catalyst.expressions.JoinedRow
import org.apache.spark.sql.catalyst.expressions.variant.VariantExpressionEvalUtils
import org.apache.spark.sql.catalyst.util.RebaseDateTime.RebaseSpec
Expand Down Expand Up @@ -249,24 +250,24 @@ object ParquetUtils extends Logging {
private[sql] def createAggInternalRowFromFooter(
footer: ParquetMetadata,
filePath: String,
dataSchema: StructType,
partitionSchema: StructType,
aggregation: Aggregation,
aggSchema: StructType,
partitionValues: InternalRow,
isCaseSensitive: Boolean,
datetimeRebaseSpec: RebaseSpec): InternalRow = {
// if there are group by columns, we will build result row first,
// and then append group by columns values (partition columns values) to the result row.
val schemaWithoutGroupBy =
AggregatePushDownUtils.getSchemaWithoutGroupingExpression(aggSchema, aggregation)

val (primitiveTypes, values) = getPushedDownAggResult(
footer, filePath, dataSchema, partitionSchema, aggregation)
footer, filePath, partitionSchema, aggregation, schemaWithoutGroupBy, isCaseSensitive)

val builder = Types.buildMessage
primitiveTypes.foreach(t => builder.addField(t))
val parquetSchema = builder.named("root")

// if there are group by columns, we will build result row first,
// and then append group by columns values (partition columns values) to the result row.
val schemaWithoutGroupBy =
AggregatePushDownUtils.getSchemaWithoutGroupingExpression(aggSchema, aggregation)

val schemaConverter = new ParquetToSparkSchemaConverter
val converter = new ParquetRowConverter(
schemaConverter,
Expand All @@ -276,8 +277,11 @@ object ParquetUtils extends Logging {
datetimeRebaseSpec,
RebaseSpec(LegacyBehaviorPolicy.CORRECTED),
NoopUpdater)
// Reset fields to null so a column absent from the file reads as null, not the zero default.
converter.start()
val primitiveTypeNames = primitiveTypes.map(_.getPrimitiveTypeName)
primitiveTypeNames.zipWithIndex.foreach {
case (_, i) if values(i) == null =>
case (PrimitiveType.PrimitiveTypeName.BOOLEAN, i) =>
val v = values(i).asInstanceOf[Boolean]
converter.getConverter(i).asPrimitiveConverter.addBoolean(v)
Expand Down Expand Up @@ -323,54 +327,65 @@ object ParquetUtils extends Logging {
private[sql] def getPushedDownAggResult(
footer: ParquetMetadata,
filePath: String,
dataSchema: StructType,
partitionSchema: StructType,
aggregation: Aggregation)
aggregation: Aggregation,
aggSchema: StructType,
isCaseSensitive: Boolean)
: (Array[PrimitiveType], Array[Any]) = {
val footerFileMetaData = footer.getFileMetaData
val fields = footerFileMetaData.getSchema.getFields
// Resolve by name in the file's own schema; a positional lookup breaks under mergeSchema.
val fileSchema = footerFileMetaData.getSchema
val resolver = if (isCaseSensitive) caseSensitiveResolution else caseInsensitiveResolution
def fileFieldIndex(colName: String): Int =
(0 until fileSchema.getFieldCount)
.indexWhere(i => resolver(fileSchema.getFieldName(i), colName))
val blocks = footer.getBlocks
val primitiveTypeBuilder = mutable.ArrayBuilder.make[PrimitiveType]
val valuesBuilder = mutable.ArrayBuilder.make[Any]
lazy val sparkToParquet = new SparkToParquetSchemaConverter(SQLConf.get)

aggregation.aggregateExpressions.foreach { agg =>
aggregation.aggregateExpressions.zipWithIndex.foreach { case (agg, aggIndex) =>
var value: Any = None
var rowCount = 0L
var isCount = false
var index = 0
var index = -1
var schemaName = ""
blocks.forEach { block =>
val blockMetaData = block.getColumns
agg match {
case max: Max if V2ColumnUtils.extractV2Column(max.column).isDefined =>
val colName = V2ColumnUtils.extractV2Column(max.column).get
index = dataSchema.getFieldIndex(colName).getOrElse(-1)
index = fileFieldIndex(colName)
schemaName = "max(" + colName + ")"
val currentMax = getCurrentBlockMaxOrMin(filePath, blockMetaData, index, true)
if (value == None || currentMax.asInstanceOf[Comparable[Any]].compareTo(value) > 0) {
value = currentMax
if (index >= 0) {
val currentMax = getCurrentBlockMaxOrMin(filePath, blockMetaData, index, true)
if (value == None || currentMax.asInstanceOf[Comparable[Any]].compareTo(value) > 0) {
value = currentMax
}
}
case min: Min if V2ColumnUtils.extractV2Column(min.column).isDefined =>
val colName = V2ColumnUtils.extractV2Column(min.column).get
index = dataSchema.getFieldIndex(colName).getOrElse(-1)
index = fileFieldIndex(colName)
schemaName = "min(" + colName + ")"
val currentMin = getCurrentBlockMaxOrMin(filePath, blockMetaData, index, false)
if (value == None || currentMin.asInstanceOf[Comparable[Any]].compareTo(value) < 0) {
value = currentMin
if (index >= 0) {
val currentMin = getCurrentBlockMaxOrMin(filePath, blockMetaData, index, false)
if (value == None || currentMin.asInstanceOf[Comparable[Any]].compareTo(value) < 0) {
value = currentMin
}
}
case count: Count if V2ColumnUtils.extractV2Column(count.column).isDefined =>
val colName = V2ColumnUtils.extractV2Column(count.column).get
schemaName = "count(" + colName + ")"
rowCount += block.getRowCount
var isPartitionCol = false
if (partitionSchema.getFieldIndex(colName).isDefined) {
isPartitionCol = true
}
val isPartitionCol = partitionSchema.getFieldIndex(colName).isDefined
isCount = true
if (!isPartitionCol) {
index = dataSchema.getFieldIndex(colName).getOrElse(-1)
// Count(*) includes the null values, but Count(colName) doesn't.
rowCount -= getNumNulls(filePath, blockMetaData, index)
index = fileFieldIndex(colName)
if (index >= 0) {
rowCount -= getNumNulls(filePath, blockMetaData, index)
} else {
rowCount -= block.getRowCount
}
}
case _: CountStar =>
schemaName = "count(*)"
Expand All @@ -381,14 +396,21 @@ object ParquetUtils extends Logging {
}
if (isCount) {
valuesBuilder += rowCount
primitiveTypeBuilder += Types.required(PrimitiveTypeName.INT64).named(schemaName);
} else {
primitiveTypeBuilder += Types.required(PrimitiveTypeName.INT64).named(schemaName)
} else if (index >= 0) {
valuesBuilder += value
val field = fields.get(index)
val field = fileSchema.getFields.get(index)
primitiveTypeBuilder += Types.required(field.asPrimitiveType.getPrimitiveTypeName)
.as(field.getLogicalTypeAnnotation)
.length(field.asPrimitiveType.getTypeLength)
.named(schemaName)
} else {
// Absent column: null result; the type comes from the requested Spark type.
valuesBuilder += null
val sparkType = aggSchema(aggIndex).dataType
val parquetType = sparkToParquet.convertField(
StructField(schemaName, sparkType, nullable = true), inShredded = false)
primitiveTypeBuilder += parquetType.asPrimitiveType()
}
}
(primitiveTypeBuilder.result(), valuesBuilder.result())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -194,7 +194,7 @@ case class OrcPartitionReaderFactory(
Utils.tryWithResource(createORCReader(filePath, conf)._1) { reader =>
OrcUtils.createAggInternalRowFromFooter(
reader, filePath.toString, dataSchema, partitionSchema, aggregation.get,
readDataSchema, file.partitionValues)
readDataSchema, file.partitionValues, isCaseSensitive, conf)
}
}

Expand Down Expand Up @@ -222,7 +222,7 @@ case class OrcPartitionReaderFactory(
Utils.tryWithResource(createORCReader(filePath, conf)._1) { reader =>
val row = OrcUtils.createAggInternalRowFromFooter(
reader, filePath.toString, dataSchema, partitionSchema, aggregation.get,
readDataSchema, file.partitionValues)
readDataSchema, file.partitionValues, isCaseSensitive, conf)
AggregatePushDownUtils.convertAggregatesRowToBatch(row, readDataSchema, offHeap = false)
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -141,8 +141,8 @@ case class ParquetPartitionReaderFactory(

if (openedFooter.footer != null && !openedFooter.footer.getBlocks.isEmpty) {
ParquetUtils.createAggInternalRowFromFooter(openedFooter.footer,
file.urlEncodedPath, dataSchema, partitionSchema, aggregation.get,
readDataSchema, file.partitionValues,
file.urlEncodedPath, partitionSchema, aggregation.get,
readDataSchema, file.partitionValues, isCaseSensitive,
getDatetimeRebaseSpec(openedFooter.footer.getFileMetaData))
} else {
null
Expand Down Expand Up @@ -187,8 +187,8 @@ case class ParquetPartitionReaderFactory(

if (openedFooter.footer != null && !openedFooter.footer.getBlocks.isEmpty) {
val row = ParquetUtils.createAggInternalRowFromFooter(openedFooter.footer,
file.urlEncodedPath, dataSchema, partitionSchema, aggregation.get,
readDataSchema, file.partitionValues,
file.urlEncodedPath, partitionSchema, aggregation.get,
readDataSchema, file.partitionValues, isCaseSensitive,
getDatetimeRebaseSpec(openedFooter.footer.getFileMetaData))
AggregatePushDownUtils.convertAggregatesRowToBatch(
row, readDataSchema, enableOffHeapColumnVector && Option(TaskContext.get()).isDefined)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -773,6 +773,56 @@ class ParquetV2AggregatePushDownSuite extends ParquetAggregatePushDownSuite {
}
}

test("SPARK-59609: aggregate push-down reads statistics from the correct column " +
"after schema merging") {
Seq("false", "true").foreach { enableVectorizedReader =>
withTempPath { dir =>
val path = dir.getCanonicalPath
spark.sql("SELECT 1 AS id, 10 AS value, 100 AS other")
.union(spark.sql("SELECT 2 AS id, 20 AS value, 200 AS other"))
.coalesce(1).write.parquet(path + "/with_value")
// Missing `value`; after merging, `other` sits where `value` is in the merged schema.
spark.sql("SELECT 3 AS id, 300 AS other")
.union(spark.sql("SELECT 4 AS id, 400 AS other"))
.coalesce(1).write.parquet(path + "/without_value")
withTempView("t") {
spark.read.option("mergeSchema", "true").option("recursiveFileLookup", "true")
.parquet(path).createOrReplaceTempView("t")
withSQLConf(
aggPushDownEnabledKey -> "true",
vectorizedReaderEnabledKey -> enableVectorizedReader) {
checkAnswer(
sql("SELECT COUNT(*), COUNT(value), MIN(value), MAX(value) FROM t"),
Row(4, 2, 10, 20))
}
}
}
}
}

test("SPARK-59609: aggregate push-down resolves the column case-insensitively") {
Seq("false", "true").foreach { enableVectorizedReader =>
withTempPath { dir =>
val path = dir.getCanonicalPath
// Physical column is `value`, but the read schema names it `VALUE`.
spark.sql("SELECT 1 AS id, 10 AS value")
.union(spark.sql("SELECT 2 AS id, 20 AS value"))
.coalesce(1).write.parquet(path)
withTempView("t") {
spark.read.schema("id INT, VALUE INT").parquet(path).createOrReplaceTempView("t")
withSQLConf(
SQLConf.CASE_SENSITIVE.key -> "false",
aggPushDownEnabledKey -> "true",
vectorizedReaderEnabledKey -> enableVectorizedReader) {
checkAnswer(
sql("SELECT COUNT(VALUE), MIN(VALUE), MAX(VALUE) FROM t"),
Row(2, 10, 20))
}
}
}
}
}

// The error originates in the executor-side partition reader, so it may be wrapped in a
// higher-level exception. Walk the cause chain to find the structured Spark exception.
private def interceptAggPushDownError(query: String): SparkUnsupportedOperationException = {
Expand Down Expand Up @@ -806,4 +856,32 @@ class OrcV2AggregatePushDownSuite extends OrcAggregatePushDownSuite {

override protected def sparkConf: SparkConf =
super.sparkConf.set(SQLConf.USE_V1_SOURCE_LIST, "")

// ORC has no mergeSchema, but a divergent physical column order exercises the same bug.
test("SPARK-59609: aggregate push-down reads statistics from the correct column " +
"when the file column order differs from the read schema") {
Seq("false", "true").foreach { enableVectorizedReader =>
withTempPath { dir =>
val path = dir.getCanonicalPath
spark.sql("SELECT 1 AS id, 10 AS value, 100 AS other")
.union(spark.sql("SELECT 2 AS id, 20 AS value, 200 AS other"))
.coalesce(1).write.orc(path + "/order1")
// Columns laid out as (id, other, value), so `value` sits at a different physical ordinal.
spark.sql("SELECT 3 AS id, 300 AS other, 30 AS value")
.union(spark.sql("SELECT 4 AS id, 400 AS other, 40 AS value"))
.coalesce(1).write.orc(path + "/order2")
withTempView("t") {
spark.read.schema("id INT, value INT, other INT")
.option("recursiveFileLookup", "true").orc(path).createOrReplaceTempView("t")
withSQLConf(
aggPushDownEnabledKey -> "true",
vectorizedReaderEnabledKey -> enableVectorizedReader) {
checkAnswer(
sql("SELECT COUNT(*), COUNT(value), MIN(value), MAX(value) FROM t"),
Row(4, 4, 10, 40))
}
}
}
}
}
}