diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpression.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpression.scala index 5d85e89e1eab..424f04b2c8c2 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpression.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpression.scala @@ -186,6 +186,8 @@ object RewriteWithExpression extends Rule[LogicalPlan] { case With(child, defs) => // For With in the conditional branches, they may not be evaluated at all and we can't // pull the common expressions into a project which will always be evaluated. Inline it. + // SPARK-58902: Inlining nondeterministic expressions can cause multiple evaluations + // per row. Lazy per-row memoization is recommended for multi-referenced expressions. val refToExpr = defs.map(d => d.id -> d.child).toMap child.transformWithPruning(_.containsPattern(COMMON_EXPR_REF)) { case ref: CommonExpressionRef => refToExpr(ref.id) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpressionSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpressionSuite.scala index 8918b58ca1b5..4ff7404554cf 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpressionSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpressionSuite.scala @@ -500,4 +500,19 @@ class RewriteWithExpressionSuite extends PlanTest { val plan = testRelation.select(expr.as("col")) comparePlans(Optimizer.execute(plan), testRelation.select((a + a + 1).as("col"))) } + + test("SPARK-58902: inlines With expression inside conditional branches") { + val a = testRelation.output.head + val exprDef = CommonExpressionDef(a + a) + val exprRef = new CommonExpressionRef(exprDef) + // CaseWhen with With inside the ELSE branch + val withExpr = With(exprRef > 0 && exprRef < 10, Seq(exprDef)) + val caseWhenExpr = CaseWhen(Seq((a < 0, Literal(false))), Some(withExpr)) + val plan = testRelation.select(caseWhenExpr.as("col")).analyze + + val expectedPlan = testRelation.select( + CaseWhen(Seq((a < 0, Literal(false))), Some((a + a > 0) && (a + a < 10))).as("col") + ).analyze + comparePlans(Optimizer.execute(plan), expectedPlan) + } }