diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala index e7762565f47ec..1fb167dda6e4e 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala @@ -429,8 +429,15 @@ abstract class InMemoryBaseTable( private var _pushedFilters: Array[Filter] = Array.empty override def build: Scan = { - val scan = InMemoryBatchScan( - data.map(_.asInstanceOf[InputPartition]).toImmutableArraySeq, schema, tableSchema, options) + val scan = if (InMemoryBaseTable.this.ordering.nonEmpty) { + new InMemoryBatchScanWithOrdering( + data.map(_.asInstanceOf[InputPartition]).toImmutableArraySeq, schema, tableSchema, + options) + } else { + InMemoryBatchScan( + data.map(_.asInstanceOf[InputPartition]).toImmutableArraySeq, schema, tableSchema, + options) + } if (evaluableFilters.nonEmpty) { scan.filter(evaluableFilters) } @@ -596,6 +603,16 @@ abstract class InMemoryBaseTable( } } + private class InMemoryBatchScanWithOrdering( + data: Seq[InputPartition], + readSchema: StructType, + tableSchema: StructType, + options: CaseInsensitiveStringMap) + extends InMemoryBatchScan(data, readSchema, tableSchema, options) + with SupportsReportOrdering { + override def outputOrdering(): Array[SortOrder] = InMemoryBaseTable.this.ordering + } + abstract class InMemoryWriterBuilder(val info: LogicalWriteInfo) extends SupportsTruncate with SupportsDynamicOverwrite with SupportsStreamingUpdateAsAppend { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanPartitioningAndOrdering.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanPartitioningAndOrdering.scala index 5d06c8786d894..7f51875b971fe 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanPartitioningAndOrdering.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2ScanPartitioningAndOrdering.scala @@ -69,7 +69,8 @@ object V2ScanPartitioningAndOrdering extends Rule[LogicalPlan] with Logging { private def ordering(plan: LogicalPlan) = plan.transformDown { case d @ DataSourceV2ScanRelation(relation, scan: SupportsReportOrdering, _, _, _) => - val ordering = V2ExpressionUtils.toCatalystOrdering(scan.outputOrdering(), relation) + val ordering = + V2ExpressionUtils.toCatalystOrdering(scan.outputOrdering(), relation, relation.funCatalog) d.copy(ordering = Some(ordering)) } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/WriteDistributionAndOrderingSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/WriteDistributionAndOrderingSuite.scala index 588490e07dfd6..ec1b34b8c2101 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/WriteDistributionAndOrderingSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/WriteDistributionAndOrderingSuite.scala @@ -22,7 +22,7 @@ import java.sql.Date import java.util.Collections import org.apache.spark.sql.{catalyst, AnalysisException, DataFrame, Row} -import org.apache.spark.sql.catalyst.expressions.{ApplyFunctionExpression, Cast, Literal} +import org.apache.spark.sql.catalyst.expressions.{ApplyFunctionExpression, Cast, Literal, TransformExpression} import org.apache.spark.sql.catalyst.expressions.objects.Invoke import org.apache.spark.sql.catalyst.plans.physical import org.apache.spark.sql.catalyst.plans.physical.{CoalescedBoundary, CoalescedHashPartitioning, HashPartitioning, RangePartitioning, UnknownPartitioning} @@ -30,9 +30,10 @@ import org.apache.spark.sql.connector.catalog.{Column, Identifier} import org.apache.spark.sql.connector.catalog.functions._ import org.apache.spark.sql.connector.distributions.{Distribution, Distributions} import org.apache.spark.sql.connector.expressions._ -import org.apache.spark.sql.connector.expressions.LogicalExpressions._ +import org.apache.spark.sql.connector.expressions.Expressions._ import org.apache.spark.sql.execution.{QueryExecution, SortExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.AQEShuffleReadExec +import org.apache.spark.sql.execution.datasources.v2.BatchScanExec import org.apache.spark.sql.execution.datasources.v2.V2TableWriteExec import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike import org.apache.spark.sql.execution.streaming.runtime.MemoryStream @@ -1532,4 +1533,22 @@ class WriteDistributionAndOrderingSuite extends DistributionAndOrderingSuiteBase Seq(None) } } + + test("SPARK-56321: Scan with SupportsReportOrdering and function-based sort order") { + val bucketById = bucket(4, "id") + val tableOrdering = Array(sort(bucketById, SortDirection.ASCENDING, NullOrdering.NULLS_FIRST)) + catalog.createTable(ident, columns, Array(bucketById), emptyProps, + Distributions.unspecified(), tableOrdering, None, None) + + sql(s"INSERT INTO testcat.ns1.test_table VALUES (1, 'a', date '2021-01-01')") + + val df = sql("SELECT id, data FROM testcat.ns1.test_table") + val scans = collect(df.queryExecution.executedPlan) { case s: BatchScanExec => s } + assert(scans.size === 1) + val ordering = scans.head.outputOrdering + assert(ordering.nonEmpty, + "scan should report non-empty outputOrdering via SupportsReportOrdering") + assert(ordering.head.child.isInstanceOf[TransformExpression], + "bucket-based sort order should resolve to a TransformExpression") + } }