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 @@ -17,7 +17,7 @@

package org.apache.spark.sql.catalyst.optimizer

import org.apache.spark.sql.catalyst.expressions.{And, EqualNullSafe, EqualTo, IsNull, Or, PredicateHelper}
import org.apache.spark.sql.catalyst.expressions.{And, EqualNullSafe, EqualTo, Expression, IsNull, Or, PredicateHelper}
import org.apache.spark.sql.catalyst.plans.logical.{Join, LogicalPlan}
import org.apache.spark.sql.catalyst.rules.Rule
import org.apache.spark.sql.catalyst.trees.TreePattern.{JOIN, OR}
Expand All @@ -30,7 +30,16 @@ object OptimizeJoinCondition extends Rule[LogicalPlan] with PredicateHelper {
override def apply(plan: LogicalPlan): LogicalPlan = plan.transformWithPruning(
_.containsPattern(JOIN), ruleId) {
case j @ Join(_, _, _, condition, _) if condition.nonEmpty =>
val newCondition = condition.map(_.transformWithPruning(_.containsPattern(OR), ruleId) {
val newCondition = condition.map(optimizeCondition)
j.copy(condition = newCondition)
}

// Rewriting the pattern to EqualNullSafe maps NULL to FALSE, so only recurse through And/Or.
private def optimizeCondition(condition: Expression): Expression = {
if (!condition.containsPattern(OR)) {
condition
} else {
condition match {
case Or(EqualTo(l, r), And(IsNull(c1), IsNull(c2)))
if (l.semanticEquals(c1) && r.semanticEquals(c2))
|| (l.semanticEquals(c2) && r.semanticEquals(c1)) =>
Expand All @@ -39,7 +48,18 @@ object OptimizeJoinCondition extends Rule[LogicalPlan] with PredicateHelper {
if (l.semanticEquals(c1) && r.semanticEquals(c2))
|| (l.semanticEquals(c2) && r.semanticEquals(c1)) =>
EqualNullSafe(l, r)
})
j.copy(condition = newCondition)
case and @ And(left, right) =>
val newLeft = optimizeCondition(left)
val newRight = optimizeCondition(right)
if (newLeft.fastEquals(left) && newRight.fastEquals(right)) and
else And(newLeft, newRight)
case or @ Or(left, right) =>
val newLeft = optimizeCondition(left)
val newRight = optimizeCondition(right)
if (newLeft.fastEquals(left) && newRight.fastEquals(right)) or
else Or(newLeft, newRight)
case other => other
}
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -46,4 +46,41 @@ class OptimizeJoinConditionSuite extends PlanTest {
comparePlans(Optimize.execute(originalQuery.analyze), correctAnswer.analyze)
})
}

test("SPARK-58384: do not replace null-safe equality pattern under NOT") {
val x = testRelation.subquery("x")
val y = testRelation1.subquery("y")
val originalQuery =
x.join(y, Inner, Option(!($"a" === $"c" || ($"a".isNull && $"c".isNull))))

comparePlans(Optimize.execute(originalQuery.analyze), originalQuery.analyze)
}

test("SPARK-58384: replace null-safe equality pattern under AND and OR") {
val x = testRelation.subquery("x")
val y = testRelation1.subquery("y")
val pattern = $"a" === $"c" || ($"a".isNull && $"c".isNull)
val optimizedPattern = $"a" <=> $"c"
val otherCondition = $"b" === $"d"
val conditions = Seq(
(pattern && otherCondition) -> (optimizedPattern && otherCondition),
(pattern || otherCondition) -> (optimizedPattern || otherCondition))

conditions.foreach { case (condition, optimizedCondition) =>
val originalQuery = x.join(y, Inner, Option(condition))
val correctAnswer = x.join(y, Inner, Option(optimizedCondition))
comparePlans(Optimize.execute(originalQuery.analyze), correctAnswer.analyze)
}
}

test("SPARK-58384: preserve unrelated AND and OR nodes") {
val x = testRelation.subquery("x")
val y = testRelation1.subquery("y")
val condition = ($"a" === $"c") && (($"b" === $"d") || ($"a" === 1))
val originalQuery = x.join(y, Inner, Option(condition)).analyze.asInstanceOf[Join]

val optimized = OptimizeJoinCondition(originalQuery).asInstanceOf[Join]

assert(optimized.condition.get eq originalQuery.condition.get)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -630,4 +630,13 @@ class DataFrameJoinSuite extends SharedSparkSession
checkAnswer(df1, df2)
})
}

test("SPARK-58384: preserve null semantics under NOT in join condition") {
val left = Seq((Some(0), 10), (None, 11)).toDF("a", "b").as("left")
val right = Seq((None, 20), (Some(0), 21), (Some(1), 22)).toDF("x", "y").as("right")
val condition = !(($"left.a" === $"right.x") ||
($"left.a".isNull && $"right.x".isNull))

checkAnswer(left.join(right, condition), Row(0, 10, 1, 22))
}
}