@@ -19,7 +19,7 @@ package org.apache.spark.sql.execution.datasources.v2
1919import org .apache .spark .rdd .RDD
2020import org .apache .spark .sql .catalyst .InternalRow
2121import org .apache .spark .sql .catalyst .expressions ._
22- import org .apache .spark .sql .catalyst .plans .physical .{KeyedPartitioning , Partitioning , SinglePartition }
22+ import org .apache .spark .sql .catalyst .plans .physical .{KeyedPartitioning , Partitioning , SinglePartition , UnknownPartitioning }
2323import org .apache .spark .sql .catalyst .util .{truncatedString , InternalRowComparableWrapper }
2424import org .apache .spark .sql .connector .catalog .Table
2525import org .apache .spark .sql .connector .read ._
@@ -108,7 +108,21 @@ abstract class AbstractBatchScanExec(
108108 }
109109
110110 override def outputPartitioning : Partitioning = {
111- super .outputPartitioning match {
111+ val basePartitioning = super .outputPartitioning match {
112+ case p : UnknownPartitioning
113+ if spjParams.keyGroupedPartitioning.isDefined &&
114+ inputPartitions.nonEmpty && inputPartitions.forall(_.isInstanceOf [HasPartitionKey ]) &&
115+ KeyedPartitioning .supportsExpressions(spjParams.keyGroupedPartitioning.get) =>
116+ val expressions = spjParams.keyGroupedPartitioning.get
117+ val keyRowOrdering = RowOrdering .createNaturalAscendingOrdering(expressions.map(_.dataType))
118+ val partitionKeys = inputPartitions
119+ .map(_.asInstanceOf [HasPartitionKey ].partitionKey())
120+ .sorted(keyRowOrdering)
121+ KeyedPartitioning (expressions, partitionKeys)
122+ case p => p
123+ }
124+
125+ basePartitioning match {
112126 case k : KeyedPartitioning if spjParams.commonPartitionValues.isDefined =>
113127 // We allow duplicated partition values if
114128 // `spark.sql.sources.v2.bucketing.partiallyClusteredDistribution.enabled` is true
0 commit comments