From a0fd72390a11a5055804689b068119917777a609 Mon Sep 17 00:00:00 2001 From: marcus Date: Sun, 9 Aug 2026 12:20:26 -0700 Subject: [PATCH] [SPARK-58208][SQL] Deep-copy stateful expressions before optimization --- .../sql/catalyst/optimizer/Optimizer.scala | 9 ++++--- .../ConvertToLocalRelationSuite.scala | 18 ++++++++++--- .../spark/sql/execution/QueryExecution.scala | 10 ++++++-- .../sql/execution/QueryExecutionSuite.scala | 25 ++++++++++++++++++- 4 files changed, 53 insertions(+), 9 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala index 802e86374f508..c5e90a8ba8520 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala @@ -2785,9 +2785,11 @@ object ConvertToLocalRelation extends Rule[LogicalPlan] { _.containsPattern(LOCAL_RELATION), ruleId) { case Project(projectList, LocalRelation(output, data, isStreaming, stream)) if !projectList.exists(hasUnevaluableExpr) => - val projection = new InterpretedMutableProjection(projectList, output) + val freshProjectList = projectList.map( + _.freshCopyIfContainsStatefulExpression().asInstanceOf[NamedExpression]) + val projection = new InterpretedMutableProjection(freshProjectList, output) projection.initialize(0) - LocalRelation(projectList.map(_.toAttribute), data.map(projection(_).copy()), + LocalRelation(freshProjectList.map(_.toAttribute), data.map(projection(_).copy()), isStreaming, stream) case Limit(IntegerLiteral(limit), LocalRelation(output, data, isStreaming, stream)) => @@ -2798,7 +2800,8 @@ object ConvertToLocalRelation extends Rule[LogicalPlan] { case Filter(condition, LocalRelation(output, data, isStreaming, stream)) if !hasUnevaluableExpr(condition) => - val predicate = Predicate.create(condition, output) + val freshCondition = condition.freshCopyIfContainsStatefulExpression() + val predicate = Predicate.create(freshCondition, output) predicate.initialize(0) LocalRelation(output, data.filter(row => predicate.eval(row)), isStreaming, stream) } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/ConvertToLocalRelationSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/ConvertToLocalRelationSuite.scala index 622af60d85d93..f4d412153b0ed 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/ConvertToLocalRelationSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/ConvertToLocalRelationSuite.scala @@ -21,12 +21,12 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.UnresolvedAttribute import org.apache.spark.sql.catalyst.dsl.expressions._ import org.apache.spark.sql.catalyst.dsl.plans._ -import org.apache.spark.sql.catalyst.expressions.{Expression, GenericInternalRow, LessThan, Literal, UnaryExpression} +import org.apache.spark.sql.catalyst.expressions.{Add, Alias, ArrayTransform, Expression, GenericInternalRow, LambdaFunction, LessThan, Literal, NamedLambdaVariable, UnaryExpression} import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} import org.apache.spark.sql.catalyst.plans.PlanTest -import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, LogicalPlan} +import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, LogicalPlan, Project} import org.apache.spark.sql.catalyst.rules.RuleExecutor -import org.apache.spark.sql.types.{DataType, StructType} +import org.apache.spark.sql.types.{ArrayType, DataType, IntegerType, StructType} class ConvertToLocalRelationSuite extends PlanTest { @@ -87,6 +87,18 @@ class ConvertToLocalRelationSuite extends PlanTest { comparePlans(optimized, correctAnswer) } + + test("SPARK-58208: ConvertToLocalRelation uses fresh stateful project expressions") { + val element = NamedLambdaVariable("x", IntegerType, nullable = false) + val transform = ArrayTransform( + Literal.create(Seq(1, 2), ArrayType(IntegerType, containsNull = false)), + LambdaFunction(Add(element, Literal(1)), Seq(element))) + val project = Project(Seq(Alias(transform, "v")()), LocalRelation(Nil, Seq(InternalRow.empty))) + + Optimize.execute(project) + + assert(element.value.get() == null) + } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala index 1fa70351f2d00..8bd12ff724af4 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala @@ -308,6 +308,12 @@ class QueryExecution( def assertCommandExecuted(): Unit = commandExecuted + private def cloneWithFreshStatefulExpressions(plan: LogicalPlan): LogicalPlan = { + plan.clone().transformWithSubqueries { + case node => node.mapExpressions(_.freshCopyIfContainsStatefulExpression()) + } + } + private val lazyOptimizedPlan = LazyTry { // We need to materialize the commandExecuted here because optimizedPlan is also tracked under // the optimizing phase @@ -315,8 +321,8 @@ class QueryExecution( executePhase(QueryPlanningTracker.OPTIMIZATION) { // clone the plan to avoid sharing the plan instance between different stages like analyzing, // optimizing and planning. - val plan = - sparkSession.sessionState.optimizer.executeAndTrack(withCachedData.clone(), tracker) + val plan = sparkSession.sessionState.optimizer.executeAndTrack( + cloneWithFreshStatefulExpressions(withCachedData), tracker) // We do not want optimized plans to be re-analyzed as literals that have been constant // folded and such can cause issues during analysis. While `clone` should maintain the // `analyzed` state of the LogicalPlan, we set the plan as analyzed here as well out of diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/QueryExecutionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/QueryExecutionSuite.scala index f7afdb5e6e537..078055f2d8b8d 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/QueryExecutionSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/QueryExecutionSuite.scala @@ -25,7 +25,7 @@ import org.apache.spark.scheduler.{SparkListener, SparkListenerEvent, SparkListe import org.apache.spark.sql.{AnalysisException, ExtendedExplainGenerator, FastOperator, SaveMode} import org.apache.spark.sql.catalyst.{QueryPlanningTracker, QueryPlanningTrackerCallback, TableIdentifier} import org.apache.spark.sql.catalyst.analysis.{CurrentNamespace, UnresolvedFunction, UnresolvedRelation} -import org.apache.spark.sql.catalyst.expressions.{Alias, UnsafeRow} +import org.apache.spark.sql.catalyst.expressions.{Alias, NamedLambdaVariable, UnsafeRow} import org.apache.spark.sql.catalyst.plans.QueryPlan import org.apache.spark.sql.catalyst.plans.logical.{CommandResult, LogicalPlan, OneRowRelation, Project, ShowTables, SubqueryAlias} import org.apache.spark.sql.catalyst.trees.TreeNodeTag @@ -55,6 +55,14 @@ class QueryExecutionSuite extends SharedSparkSession { override protected def sparkConf = super.sparkConf.set(SQLConf.ADAPTIVE_MAX_SHUFFLE_HASH_JOIN_LOCAL_MAP_THRESHOLD.key, "0") + private def collectLambdaVariables(plan: LogicalPlan): Seq[NamedLambdaVariable] = { + plan.collect { + case node => node.expressions.flatMap(_.collect { + case variable: NamedLambdaVariable => variable + }) + }.flatten + } + def checkDumpedPlans(path: String, expected: Int): Unit = Utils.tryWithResource( Source.fromFile(path)) { source => assert(source.getLines().toList @@ -105,6 +113,21 @@ class QueryExecutionSuite extends SharedSparkSession { } } + test("SPARK-58208: optimizedPlan uses fresh stateful expressions") { + val df = spark.range(1).selectExpr("transform(array(id), x -> x + 1) AS v") + val queryExecution = df.queryExecution + + val beforeOptimize = collectLambdaVariables(queryExecution.withCachedData) + val optimized = collectLambdaVariables(queryExecution.optimizedPlan) + + assert(beforeOptimize.nonEmpty) + assert(beforeOptimize.size == optimized.size) + beforeOptimize.zip(optimized).foreach { case (before, after) => + assert(before.exprId == after.exprId) + assert(before.value ne after.value) + } + } + test("dumping query execution info by invalid path") { val path = "1234567890://plans.txt" val exception = intercept[IllegalArgumentException] {