From c03dfadf9af76d2d90f40ec788e7ee91b6b4701d Mon Sep 17 00:00:00 2001 From: Eric Yang Date: Sun, 20 Sep 2026 22:22:54 -0700 Subject: [PATCH] [SPARK-59684][SQL] Fix pivot() failing on a struct column when the values are collected --- .../connect/planner/SparkConnectPlanner.scala | 1 - .../planner/SparkConnectProtoSuite.scala | 18 +++++++++++++++++- .../sql/classic/RelationalGroupedDataset.scala | 12 ++++++++---- .../apache/spark/sql/DataFramePivotSuite.scala | 18 +++++++++++++++++- .../sql/errors/QueryExecutionErrorsSuite.scala | 4 +++- 5 files changed, 45 insertions(+), 8 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala index 0d11148a2441d..f7b56e20aac8d 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala @@ -2788,7 +2788,6 @@ class SparkConnectPlanner( } else { RelationalGroupedDataset .collectPivotValues(Dataset.ofRows(session, logicalPlan), Column(pivotExpr)) - .map(expressions.Literal.apply) } logical.Pivot( groupByExprsOpt = Some(groupingExpressionsWithOrdinals.map(toNamedExpression)), diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/planner/SparkConnectProtoSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/planner/SparkConnectProtoSuite.scala index 1a5189e243049..20c455f1c5e97 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/planner/SparkConnectProtoSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/planner/SparkConnectProtoSuite.scala @@ -32,8 +32,8 @@ import org.apache.spark.sql.catalyst.expressions.{AttributeReference, GenericInt import org.apache.spark.sql.catalyst.plans.{FullOuter, Inner, LeftAnti, LeftOuter, LeftSemi, PlanTest, RightOuter} import org.apache.spark.sql.catalyst.plans.logical.{CollectMetrics, Deduplicate, DeduplicateWithinWatermark, Distinct, LocalRelation, LogicalPlan} import org.apache.spark.sql.catalyst.types.DataTypeUtils +import org.apache.spark.sql.classic.{DataFrame, Dataset} import org.apache.spark.sql.classic.ClassicConversions._ -import org.apache.spark.sql.classic.DataFrame import org.apache.spark.sql.connect.common.InvalidPlanInput import org.apache.spark.sql.connect.common.LiteralValueProtoConverter.toLiteralProto import org.apache.spark.sql.connect.dsl.MockRemoteSession @@ -325,6 +325,22 @@ class SparkConnectProtoSuite extends PlanTest with SparkConnectPlanTest { comparePlans(connectPlan2, sparkPlan2) } + test("SPARK-59684: pivot by a struct column without explicit values") { + val schema = new StructType() + .add("v", IntegerType) + .add("s", new StructType().add("a", IntegerType)) + val data = Seq(1, 2).map { i => + new GenericInternalRow(Array[Any](i, new GenericInternalRow(Array[Any](i * 10)))) + } + val connectPlan = + createLocalRelationProto(schema, data).pivot("v".protoAttr)("s".protoAttr, Seq.empty)( + proto_min(proto.Expression.newBuilder().setLiteral(toLiteralProto(1)).build()) + .as("agg1")) + val result = Dataset.ofRows(spark, transform(connectPlan)) + assert(result.columns.toSeq === Seq("v", "{10}", "{20}")) + assert(result.orderBy("v").collect().toSeq === Seq(Row(1, 1, null), Row(2, null, 1))) + } + test("GroupingSets expressions") { val connectPlan1 = connectTestRelation.groupingSets(Seq(Seq("id".protoAttr), Seq.empty), "id".protoAttr)( diff --git a/sql/core/src/main/scala/org/apache/spark/sql/classic/RelationalGroupedDataset.scala b/sql/core/src/main/scala/org/apache/spark/sql/classic/RelationalGroupedDataset.scala index bd7b3348b9f09..d93472fe5ea88 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/classic/RelationalGroupedDataset.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/classic/RelationalGroupedDataset.scala @@ -23,6 +23,7 @@ import org.apache.spark.api.python.PythonEvalType import org.apache.spark.broadcast.Broadcast import org.apache.spark.sql import org.apache.spark.sql.{AnalysisException, Column, Encoder} +import org.apache.spark.sql.catalyst.CatalystTypeConverters import org.apache.spark.sql.catalyst.analysis.{UnresolvedAlias, UnresolvedAttribute, UnresolvedOrdinal} import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate._ @@ -193,7 +194,7 @@ class RelationalGroupedDataset protected[sql]( /** @inheritdoc */ override def pivot(pivotColumn: Column): RelationalGroupedDataset = - pivot(pivotColumn, collectPivotValues(df, pivotColumn)) + pivot(pivotColumn, collectPivotValues(df, pivotColumn).map(Column(_))) /** @inheritdoc */ def pivot(pivotColumn: Column, values: Seq[Any]): RelationalGroupedDataset = { @@ -664,7 +665,7 @@ private[sql] object RelationalGroupedDataset { case expr: Expression => Alias(expr, toPrettySQL(expr))() } - private[sql] def collectPivotValues(df: DataFrame, pivotColumn: Column): Seq[Any] = { + private[sql] def collectPivotValues(df: DataFrame, pivotColumn: Column): Seq[Literal] = { if (df.isStreaming) { throw new AnalysisException( errorClass = "_LEGACY_ERROR_TEMP_3063", @@ -673,12 +674,15 @@ private[sql] object RelationalGroupedDataset { // This is to prevent unintended OOM errors when the number of distinct values is large val maxValues = df.sparkSession.sessionState.conf.dataFramePivotMaxValues // Get the distinct values of the column and sort them so its consistent - val values = df.select(pivotColumn) + val pivotDf = df.select(pivotColumn) + val dataType = pivotDf.schema.head.dataType + val toCatalyst = CatalystTypeConverters.createToCatalystConverter(dataType) + val values = pivotDf .distinct() .limit(maxValues + 1) .sort(pivotColumn) // ensure that the output columns are in a consistent logical order .collect() - .map(_.get(0)) + .map(row => Literal(toCatalyst(row.get(0)), dataType)) .toImmutableArraySeq if (values.length > maxValues) { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFramePivotSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFramePivotSuite.scala index ee1ad5094f6a5..affe08f0a7c26 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFramePivotSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFramePivotSuite.scala @@ -17,7 +17,7 @@ package org.apache.spark.sql -import java.time.LocalDateTime +import java.time.{LocalDateTime, Year} import java.util.Locale import org.apache.spark.sql.catalyst.expressions.aggregate.PivotFirst @@ -342,6 +342,22 @@ class DataFramePivotSuite extends SharedSparkSession { checkAnswer(actual, expected) } + test("SPARK-59684: pivoting by a struct column") { + val df = Seq(1.0d, 2.0d).toDF("v").selectExpr("v", "struct(v, v) AS s", "array(struct(v)) AS a") + checkAnswer( + df.groupBy("v").pivot("s").count(), + Row(1.0d, 1L, null) :: Row(2.0d, null, 1L) :: Nil) + checkAnswer( + df.groupBy("v").pivot("a").agg(first("v")), + Row(1.0d, 1.0d, null) :: Row(2.0d, null, 2.0d) :: Nil) + val udtDf = spark.createDataFrame( + spark.sparkContext.parallelize(Seq(Row(1, Row(Year.of(2020))), Row(2, Row(Year.of(2021))))), + new StructType().add("v", IntegerType).add("s", new StructType().add("y", new YearUDT))) + checkAnswer( + udtDf.groupBy("v").pivot("s").count(), + Row(1, 1L, null) :: Row(2, null, 1L) :: Nil) + } + test("SPARK-35480: percentile_approx should work with pivot") { val actual = Seq( ("a", -1.0), ("a", 5.5), ("a", 2.5), ("b", 3.0), ("b", 5.2)).toDF("type", "value") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala index 218c74c6efcdd..5c58ba302fcc5 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala @@ -270,7 +270,9 @@ class QueryExecutionErrorsSuite val e2 = intercept[SparkRuntimeException] { trainingSales .groupBy($"sales.year") - .pivot(struct(lower(trainingSales("sales.course")), trainingSales("training"))) + .pivot( + struct(lower(trainingSales("sales.course")), trainingSales("training")), + Seq(Row("dotnet", "Dummies"))) .agg(sum($"sales.earnings")) .collect() }