From 82a61b02040ea42295024878c94ec22905a2a5aa Mon Sep 17 00:00:00 2001 From: Nick Young Date: Wed, 16 Sep 2026 16:38:36 +0000 Subject: [PATCH 1/3] [SQL] Correct expression foldability and determinism --- .../catalyst/expressions/DynamicPruning.scala | 3 ++ .../spark/sql/catalyst/expressions/misc.scala | 29 +++++++++++++++- .../expressions/MiscExpressionsSuite.scala | 33 +++++++++++++++++++ 3 files changed, 64 insertions(+), 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/DynamicPruning.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/DynamicPruning.scala index 959acbc762b4e..6c7fd1f7c8a73 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/DynamicPruning.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/DynamicPruning.scala @@ -131,6 +131,9 @@ case class DynamicPruningExpression(child: Expression) extends UnaryExpression with DynamicPruning { override def eval(input: InternalRow): Any = child.eval(input) + + override def foldable: Boolean = false + final override val nodePatterns: Seq[TreePattern] = Seq(DYNAMIC_PRUNING_EXPRESSION) override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala index 61e04b9f761ce..e677a5db0d5c0 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala @@ -38,6 +38,10 @@ case class PrintToStderr(child: Expression) extends UnaryExpression { override def dataType: DataType = child.dataType + override lazy val deterministic: Boolean = false + + override def foldable: Boolean = false + protected override def nullSafeEval(input: Any): Any = { // scalastyle:off println System.err.println(outputPrefix + input) @@ -412,6 +416,26 @@ case class CurrentUser() final override val nodePatterns: Seq[TreePattern] = Seq(CURRENT_LIKE) } +private object AesEncryptDeterminism { + def apply(arguments: Seq[Expression]): Boolean = { + arguments.forall(_.deterministic) && + (hasNonEmptyLiteral(arguments(4)) || isEcbLiteral(arguments(2))) + } + + private def hasNonEmptyLiteral(expression: Expression): Boolean = expression match { + case Literal(value: Array[Byte], _) => value.nonEmpty + case Literal(value: UTF8String, _) => value.numBytes() > 0 + case cast: Cast if cast.dataType == BinaryType => hasNonEmptyLiteral(cast.child) + case _ => false + } + + private def isEcbLiteral(expression: Expression): Boolean = expression match { + case Literal(value: UTF8String, _) => value.toString.equalsIgnoreCase("ECB") + case cast: Cast if cast.dataType.isInstanceOf[StringType] => isEcbLiteral(cast.child) + case _ => false + } +} + /** * A function that encrypts input using AES. Key lengths of 128, 192 or 256 bits can be used. * If either argument is NULL or the key length is not one of the permitted values, @@ -471,12 +495,15 @@ case class AesEncrypt( aad: Expression) extends RuntimeReplaceable with ImplicitCastInputTypes { + override lazy val deterministic: Boolean = AesEncryptDeterminism(children) + override lazy val replacement: Expression = StaticInvoke( classOf[ExpressionImplUtils], BinaryType, "aesEncrypt", Seq(input, key, mode, padding, iv, aad), - inputTypes) + inputTypes, + isDeterministic = deterministic) def this(input: Expression, key: Expression, mode: Expression, padding: Expression, iv: Expression) = this(input, key, mode, padding, iv, Literal("")) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala index a586e33afdd26..663651bb029ef 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala @@ -96,6 +96,9 @@ class MiscExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { test("PrintToStderr") { val inputExpr = Literal(1) + assert(!PrintToStderr(inputExpr).foldable) + assert(!PrintToStderr(inputExpr).deterministic) + val systemErr = System.err val (outputEval, outputCodegen) = try { @@ -117,6 +120,36 @@ class MiscExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { assert(outputEval.contains(s"Result of $inputExpr is 1")) } + test("DynamicPruningExpression is not foldable") { + assert(!DynamicPruningExpression(Literal.TrueLiteral).foldable) + } + + test("AesEncrypt determinism reflects whether it generates an IV") { + val randomIvExpression = new AesEncrypt(Literal("Spark"), Literal("0000111122223333")) + assert(!randomIvExpression.deterministic) + assert(!randomIvExpression.replacement.deterministic) + assert(!randomIvExpression.replacement.foldable) + + val explicitIvExpression = new AesEncrypt( + Literal("Spark"), + Literal("0000111122223333"), + Literal("GCM"), + Literal("DEFAULT"), + Literal(Array.fill[Byte](12)(0))) + assert(explicitIvExpression.deterministic) + assert(explicitIvExpression.replacement.deterministic) + assert(explicitIvExpression.replacement.foldable) + + val ecbExpression = new AesEncrypt( + Literal("Spark"), + Literal("0000111122223333"), + Literal("ECB"), + Literal("PKCS")) + assert(ecbExpression.deterministic) + assert(ecbExpression.replacement.deterministic) + assert(ecbExpression.replacement.foldable) + } + test("Hmac") { def bytes(hex: String): Array[Byte] = hex.grouped(2).map(Integer.parseInt(_, 16).toByte).toArray From 768a33bc4a24aa6d882ff9a19c038feef622a6f6 Mon Sep 17 00:00:00 2001 From: Nick Young Date: Fri, 18 Sep 2026 18:40:16 +0000 Subject: [PATCH 2/3] Address AES nondeterminism review feedback --- ...NondeterministicExpressionCollection.scala | 12 ++++- .../spark/sql/catalyst/expressions/misc.scala | 32 ++++++++---- .../PullOutNondeterministicSuite.scala | 52 ++++++++++++++++++- .../expressions/MiscExpressionsSuite.scala | 11 +++- .../RewriteWithExpressionSuite.scala | 14 ----- 5 files changed, 95 insertions(+), 26 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala index d8a29b984859c..4338f6b52acb8 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala @@ -33,8 +33,18 @@ object NondeterministicExpressionCollection { case udf: UserDefinedExpression if !udf.deterministic => udf case udf: ExternalUserDefinedFunction if !udf.deterministic => udf } + val expressionsToCollect = if (leafNondeterministic.nonEmpty) { + leafNondeterministic + } else { + expr.collect { + case nondeterministicExpr + if !nondeterministicExpr.deterministic && + nondeterministicExpr.children.forall(_.deterministic) => + nondeterministicExpr + } + } - for (nondeterministicExpr <- leafNondeterministic.distinct) { + for (nondeterministicExpr <- expressionsToCollect.distinct) { val namedExpression = nondeterministicExpr match { case namedExpression: NamedExpression => namedExpression case _ => Alias(nondeterministicExpr, "_nondeterministic")() diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala index e677a5db0d5c0..423063a19250b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/misc.scala @@ -17,6 +17,8 @@ package org.apache.spark.sql.catalyst.expressions +import scala.util.control.NonFatal + import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.{ExpressionBuilder, FunctionRegistry, UnresolvedSeed} import org.apache.spark.sql.catalyst.expressions.codegen._ @@ -419,19 +421,31 @@ case class CurrentUser() private object AesEncryptDeterminism { def apply(arguments: Seq[Expression]): Boolean = { arguments.forall(_.deterministic) && - (hasNonEmptyLiteral(arguments(4)) || isEcbLiteral(arguments(2))) + (hasNonEmptyFixedValue(arguments(4)) || isEcbFixedValue(arguments(2))) } - private def hasNonEmptyLiteral(expression: Expression): Boolean = expression match { - case Literal(value: Array[Byte], _) => value.nonEmpty - case Literal(value: UTF8String, _) => value.numBytes() > 0 - case cast: Cast if cast.dataType == BinaryType => hasNonEmptyLiteral(cast.child) - case _ => false + private def fixedValue(expression: Expression): Option[Any] = { + if (expression.resolved && expression.foldable && expression.deterministic && + expression.contextIndependentFoldable) { + try { + Option(expression.eval(EmptyRow)) + } catch { + case NonFatal(_) => None + } + } else { + None + } } - private def isEcbLiteral(expression: Expression): Boolean = expression match { - case Literal(value: UTF8String, _) => value.toString.equalsIgnoreCase("ECB") - case cast: Cast if cast.dataType.isInstanceOf[StringType] => isEcbLiteral(cast.child) + private def hasNonEmptyFixedValue(expression: Expression): Boolean = + fixedValue(expression) match { + case Some(value: Array[Byte]) => value.nonEmpty + case Some(value: UTF8String) => value.numBytes() > 0 + case _ => false + } + + private def isEcbFixedValue(expression: Expression): Boolean = fixedValue(expression) match { + case Some(value: UTF8String) => value.toString.equalsIgnoreCase("ECB") case _ => false } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala index 5183c90f57221..1d338a7d6fe0a 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala @@ -17,10 +17,12 @@ package org.apache.spark.sql.catalyst.analysis +import org.apache.spark.sql.catalyst.QueryPlanningTracker import org.apache.spark.sql.catalyst.dsl.expressions._ import org.apache.spark.sql.catalyst.dsl.plans._ import org.apache.spark.sql.catalyst.expressions._ -import org.apache.spark.sql.catalyst.plans.logical.LocalRelation +import org.apache.spark.sql.catalyst.parser.CatalystSqlParser.parsePlan +import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, LocalRelation, LogicalPlan, Sort} /** * Test suite for moving non-deterministic expressions into Project. @@ -33,6 +35,10 @@ class PullOutNondeterministicSuite extends AnalysisTest { private lazy val rnd = Rand(10).as("_nondeterministic") private lazy val rndref = rnd.toAttribute + private def analyze(sqlText: String): LogicalPlan = { + getAnalyzer.executeAndCheck(parsePlan(sqlText), new QueryPlanningTracker) + } + test("no-op on filter") { checkAnalysis( r.where(Rand(10) > Literal(1.0)), @@ -53,4 +59,48 @@ class PullOutNondeterministicSuite extends AnalysisTest { r.select(a, b, rnd).groupBy(rndref)(rndref.as("rnd")) ) } + + test("aes_encrypt with a foldable fixed IV is deterministic") { + val analyzed = analyze( + """SELECT aes_encrypt( + | 'Spark', + | '0000111122223333', + | 'GCM', + | 'DEFAULT', + | unhex('000000000000000000000000')) + |""".stripMargin) + val aesEncrypt = analyzed.expressions.flatMap(_.collect { + case aesEncrypt: AesEncrypt => aesEncrypt + }).head + + assert(aesEncrypt.deterministic) + assert(aesEncrypt.replacement.deterministic) + assert(aesEncrypt.replacement.foldable) + } + + test("pull out aes_encrypt that generates a random IV") { + val queries = Seq( + """SELECT aes_encrypt('Spark', '0000111122223333') AS encrypted + |FROM TaBlE + |GROUP BY aes_encrypt('Spark', '0000111122223333') + |""".stripMargin, + """SELECT * FROM TaBlE + |ORDER BY aes_encrypt('Spark', '0000111122223333') + |""".stripMargin) + + queries.foreach { sqlText => + val analyzed = analyze(sqlText) + + val nondeterministicOperatorExpressions = analyzed.collect { + case aggregate: Aggregate => aggregate.groupingExpressions + case sort: Sort => sort.order + }.flatten.filterNot(_.deterministic) + assert(nondeterministicOperatorExpressions.isEmpty) + + val aesEncryptCount = analyzed.collect { + case plan => plan.expressions.flatMap(_.collect { case _: AesEncrypt => 1 }).sum + }.sum + assert(aesEncryptCount == 1) + } + } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala index 663651bb029ef..15c1523aae4b0 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/MiscExpressionsSuite.scala @@ -135,7 +135,7 @@ class MiscExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { Literal("0000111122223333"), Literal("GCM"), Literal("DEFAULT"), - Literal(Array.fill[Byte](12)(0))) + Unhex(Literal("000000000000000000000000"))) assert(explicitIvExpression.deterministic) assert(explicitIvExpression.replacement.deterministic) assert(explicitIvExpression.replacement.foldable) @@ -148,6 +148,15 @@ class MiscExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { assert(ecbExpression.deterministic) assert(ecbExpression.replacement.deterministic) assert(ecbExpression.replacement.foldable) + + val nondeterministicKeyExpression = new AesEncrypt( + Literal("Spark"), + Uuid(Some(0)), + Literal("ECB"), + Literal("PKCS")) + assert(!nondeterministicKeyExpression.deterministic) + assert(!nondeterministicKeyExpression.replacement.deterministic) + assert(!nondeterministicKeyExpression.replacement.foldable) } test("Hmac") { 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 11bd8623033c9..a7d0154f08c81 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 @@ -127,20 +127,6 @@ class RewriteWithExpressionSuite extends PlanTest { assert(rewritten == (Literal(1) + Literal(1)) * (Literal(1) + Literal(1))) } - test("applyForExpression rejects an impure foldable definition referenced more than once") { - // aes_encrypt becomes a foldable StaticInvoke that draws a fresh random IV on every eval, so - // two inlined copies would encrypt to different values. - val aes = ReplaceExpressions.replace( - new AesEncrypt(Literal("abc".getBytes), Literal("1234567890123456".getBytes))) - assert(aes.foldable, "the AES rewrite is only interesting while it stays foldable") - val expr = With(aes) { case Seq(ref) => - EqualTo(ref, ref) - } - intercept[SparkException] { - RewriteWithExpression.applyForExpression(expr) - } - } - test("applyForExpression rejects canonicalized common expression ids") { // Canonicalization re-numbers ids per `With`, starting from 1, so these two siblings both get // id 1: the safe (literal) definition would otherwise mark id 1 safe in the flat `safeIds` set From 144e984b890d93c3ae066bdefe10954974646709 Mon Sep 17 00:00:00 2001 From: Nick Young Date: Wed, 23 Sep 2026 18:51:02 +0000 Subject: [PATCH 3/3] Address mixed nondeterministic extraction --- ...NondeterministicExpressionCollection.scala | 43 ++++++++++++------- .../PullOutNondeterministicSuite.scala | 30 +++++++++++++ .../RewriteWithExpressionSuite.scala | 21 ++++++++- 3 files changed, 78 insertions(+), 16 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala index 4338f6b52acb8..ba53d0b02e597 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/NondeterministicExpressionCollection.scala @@ -22,27 +22,40 @@ import java.util.LinkedHashMap import org.apache.spark.sql.catalyst.expressions._ object NondeterministicExpressionCollection { + /** + * Returns a copy used only to classify the parent, plus the minimal nondeterministic + * evaluation frontier below `expression`. + * + * A selected expression is represented by a deterministic attribute while walking upward. If + * its parent is still nondeterministic after that substitution, the parent is the evaluation + * unit and replaces the selected descendants in the frontier. + */ + private def collectNondeterministicFrontier( + expression: Expression): (Expression, Seq[Expression]) = { + if (expression.deterministic) { + (expression, Seq.empty) + } else { + val childResults = expression.children.map(collectNondeterministicFrontier) + val expressionWithDeterministicChildren = + expression.withNewChildren(childResults.map(_._1)) + + if (!expressionWithDeterministicChildren.deterministic) { + val placeholder = AttributeReference( + "_nondeterministic", expression.dataType, expression.nullable)() + (placeholder, Seq(expression)) + } else { + (expressionWithDeterministicChildren, childResults.flatMap(_._2)) + } + } + } + def getNondeterministicToAttributes( expressions: Seq[Expression]): LinkedHashMap[Expression, NamedExpression] = { val nonDeterministicToAttributes = new LinkedHashMap[Expression, NamedExpression] for (expr <- expressions) { if (!expr.deterministic) { - val leafNondeterministic = expr.collect { - case nondeterministicExpr: Nondeterministic => nondeterministicExpr - case udf: UserDefinedExpression if !udf.deterministic => udf - case udf: ExternalUserDefinedFunction if !udf.deterministic => udf - } - val expressionsToCollect = if (leafNondeterministic.nonEmpty) { - leafNondeterministic - } else { - expr.collect { - case nondeterministicExpr - if !nondeterministicExpr.deterministic && - nondeterministicExpr.children.forall(_.deterministic) => - nondeterministicExpr - } - } + val expressionsToCollect = collectNondeterministicFrontier(expr)._2 for (nondeterministicExpr <- expressionsToCollect.distinct) { val namedExpression = nondeterministicExpr match { diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala index 1d338a7d6fe0a..d42b8e568af9e 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/PullOutNondeterministicSuite.scala @@ -103,4 +103,34 @@ class PullOutNondeterministicSuite extends AnalysisTest { assert(aesEncryptCount == 1) } } + + test("pull out random-IV aes_encrypt with a nondeterministic child") { + val queries = Seq( + """SELECT count(*) + |FROM TaBlE + |GROUP BY aes_encrypt(uuid(), '0000111122223333') + |""".stripMargin, + """SELECT * FROM TaBlE + |ORDER BY aes_encrypt(uuid(), '0000111122223333') + |""".stripMargin) + + queries.foreach { sqlText => + val analyzed = analyze(sqlText) + + val nondeterministicOperatorExpressions = analyzed.collect { + case aggregate: Aggregate => aggregate.groupingExpressions + case sort: Sort => sort.order + }.flatten.filterNot(_.deterministic) + assert(nondeterministicOperatorExpressions.isEmpty) + + val extractedExpressions = analyzed.collect { + case plan => plan.expressions.flatMap(_.collect { + case _: AesEncrypt => "aes_encrypt" + case _: Uuid => "uuid" + }) + }.flatten + assert(extractedExpressions.count(_ == "aes_encrypt") == 1) + assert(extractedExpressions.count(_ == "uuid") == 1) + } + } } 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 a7d0154f08c81..c513572c75493 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 @@ -27,6 +27,7 @@ import org.apache.spark.sql.catalyst.catalog.{InMemoryCatalog, SessionCatalog} import org.apache.spark.sql.catalyst.dsl.expressions._ import org.apache.spark.sql.catalyst.dsl.plans._ import org.apache.spark.sql.catalyst.expressions._ +import org.apache.spark.sql.catalyst.expressions.codegen.CodegenFallback import org.apache.spark.sql.catalyst.plans.PlanTest import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, LogicalPlan, Project} import org.apache.spark.sql.catalyst.rules.RuleExecutor @@ -34,10 +35,19 @@ import org.apache.spark.sql.catalyst.util.DateTimeUtils import org.apache.spark.sql.connector.catalog.CatalogV2Implicits._ import org.apache.spark.sql.connector.catalog.DefaultCatalogManager import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{DateType, IntegerType, StringType, TimestampNTZType, TimestampType, TimeType} +import org.apache.spark.sql.types.{DataType, DateType, IntegerType, StringType, TimestampNTZType, TimestampType, TimeType} class RewriteWithExpressionSuite extends PlanTest { + /** A foldable expression whose evaluation must not be duplicated. */ + private case class ImpureFoldable(child: Expression) + extends UnaryExpression with NonSQLExpression with CodegenFallback { + override def dataType: DataType = child.dataType + override def eval(input: InternalRow): Any = child.eval(input) + override protected def withNewChildInternal(newChild: Expression): ImpureFoldable = + copy(child = newChild) + } + object Optimizer extends RuleExecutor[LogicalPlan] { val batches = Batch("Rewrite With expression", FixedPoint(5), PullOutGroupingExpressions, @@ -127,6 +137,15 @@ class RewriteWithExpressionSuite extends PlanTest { assert(rewritten == (Literal(1) + Literal(1)) * (Literal(1) + Literal(1))) } + test("applyForExpression rejects an impure foldable definition referenced more than once") { + val expr = With(ImpureFoldable(Literal(1))) { case Seq(ref) => + ref + ref + } + intercept[SparkException] { + RewriteWithExpression.applyForExpression(expr) + } + } + test("applyForExpression rejects canonicalized common expression ids") { // Canonicalization re-numbers ids per `With`, starting from 1, so these two siblings both get // id 1: the safe (literal) definition would otherwise mark id 1 safe in the flat `safeIds` set