Skip to content
Closed
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 @@ -24,11 +24,11 @@ import org.apache.spark.sql.catalyst.TableIdentifier
import org.apache.spark.sql.catalyst.analysis.UnresolvedRelation
import org.apache.spark.sql.catalyst.dsl.expressions._
import org.apache.spark.sql.catalyst.dsl.plans._
import org.apache.spark.sql.catalyst.expressions.{Alias, AttributeReference, Expression, ListQuery, Literal, NamedExpression, Rand}
import org.apache.spark.sql.catalyst.expressions.{Alias, AttributeReference, AttributeSeq, Expression, ExprId, ListQuery, Literal, NamedExpression, Rand}
import org.apache.spark.sql.catalyst.plans.logical.{Filter, LocalRelation, LogicalPlan, Project, Union}
import org.apache.spark.sql.catalyst.rules.Rule
import org.apache.spark.sql.catalyst.trees.{CurrentOrigin, Origin, TreePattern}
import org.apache.spark.sql.types.IntegerType
import org.apache.spark.sql.types.{IntegerType, StringType}

class QueryPlanSuite extends SparkFunSuite {

Expand Down Expand Up @@ -185,4 +185,19 @@ class QueryPlanSuite extends SparkFunSuite {
assert(visited.size == 2)
assert(visited.forall(_.containsPattern(TreePattern.FILTER)))
}

test("SPARK-58120: normalizePredicates reorders expressions by hashCode via orderCommutative") {
val id = AttributeReference("id", IntegerType)(ExprId(0))
val data = AttributeReference("data", StringType)(ExprId(1))
val output: AttributeSeq = Seq(id, data)
val exprs: Seq[Expression] = Seq(id, data)

val normalizedPredicates = QueryPlan.normalizePredicates(exprs, output)
assert(normalizedPredicates.map(_.dataType) == Seq(StringType, IntegerType),
"normalizePredicates reorders expressions by hashCode")

val normalizedExpressions = exprs.map(QueryPlan.normalizeExpressions(_, output))
assert(normalizedExpressions.map(_.dataType) == Seq(IntegerType, StringType),
"normalizeExpressions preserves original expression order")
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,8 @@ case class BatchScanExec(
runtimeFilters = QueryPlan.normalizePredicates(
runtimeFilters.filterNot(_ == DynamicPruningExpression(Literal.TrueLiteral)),
output),
keyGroupedPartitioning = keyGroupedPartitioning.map(QueryPlan.normalizePredicates(_, output)))
keyGroupedPartitioning = keyGroupedPartitioning.map(_.map(
QueryPlan.normalizeExpressions(_, output))))
}

override def simpleString(maxFields: Int): String = {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ import org.apache.spark.sql.catalyst.trees.LeafLike
import org.apache.spark.sql.execution.datasources.v2.BatchScanExec
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.test.SharedSparkSession
import org.apache.spark.sql.types.{DataType, IntegerType}
import org.apache.spark.sql.types.{DataType, IntegerType, StringType}
import org.apache.spark.sql.vectorized.ColumnarBatch
import org.apache.spark.util.ThreadUtils

Expand Down Expand Up @@ -232,6 +232,30 @@ class SparkPlanSuite extends SharedSparkSession {
executor.shutdown()
}
}

test("SPARK-58120: BatchScanExec doCanonicalize preserves keyGroupedPartitioning order") {
// The int/string pair is known to reorder under normalizePredicates hashCode sorting.
// See QueryPlanSuite "SPARK-58120: normalizePredicates reorders expressions by hashCode
// via orderCommutative".
val id = AttributeReference("id", IntegerType)(ExprId(0))
val data = AttributeReference("data", StringType)(ExprId(1))
val scan = BatchScanExec(
output = Seq(id, data),
scan = null,
runtimeFilters = Seq.empty,
table = null,
keyGroupedPartitioning = Some(Seq(id, data)))
val canonicalized = scan.canonicalized.asInstanceOf[BatchScanExec]
assert(scan.keyGroupedPartitioning.isDefined,
"Expected BatchScanExec to have keyGroupedPartitioning set")
val originalTypes = scan.keyGroupedPartitioning.get.map(_.dataType)
val canonicalizedTypes = canonicalized.keyGroupedPartitioning.get.map(_.dataType)
assert(originalTypes == canonicalizedTypes,
s"Expression order changed after canonicalization: " +
s"original types=$originalTypes, canonicalized types=$canonicalizedTypes")
assert(canonicalizedTypes == Seq(IntegerType, StringType),
s"Expected [IntegerType, StringType] but got $canonicalizedTypes")
}
}

case class ColumnarOp(child: SparkPlan) extends UnaryExecNode {
Expand Down