diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 04dcd223275ad..5e5f9dcaca564 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -49,7 +49,7 @@ import org.apache.spark.sql.catalyst.trees.AlwaysProcess import org.apache.spark.sql.catalyst.trees.CurrentOrigin.withOrigin import org.apache.spark.sql.catalyst.trees.TreePattern._ import org.apache.spark.sql.catalyst.types.DataTypeUtils -import org.apache.spark.sql.catalyst.util.{toPrettySQL, trimTempResolvedColumn, CharVarcharUtils, GeneratedColumn} +import org.apache.spark.sql.catalyst.util.{toPrettySQL, trimTempResolvedColumn, CharVarcharUtils, GeneratedColumn, MetadataColumnHelper} import org.apache.spark.sql.catalyst.util.ResolveDefaultColumns._ // `View` is aliased to `V2View` to avoid clashing with the logical-plan `View` imported via // `org.apache.spark.sql.catalyst.plans.logical._`. @@ -700,6 +700,7 @@ class Analyzer( UpdateOuterReferences), Batch("Cleanup", fixedPoint, CleanupAliases), + Batch("Eliminate Resolved Pipe SET Inputs", Once, EliminateResolvedPipeSetInputs), Batch("HandleSpecialCommand", Once, HandleSpecialCommand), Batch("Remove watermark for batch query", Once, @@ -2152,7 +2153,9 @@ class Analyzer( resolvesToCountBuiltin && f.arguments.length == 1) { f.arguments.foreach { - case u: UnresolvedStar if u.isQualifiedByTable(child.output, resolver) => + case u: UnresolvedStar if u.isQualifiedByTable( + child.output ++ child.metadataOutput.filter(_.qualifiedAccessOnly), + resolver) => throw QueryCompilationErrors .singleTableStarInCountNotAllowedError(u.target.get.mkString(".")) case _ => // do nothing diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/ColumnResolutionHelper.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/ColumnResolutionHelper.scala index 08bfcc446029d..77d267e1d71a9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/ColumnResolutionHelper.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/ColumnResolutionHelper.scala @@ -687,9 +687,11 @@ trait ColumnResolutionHelper extends Logging with DataTypeErrorsBase { // ancestor plan (e.g. a natural/USING join wrapper that hides a join key // via `Project.hiddenOutputTag`). We accept that here but tag the candidate // as `hidden` so the top-level merge in `resolveDataFrameColumn` can prefer - // a regular (p.output) match over hidden (p.metadataOutput) ones. + // a regular (p.output) match over hidden (p.metadataOutput) ones. An attribute can be in both, + // for example when pipe SET retains its input row, and then counts as a regular match. val filtered = candidates.flatMap { c => - val hidden = c.hidden || c.expr.references.subsetOf(AttributeSet(p.metadataOutput)) + val hidden = c.hidden || (!c.expr.references.subsetOf(p.outputSet) && + c.expr.references.subsetOf(AttributeSet(p.metadataOutput))) if (c.expr.references.subsetOf(AttributeSet(p.output ++ p.metadataOutput))) { Some(c.copy(hidden = hidden)) } else { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ExpressionResolver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ExpressionResolver.scala index c763d83234389..d6d90df961cae 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ExpressionResolver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ExpressionResolver.scala @@ -332,6 +332,12 @@ class ExpressionResolver( resolveLiteral(unresolvedLiteral) case unresolvedOrdinal: UnresolvedOrdinal => ordinalResolver.resolve(unresolvedOrdinal) + case pipeExpression: PipeExpression + if pipeExpression.clause == PipeOperators.setClause => + val resolvedChild = resolve(pipeExpression.child) + ValidateAndStripPipeExpressions.validateAndStripPipeExpression( + pipeExpression, + resolvedChild) case unresolvedPredicate: Predicate => resolvePredicate(unresolvedPredicate) case unresolvedScalarSubquery: ScalarSubquery => diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/JoinResolver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/JoinResolver.scala index 1fbd43074154d..6934d4c59c4c0 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/JoinResolver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/JoinResolver.scala @@ -175,13 +175,18 @@ class JoinResolver( val resolvedCondition = resolveJoinCondition(unresolvedJoin, newCondition, leftNameScope, rightNameScope) + val filteredHiddenOutput = filterHiddenOutputMetadataForJoin( + joinType = joinType, + oldHiddenOutput = scopes.current.hiddenOutput, + rightHiddenOutput = rightNameScope.hiddenOutput + ) scopes.overwriteCurrent( output = Some(newOutputList.map(_.toAttribute)), hiddenOutput = Some( computeHiddenOutputForNaturalAndUsingJoin( newHiddenOutput = hiddenList, - oldHiddenOutput = scopes.current.hiddenOutput + oldHiddenOutput = filteredHiddenOutput ) ) ) @@ -263,10 +268,15 @@ class JoinResolver( leftOutput = leftNameScope.output, rightOutput = rightNameScope.output ) + val filteredHiddenOutput = filterHiddenOutputMetadataForJoin( + joinType = partiallyResolvedJoin.joinType, + oldHiddenOutput = scopes.current.hiddenOutput, + rightHiddenOutput = rightNameScope.hiddenOutput + ) val newHiddenOutput = computeHiddenOutputForRegularJoin( mainOutput = newOutput, - oldHiddenOutput = scopes.current.hiddenOutput + oldHiddenOutput = filteredHiddenOutput ) scopes.overwriteCurrent( diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/NameScope.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/NameScope.scala index ad232adfd2733..4e0e446c7e5b9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/NameScope.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/NameScope.scala @@ -417,7 +417,7 @@ class NameScope( */ def isStarQualifiedByTable(unresolvedStar: UnresolvedStarBase): Boolean = { unresolvedStar.isQualifiedByTable( - childOperatorOutput = output, + childOperatorOutput = output ++ hiddenOutput.filter(_.qualifiedAccessOnly), resolver = nameComparator ) } @@ -1101,6 +1101,7 @@ class NameScopeStack( val refreshedHiddenOutput = prevScope.hiddenOutput.map { attribute => outputLookup.get(attribute.exprId) match { + case _ if attribute.qualifiedAccessOnly => attribute case null => attribute case outputAttribute => outputAttribute } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ProjectResolver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ProjectResolver.scala index 4a8b95cffec9b..be529e6afd3d9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ProjectResolver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ProjectResolver.scala @@ -17,8 +17,17 @@ package org.apache.spark.sql.catalyst.analysis.resolver -import org.apache.spark.sql.catalyst.expressions.Expression +import scala.collection.mutable + +import org.apache.spark.sql.catalyst.expressions.{ + Attribute, + Expression, + ExprId, + NamedExpression, + PipeSetInput +} import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, LogicalPlan, Project} +import org.apache.spark.sql.catalyst.util._ import org.apache.spark.sql.internal.SQLConf /** @@ -73,6 +82,13 @@ class ProjectResolver(operatorResolver: Resolver, expressionResolver: Expression val resolvedProjectList = expressionResolver.resolveProjectList(unresolvedProject.projectList, unresolvedProject) + val retainedPipeSetOutput = resolvedChild match { + case pipeSetInput: PipeSetInput => + pipeSetInput.metadataOutput.filter(_.qualifiedAccessOnly) + case _ => + Seq.empty + } + val resolvedChildWithMetadataColumns = retainOriginalJoinOutput( plan = resolvedChild, outputExpressions = resolvedProjectList.expressions, @@ -91,8 +107,17 @@ class ProjectResolver(operatorResolver: Resolver, expressionResolver: Expression resolvedChildWithMetadataColumns = resolvedChildWithMetadataColumns ) } else { + val retainedProjectList = + if (retainedPipeSetOutput.isEmpty || + unresolvedProject.containsTag(ResolverTag.TOP_LEVEL_OPERATOR)) { + Seq.empty + } else { + missingRetainedOutput(retainedPipeSetOutput, resolvedProjectList.expressions) + } val resolvedProject = - Project(resolvedProjectList.expressions, resolvedChildWithMetadataColumns) + Project( + resolvedProjectList.expressions ++ retainedProjectList, + resolvedChildWithMetadataColumns) (resolvedProject, resolvedProjectList) } @@ -109,6 +134,32 @@ class ProjectResolver(operatorResolver: Resolver, expressionResolver: Expression resolvedOperator } + /** + * Returns retained attributes that are not already produced by the visible project list. + * Matching is occurrence-based because a projection can repeat one expression ID. + */ + private def missingRetainedOutput( + retainedOutput: Seq[Attribute], + visibleOutput: Seq[NamedExpression]): Seq[Attribute] = { + val visibleOccurrencesByExprId = mutable.HashMap.empty[ExprId, Int] + visibleOutput.foreach { visibleExpression => + val exprId = visibleExpression.exprId + visibleOccurrencesByExprId.updateWith(exprId) { + case Some(occurrences) => Some(occurrences + 1) + case None => Some(1) + } + } + retainedOutput.filter { retainedAttribute => + visibleOccurrencesByExprId.get(retainedAttribute.exprId) match { + case Some(occurrences) if occurrences > 0 => + visibleOccurrencesByExprId.update(retainedAttribute.exprId, occurrences - 1) + false + case _ => + true + } + } + } + /** * Resolve the original [[Project]] node with aggregate expressions to an appropriate node * ([[Project]], [[Aggregate]] or [[Window]]). diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/PruneMetadataColumns.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/PruneMetadataColumns.scala index 7a437562e32e6..8b641a7265dd8 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/PruneMetadataColumns.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/PruneMetadataColumns.scala @@ -94,20 +94,16 @@ object PruneMetadataColumns extends Rule[LogicalPlan] { */ private def pruneMetadataColumnsInProject(project: Project, neededAttributes: HashSet[ExprId]) = { val existingExprIds = new HashSet[ExprId] - val newProjectList = if (!neededAttributes.isEmpty) { - project.projectList.collect { - case namedExpression: NamedExpression if !namedExpression.toAttribute.qualifiedAccessOnly => - existingExprIds.add(namedExpression.exprId) - namedExpression - case namedExpression: NamedExpression - if namedExpression.toAttribute.qualifiedAccessOnly && neededAttributes.contains( - namedExpression.exprId - ) && !existingExprIds.contains(namedExpression.exprId) => - existingExprIds.add(namedExpression.exprId) - namedExpression - } - } else { - project.projectList + val newProjectList = project.projectList.collect { + case namedExpression: NamedExpression if !namedExpression.toAttribute.qualifiedAccessOnly => + existingExprIds.add(namedExpression.exprId) + namedExpression + case namedExpression: NamedExpression + if namedExpression.toAttribute.qualifiedAccessOnly && neededAttributes.contains( + namedExpression.exprId + ) && !existingExprIds.contains(namedExpression.exprId) => + existingExprIds.add(namedExpression.exprId) + namedExpression } val projectWithNewChildren = withNewChildrenPrunedByNeededAttributes(project, newProjectList).asInstanceOf[Project] diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolutionValidator.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolutionValidator.scala index 452f6d9555aeb..ff08a122eed91 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolutionValidator.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolutionValidator.scala @@ -24,7 +24,7 @@ import org.apache.spark.sql.catalyst.analysis.{ MultiInstanceRelation, ResolvedInlineTable } -import org.apache.spark.sql.catalyst.expressions.AttributeReference +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, PipeSetInput} import org.apache.spark.sql.catalyst.plans.logical._ import org.apache.spark.sql.errors.QueryCompilationErrors import org.apache.spark.sql.types.BooleanType @@ -71,6 +71,8 @@ class ResolutionValidator { validateAggregate(aggregate) case project: Project => validateProject(project) + case pipeSetInput: PipeSetInput => + validatePipeSetInput(pipeSetInput) case filter: Filter => validateFilter(filter) case subqueryAlias: SubqueryAlias => @@ -162,6 +164,11 @@ class ResolutionValidator { handleOperatorOutput(window) } + private def validatePipeSetInput(pipeSetInput: PipeSetInput): Unit = { + validate(pipeSetInput.child) + handleOperatorOutput(pipeSetInput) + } + private def validateCteRelationDef(cteRelationDef: CTERelationDef): Unit = { validate(cteRelationDef.child) } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/Resolver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/Resolver.scala index 458836388a249..54e187436d482 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/Resolver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/Resolver.scala @@ -45,9 +45,12 @@ import org.apache.spark.sql.catalyst.catalog.HiveTableRelation import org.apache.spark.sql.catalyst.expressions.{ Alias, Attribute, + AttributeSeq, AttributeSet, + EliminateResolvedPipeSetInputs, Expression, - ExprId + ExprId, + PipeSetInput } import org.apache.spark.sql.catalyst.plans.logical._ import org.apache.spark.sql.catalyst.rules.Rule @@ -90,7 +93,8 @@ class Resolver( extends LogicalPlanResolver with ResolverMetricTracker with DelegatesResolutionToExtensions - with QueryErrorsBase { + with QueryErrorsBase + with RetainsOriginalJoinOutput { private val planLogger = new PlanLogger private val subqueryRegistry = new SubqueryRegistry private val scopes = new NameScopeStack( @@ -134,7 +138,9 @@ class Resolver( * `planRewriter` is used to rewrite the plan and the subqueries inside by applying * `planRewriteRules`. */ - private val planRewriter = new PlanRewriter(planRewriteRules, extendedRewriteRules) + private val planRewriter = new PlanRewriter( + planRewriteRules, + extendedRewriteRules :+ EliminateResolvedPipeSetInputs) /** * [[relationMetadataProvider]] is used to resolve metadata for relations. It's initialized with @@ -233,9 +239,13 @@ class Resolver( recordProfile("resolve") { resolve(planAfterSubstitution) } + val resolvedPlanWithOriginalJoinOutput = retainOriginalJoinOutputAtBoundary( + plan = resolvedPlan, + outputExpressions = scopes.current.output + ) recordProfile("rewrite") { - planRewriter.rewriteWithSubqueries(resolvedPlan) + planRewriter.rewriteWithSubqueries(resolvedPlanWithOriginalJoinOutput) } } } @@ -273,6 +283,8 @@ class Resolver( handleResolvedWithCte(withCte) case unresolvedProject: Project => projectResolver.resolve(unresolvedProject) + case unresolvedPipeSetInput: PipeSetInput => + resolvePipeSetInput(unresolvedPipeSetInput) case unresolvedAggregate: Aggregate => aggregateResolver.resolve(unresolvedAggregate) case unresolvedFilter: Filter => @@ -348,6 +360,22 @@ class Resolver( } } + /** + * Resolves the input marker of a pipe SET assignment and exposes its original qualified row as + * hidden output. The visible output remains unchanged so unqualified references continue to see + * the values produced by earlier assignments. + */ + private def resolvePipeSetInput(unresolvedPipeSetInput: PipeSetInput): LogicalPlan = { + val resolvedPipeSetInput = + unresolvedPipeSetInput.copy(child = resolve(unresolvedPipeSetInput.child)) + val hiddenOutput = AttributeSeq.mergeHiddenAndVisibleOutput( + scopes.current.hiddenOutput, + resolvedPipeSetInput.metadataOutput + ) + scopes.overwriteCurrent(hiddenOutput = Some(hiddenOutput)) + resolvedPipeSetInput + } + /** * [[UnresolvedWith]] contains a list of unresolved CTE definitions, which are represented by * (name, subquery) pairs, and an actual child query. First we resolve the CTE definitions @@ -367,7 +395,11 @@ class Resolver( cteRegistry.pushScope() val resolvedCtePlan = try { - resolve(cteRelation.plan) + val plan = resolve(cteRelation.plan) + retainOriginalJoinOutputAtBoundary( + plan = plan, + outputExpressions = scopes.current.output + ) } finally { cteRegistry.popScope() scopes.popScope() @@ -472,8 +504,13 @@ class Resolver( * {{{ spark.sql("SELECT * FROM VALUES (1, 2)").select("col1").as("q1").select("col2"); }}} */ private def resolveSubqueryAlias(unresolvedSubqueryAlias: SubqueryAlias): LogicalPlan = { + val resolvedChild = resolve(unresolvedSubqueryAlias.child) + val resolvedChildWithOriginalOutput = retainOriginalJoinOutputAtBoundary( + plan = resolvedChild, + outputExpressions = scopes.current.output + ) val resolvedSubqueryAlias = - unresolvedSubqueryAlias.copy(child = resolve(unresolvedSubqueryAlias.child)) + unresolvedSubqueryAlias.copy(child = resolvedChildWithOriginalOutput) val qualifier = resolvedSubqueryAlias.identifier.qualifier :+ resolvedSubqueryAlias.alias val output = scopes.current.output.map(attribute => attribute.withQualifier(qualifier)) @@ -627,7 +664,12 @@ class Resolver( * programs. In that case we simply recurse into the child plan. */ private def handleResolvedCteRelationDef(cteRelationDef: CTERelationDef): LogicalPlan = { - cteRelationDef.copy(child = resolve(cteRelationDef.child)) + val resolvedChild = resolve(cteRelationDef.child) + val resolvedChildWithOriginalOutput = retainOriginalJoinOutputAtBoundary( + plan = resolvedChild, + outputExpressions = scopes.current.output + ) + cteRelationDef.copy(child = resolvedChildWithOriginalOutput) } /** diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolverGuard.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolverGuard.scala index 96c48843ea364..64b117ee698c6 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolverGuard.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/ResolverGuard.scala @@ -105,6 +105,8 @@ class ResolverGuard( checkWithCte(withCte) case project: Project => checkProject(project) + case pipeSetInput: PipeSetInput => + checkOperator(pipeSetInput.child) case aggregate: Aggregate => checkAggregate(aggregate) case filter: Filter => @@ -198,6 +200,9 @@ class ResolverGuard( checkLiteral(literal) case unresolvedOrdinal: UnresolvedOrdinal => checkUnresolvedOrdinal(unresolvedOrdinal) + case pipeExpression: PipeExpression + if pipeExpression.clause == PipeOperators.setClause => + checkExpression(pipeExpression.child) case unresolvedPredicate: Predicate => checkUnresolvedPredicate(unresolvedPredicate) case scalarSubquery: ScalarSubquery => diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/RetainsOriginalJoinOutput.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/RetainsOriginalJoinOutput.scala index 11ac01b0a08a6..4f3dae23837b1 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/RetainsOriginalJoinOutput.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/RetainsOriginalJoinOutput.scala @@ -21,8 +21,16 @@ import java.util.{HashMap, HashSet} import scala.jdk.CollectionConverters._ -import org.apache.spark.sql.catalyst.expressions.{Attribute, ExprId, NamedExpression} -import org.apache.spark.sql.catalyst.plans.logical.{Join, LogicalPlan, Project} +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, ExprId, NamedExpression} +import org.apache.spark.sql.catalyst.plans.logical.{ + Join, + LogicalPlan, + Project, + SubqueryAlias, + UnaryNode, + WithCTE +} +import org.apache.spark.sql.catalyst.trees.CurrentOrigin import org.apache.spark.sql.catalyst.util._ trait RetainsOriginalJoinOutput { @@ -30,9 +38,11 @@ trait RetainsOriginalJoinOutput { /** * This method adds a [[Project]] node on top of a [[Join]], if [[Join]]'s output has been * changed when metadata columns are added to [[Project]] nodes below the [[Join]]. This is - * necessary in order to stay compatible with fixed-point analyzer. Instead of doing this in - * [[JoinResolver]] we must do this here, because while resolving [[Join]] we still don't know if - * we should add a [[Project]] or not. For example consider the following query: + * necessary in order to stay compatible with fixed-point analyzer. For an intermediate + * [[Join]], this must be done after resolving the parent because the parent determines whether + * the metadata column is part of its output. A top-level [[Join]] or an alias boundary can use + * this method directly because no parent can consume the hidden output. For example consider the + * following query: * * {{{ * -- tables: nt1(k, v1), nt2(k, v2), nt3(k, v3) @@ -131,6 +141,64 @@ trait RetainsOriginalJoinOutput { } } + /** + * Restores the visible output of a [[Join]] at a boundary that discards hidden output. Descend + * through operators that preserve their child's output so the [[Project]] is placed directly on + * top of the [[Join]], matching fixed-point analysis. A [[SubqueryAlias]] is already a boundary + * and restores its child while it is being resolved. An existing [[Project]] is trimmed at the + * boundary instead of descending through it because operators below may still reference its + * hidden attributes. + */ + def retainOriginalJoinOutputAtBoundary( + plan: LogicalPlan, + outputExpressions: Seq[NamedExpression]): LogicalPlan = plan match { + case join: Join => + val project = Project(outputExpressions, join) + if (project.sameOutput(join)) { + join + } else { + project + } + case withCte: WithCTE => + val newPlan = retainOriginalJoinOutputAtBoundary(withCte.plan, outputExpressions) + if (newPlan eq withCte.plan) { + withCte + } else { + withCte.withNewPlan(newPlan) + } + case _: SubqueryAlias => + plan + case project: Project => + val projectListByExpressionId = project.projectList.groupBy(_.exprId) + val newProjectList = outputExpressions.map { outputExpression => + projectListByExpressionId + .get(outputExpression.exprId) + .flatMap(_.headOption) + .getOrElse(outputExpression) + } + if (newProjectList == project.projectList) { + project + } else { + val newProject = CurrentOrigin.withOrigin(project.origin) { + project.copy(projectList = newProjectList) + } + newProject.copyTagsFrom(project) + newProject + } + case unaryNode: UnaryNode + if unaryNode.sameOutput(unaryNode.child) && unaryNode.references.subsetOf( + AttributeSet(outputExpressions.map(_.toAttribute)) + ) => + val newChild = retainOriginalJoinOutputAtBoundary(unaryNode.child, outputExpressions) + if (newChild eq unaryNode.child) { + unaryNode + } else { + unaryNode.withNewChildren(Seq(newChild)) + } + case _ => + plan + } + /** * Returns true if a node has missing attributes that can be resolved from * [[NameScope.hiddenOutput]] and those attributes are not present in the output. diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/SortResolver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/SortResolver.scala index 1219d16e4c831..44eb57e1dcce2 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/SortResolver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/resolver/SortResolver.scala @@ -40,6 +40,8 @@ import org.apache.spark.sql.catalyst.plans.logical.{ Project, Sort } +import org.apache.spark.sql.catalyst.util.MetadataColumnHelper + /** * Resolves a [[Sort]] by resolving its child and order expressions. */ @@ -178,7 +180,16 @@ class SortResolver(operatorResolver: Resolver, expressionResolver: ExpressionRes val resolvedChildWithMissingAttributes = insertMissingExpressions(resolvedChild, filteredMissingExpressions) - val isChildChangedByMissingExpressions = !resolvedChildWithMissingAttributes.eq(resolvedChild) + val isChildChangedByMissingExpressions = + !resolvedChildWithMissingAttributes.eq(resolvedChild) + val hasAlreadyMaterializedRetainedMissingAttribute = missingAttributes.exists { attribute => + attribute.qualifiedAccessOnly && + attribute.pipeSetRetained && + filteredMissingExpressions.exists(_.exprId == attribute.exprId) && + resolvedChild.outputSet.contains(attribute) + } + val shouldRetainOriginalOutput = + isChildChangedByMissingExpressions || hasAlreadyMaterializedRetainedMissingAttribute val (finalChild, finalOrderExpressions) = resolvedChildWithMissingAttributes match { case project: Project if scopes.current.baseAggregate.isDefined => @@ -197,7 +208,7 @@ class SortResolver(operatorResolver: Resolver, expressionResolver: ExpressionRes order = finalOrderExpressions ) - if (isChildChangedByMissingExpressions) { + if (shouldRetainOriginalOutput) { retainOriginalOutput( operator = resolvedSort, missingExpressions = missingExpressions, diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/unresolved.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/unresolved.scala index 6765af018b2ab..ccd115cadad45 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/unresolved.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/unresolved.scala @@ -590,7 +590,8 @@ trait UnresolvedStarBase extends Star with Unevaluable { // keep any restrictions that may break column resolution for normal attributes. // See SPARK-42084 for more details. .map(_.markAsAllowAnyAccess()) - val expandedAttributes = (hiddenOutput ++ parameters.childOperatorOutput) + val expandedAttributes = AttributeSeq + .mergeHiddenAndVisibleOutput(hiddenOutput, parameters.childOperatorOutput) .filter(matchedQualifier(_, target.get, parameters.resolver)) if (expandedAttributes.nonEmpty) return expandedAttributes diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/package.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/package.scala index b3525da97f8de..4ef8809e77fe5 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/package.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/package.scala @@ -19,6 +19,8 @@ package org.apache.spark.sql.catalyst import java.util.Locale +import scala.collection.mutable + import com.google.common.collect.Maps import org.apache.spark.sql.catalyst.analysis.{Resolver, UnresolvedAttribute} @@ -84,6 +86,44 @@ package object expressions { // we explicitly remove that special flag to be safe. new AttributeSeq(attr.map(_.markAsAllowAnyAccess())) } + + /** + * Merges hidden and visible output in the order used by qualified star expansion. + * + * Hidden output can repeat visible attributes to preserve their position in an original row. + * Each hidden occurrence therefore consumes at most one visible occurrence with the same + * expression ID. Any visible occurrences that were not consumed are appended in their + * original order. + */ + def mergeHiddenAndVisibleOutput( + hiddenOutput: Seq[Attribute], + visibleOutput: Seq[Attribute]): Seq[Attribute] = { + if (hiddenOutput.isEmpty) { + return visibleOutput + } + + val visible = visibleOutput.toIndexedSeq + val visibleIndicesByExprId = mutable.HashMap.empty[ExprId, mutable.Queue[Int]] + visible.zipWithIndex.foreach { case (attribute, index) => + visibleIndicesByExprId + .getOrElseUpdate(attribute.exprId, mutable.Queue.empty[Int]) + .enqueue(index) + } + val consumedVisible = Array.fill(visible.length)(false) + val mergedHidden = hiddenOutput.map { hiddenAttribute => + visibleIndicesByExprId.get(hiddenAttribute.exprId) match { + case Some(visibleIndices) if visibleIndices.nonEmpty => + val visibleIndex = visibleIndices.dequeue() + consumedVisible(visibleIndex) = true + visible(visibleIndex) + case _ => + hiddenAttribute + } + } + mergedHidden ++ visible.zipWithIndex.collect { + case (attribute, index) if !consumedVisible(index) => attribute + } + } } /** diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/pipeOperators.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/pipeOperators.scala index cfbd403d66fdb..211a89e8376a7 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/pipeOperators.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/pipeOperators.scala @@ -18,9 +18,15 @@ package org.apache.spark.sql.catalyst.expressions import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateFunction -import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode} +import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, Project, UnaryNode} import org.apache.spark.sql.catalyst.rules.Rule -import org.apache.spark.sql.catalyst.trees.TreePattern.{PIPE_EXPRESSION, PIPE_OPERATOR, TreePattern} +import org.apache.spark.sql.catalyst.trees.TreePattern.{ + PIPE_EXPRESSION, + PIPE_OPERATOR, + PLAN_EXPRESSION, + TreePattern +} +import org.apache.spark.sql.catalyst.util._ import org.apache.spark.sql.errors.QueryCompilationErrors import org.apache.spark.sql.types.DataType @@ -53,11 +59,72 @@ case class PipeOperator(child: LogicalPlan) extends UnaryNode { override def withNewChildInternal(newChild: LogicalPlan): PipeOperator = copy(child = newChild) } -/** This rule removes all PipeOperator nodes from a logical plan at the end of analysis. */ +/** + * Preserves the qualified input row of a SQL pipe SET assignment as metadata output. + * + * The SET projection replaces assigned columns in its visible output. Keeping the original row as + * qualified-access-only metadata lets a later `t.a` or `t.*` continue to refer to the input row + * without changing the SET output schema or unqualified star expansion. + */ +case class PipeSetInput(child: LogicalPlan) extends UnaryNode { + final override val nodePatterns: Seq[TreePattern] = Seq(PIPE_OPERATOR) + + override def output: Seq[Attribute] = child.output + override def maxRows: Option[Long] = child.maxRows + override def maxRowsPerPartition: Option[Long] = child.maxRowsPerPartition + + override def metadataOutput: Seq[Attribute] = { + val childMetadataOutput = child.metadataOutput + val retainedQualifiedOutput = AttributeSeq + .mergeHiddenAndVisibleOutput( + childMetadataOutput.filter(_.qualifiedAccessOnly), child.output) + .filter(_.qualifier.nonEmpty) + .map(_.markAsQualifiedAccessOnly().markAsPipeSetRetained()) + val retainedQualifiedOutputIds = retainedQualifiedOutput.iterator.map(_.exprId).toSet + retainedQualifiedOutput ++ childMetadataOutput.filterNot { attribute => + attribute.qualifiedAccessOnly || retainedQualifiedOutputIds.contains(attribute.exprId) + } + } + + override def withNewChildInternal(newChild: LogicalPlan): PipeSetInput = copy(child = newChild) +} + +/** + * Removes resolved [[PipeSetInput]] nodes before analyzed plans cross a Dataset boundary. + * + * Qualified source attributes referenced later in the same pipe query have already been + * materialized by this point. Removing the marker and any retained attributes copied to hidden + * output prevents them from being resolved by a subsequent Dataset operation. + */ +object EliminateResolvedPipeSetInputs extends Rule[LogicalPlan] { + override def apply(plan: LogicalPlan): LogicalPlan = { + val planWithCleanedHiddenOutput = + plan.transformDownWithSubqueriesAndReferenceEquality { + case project: Project => + project.getTagValue(Project.hiddenOutputTag) match { + case Some(hiddenOutput) if hiddenOutput.exists(_.pipeSetRetained) => + val cleanedHiddenOutput = hiddenOutput.filterNot(_.pipeSetRetained) + val cleanedProject = project.copy() + cleanedProject.copyTagsFrom(project) + cleanedProject.setTagValue(Project.hiddenOutputTag, cleanedHiddenOutput) + cleanedProject + case _ => project + } + } + + planWithCleanedHiddenOutput.resolveOperatorsUpWithSubqueriesAndPruning( + _.containsAnyPattern(PIPE_OPERATOR, PLAN_EXPRESSION), ruleId) { + case pipeSetInput: PipeSetInput if pipeSetInput.resolved => pipeSetInput.child + } + } +} + +/** This rule removes transparent pipe-operator nodes from a logical plan after analysis. */ object EliminatePipeOperators extends Rule[LogicalPlan] { def apply(plan: LogicalPlan): LogicalPlan = plan.transformWithPruning( _.containsPattern(PIPE_OPERATOR), ruleId) { case PipeOperator(child) => child + case PipeSetInput(child) => child } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index 8f09dde0c03c1..0ef50133801c4 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -8215,7 +8215,8 @@ class AstBuilder extends DataTypeAstBuilder // Add a projection to implement the SET operator using the UnresolvedStarExceptOrReplace // expression. We do this once per SET assignment to allow for multiple SET assignments with // optional lateral references to previous ones. - plan = Project(projectList, plan) + // PipeSetInput retains qualified access to the row before the assignment. + plan = Project(projectList, PipeSetInput(plan)) } plan } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/LogicalPlan.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/LogicalPlan.scala index 757da1dccdd9e..1ab42b4122e19 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/LogicalPlan.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/LogicalPlan.scala @@ -157,13 +157,30 @@ abstract class LogicalPlan } } - private[this] lazy val childAttributes = AttributeSeq.fromNormalOutput(children.flatMap(_.output)) + private def attributesForResolution( + output: Seq[Attribute], + metadataOutput: Seq[Attribute]): AttributeSeq = { + lazy val outputSet = AttributeSet(output) + new AttributeSeq( + output.map(_.markAsAllowAnyAccess()) ++ + metadataOutput.filter(attribute => + attribute.qualifiedAccessOnly && !outputSet.contains(attribute))) + } + + private[this] lazy val childOutput = children.flatMap(_.output) + + private[this] lazy val childMetadataOutput = children.flatMap(_.metadataOutput) + + private[this] lazy val childAttributes = + attributesForResolution(childOutput, childMetadataOutput) - private[this] lazy val childMetadataAttributes = AttributeSeq(children.flatMap(_.metadataOutput)) + private[this] lazy val childMetadataAttributes = + new AttributeSeq(childMetadataOutput.filterNot(_.qualifiedAccessOnly)) - private[this] lazy val outputAttributes = AttributeSeq.fromNormalOutput(output) + private[this] lazy val outputAttributes = attributesForResolution(output, metadataOutput) - private[this] lazy val outputMetadataAttributes = AttributeSeq(metadataOutput) + private[this] lazy val outputMetadataAttributes = + new AttributeSeq(metadataOutput.filterNot(_.qualifiedAccessOnly)) /** * Optionally resolves the given strings to a [[NamedExpression]] using the input from all child diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/basicLogicalOperators.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/basicLogicalOperators.scala index 0b70b973aec65..02d18bf3794ad 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/basicLogicalOperators.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/basicLogicalOperators.scala @@ -2243,6 +2243,7 @@ case class Sample( */ case class Distinct(child: LogicalPlan) extends UnaryNode { override def maxRows: Option[Long] = child.maxRows + override def metadataOutput: Seq[Attribute] = child.metadataOutput.filterNot(_.pipeSetRetained) override def output: Seq[Attribute] = { val base = child.output if (isStateful) WidenStatefulOpNullability.widenOutputForStatefulOp(base) else base @@ -2441,6 +2442,7 @@ case class Deduplicate( child: LogicalPlan, dedupSpec: Option[DeduplicateSpec] = None) extends UnaryNode { override def maxRows: Option[Long] = child.maxRows + override def metadataOutput: Seq[Attribute] = child.metadataOutput.filterNot(_.pipeSetRetained) override def output: Seq[Attribute] = { val base = child.output if (isStateful) WidenStatefulOpNullability.widenOutputForStatefulOp(base) else base @@ -2462,6 +2464,7 @@ case class DeduplicateWithinWatermark( override def references: AttributeSet = AttributeSet(keys) ++ AttributeSet(child.output.filter(_.metadata.contains(EventTimeWatermark.delayKey))) override def maxRows: Option[Long] = child.maxRows + override def metadataOutput: Seq[Attribute] = child.metadataOutput.filterNot(_.pipeSetRetained) override def output: Seq[Attribute] = { val base = child.output if (isStateful) WidenStatefulOpNullability.widenOutputForStatefulOp(base) else base diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala index cd1d14e1bd6c1..416c502f232c6 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala @@ -119,6 +119,7 @@ object RuleIdCollection { "org.apache.spark.sql.catalyst.analysis.UpdateAttributeNullability" :: "org.apache.spark.sql.catalyst.analysis.UpdateOuterReferences" :: "org.apache.spark.sql.catalyst.expressions.EliminatePipeOperators" :: + "org.apache.spark.sql.catalyst.expressions.EliminateResolvedPipeSetInputs" :: "org.apache.spark.sql.catalyst.expressions.ExtractSemiStructuredFields" :: "org.apache.spark.sql.catalyst.expressions.ValidateAndStripPipeExpressions" :: // Catalyst Optimizer rules diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/package.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/package.scala index 2a412c538b0c1..1d827ad6f76b1 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/package.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/package.scala @@ -192,6 +192,9 @@ package object util extends Logging { */ val QUALIFIED_ACCESS_ONLY = "__qualified_access_only" + /** If set, this metadata column retains an input attribute for a SQL pipe SET assignment. */ + val PIPE_SET_RETAINED = "__pipe_set_retained" + /** * If set, this column can only be accessed under [[AggregateExpression]]. This is important when * resolving columns in ORDER BY and HAVING clauses on top of [[Aggregate]]. In this case we can @@ -208,6 +211,9 @@ package object util extends Logging { attr.metadata.contains(QUALIFIED_ACCESS_ONLY) && attr.metadata.getBoolean(QUALIFIED_ACCESS_ONLY) + def pipeSetRetained: Boolean = attr.metadata.contains(PIPE_SET_RETAINED) && + attr.metadata.getBoolean(PIPE_SET_RETAINED) + def aggregatedAccessOnly: Boolean = attr.metadata.contains(AGGREGATED_ACCESS_ONLY) && attr.metadata.getBoolean(AGGREGATED_ACCESS_ONLY) @@ -226,13 +232,21 @@ package object util extends Logging { .build() ) + def markAsPipeSetRetained(): Attribute = attr.withMetadata( + new MetadataBuilder() + .withMetadata(attr.metadata) + .putBoolean(PIPE_SET_RETAINED, true) + .build() + ) + def markAsAllowAnyAccess(): Attribute = { - if (qualifiedAccessOnly) { + if (qualifiedAccessOnly || pipeSetRetained) { attr.withMetadata( new MetadataBuilder() .withMetadata(attr.metadata) .remove(QUALIFIED_ACCESS_ONLY) .remove(AGGREGATED_ACCESS_ONLY) + .remove(PIPE_SET_RETAINED) .build() ) } else { @@ -247,6 +261,7 @@ package object util extends Logging { AUTO_GENERATED_ALIAS, METADATA_COL_ATTR_KEY, QUALIFIED_ACCESS_ONLY, + PIPE_SET_RETAINED, FIELD_ID_METADATA_KEY, FileSourceMetadataAttribute.FILE_SOURCE_METADATA_COL_ATTR_KEY, FileSourceConstantMetadataStructField.FILE_SOURCE_CONSTANT_METADATA_COL_ATTR_KEY, diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index 8ace371566ff0..98168f9eab9e9 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1893,6 +1893,152 @@ class AnalysisSuite extends AnalysisTest with Matchers { val expectedPlan = Project(Seq(UnresolvedAttribute("i")), addColumnF).analyze checkAnalysis(inputPlan, expectedPlan) } + + test("SPARK-59146: pipe SET retained output remains a regular plan-ID candidate") { + val source = testRelation2.subquery("t") + val Seq(_, b) = source.output.take(2) + val set = Project(Seq(Alias(Literal("x"), "a")(), b), PipeSetInput(source)) + set.setTagValue(LogicalPlan.PLAN_ID_TAG, 1L) + + val otherRelation = LocalRelation(AttributeReference("b", StringType)()) + val other = Project(otherRelation.output, otherRelation) + other.setTagValue(LogicalPlan.PLAN_ID_TAG, 1L) + + val column = UnresolvedAttribute("b") + column.setTagValue(LogicalPlan.PLAN_ID_TAG, 1L) + val plan = Project(Seq(column), Join(set, other, Inner, None, JoinHint.NONE)) + + checkError( + exception = intercept[AnalysisException](getAnalyzer.execute(plan)), + condition = "AMBIGUOUS_COLUMN_REFERENCE", + parameters = Map("name" -> "\"b\"")) + } + + test("SPARK-59146: pipe SET does not duplicate retained metadata output") { + val a = AttributeReference("a", IntegerType)().withQualifier(Seq("t")) + val metadata = MetadataAttribute("_metadata", StringType).withQualifier(Seq("t")) + val child = Project(Seq(a, metadata), LocalRelation(a, metadata)) + child.setTagValue(Project.hiddenOutputTag, Seq(metadata)) + + val metadataOutput = PipeSetInput(child).metadataOutput + assert(metadataOutput.map(_.exprId) === Seq(a.exprId, metadata.exprId)) + assert(metadataOutput.forall(_.qualifiedAccessOnly)) + } + + test("SPARK-59146: pipe SET input forwards cardinality bounds") { + case class CardinalityLeaf(override val output: Seq[Attribute]) extends LeafNode { + override def maxRows: Option[Long] = Some(2L) + override def maxRowsPerPartition: Option[Long] = Some(1L) + } + + val child = CardinalityLeaf(testRelation.output) + val setInput = PipeSetInput(child) + + assert(setInput.maxRows === child.maxRows) + assert(setInput.maxRowsPerPartition === child.maxRowsPerPartition) + } + + test("SPARK-59146: pipe SET cleanup reaches subqueries") { + val attribute = AttributeReference("a", IntegerType)() + val subquery = ScalarSubquery(PipeSetInput(LocalRelation(attribute))) + val plan = Project(Seq(Alias(subquery, "a")()), OneRowRelation()) + + val cleaned = EliminateResolvedPipeSetInputs(plan) + + assert(cleaned.collectWithSubqueries { case _: PipeSetInput => () }.isEmpty) + } + + test("SPARK-59146: pipe SET cleanup removes nested retained hidden output") { + val attribute = AttributeReference("a", IntegerType)().withQualifier(Seq("t")) + val retained = attribute.markAsQualifiedAccessOnly().markAsPipeSetRetained() + val ordinary = AttributeReference("a", IntegerType)() + .withQualifier(Seq("u")) + .markAsQualifiedAccessOnly() + val taggedProject = Project(Seq(attribute), LocalRelation(attribute)) + taggedProject.setTagValue(Project.hiddenOutputTag, Seq(retained, ordinary)) + val plan = SubqueryAlias("outer", taggedProject) + + val cleaned = EliminateResolvedPipeSetInputs(plan) + val cleanedProject = cleaned.collectFirst { case project: Project => project }.get + + assert(cleanedProject.getTagValue(Project.hiddenOutputTag).contains(Seq(ordinary))) + } + + test("SPARK-59146: pipe SET cleanup invalidates cached hidden output") { + val visible = AttributeReference("a", IntegerType)() + val retained = AttributeReference("a", IntegerType)() + .withQualifier(Seq("t")) + .markAsQualifiedAccessOnly() + .markAsPipeSetRetained() + val project = Project(Seq(visible), LocalRelation(visible, retained)) + project.setTagValue(Project.hiddenOutputTag, Seq(retained)) + + assert(project.resolve(Seq("t", "a"), caseInsensitiveResolution).nonEmpty) + + val cleanedProject = EliminateResolvedPipeSetInputs(project).asInstanceOf[Project] + + assert(cleanedProject ne project) + assert(cleanedProject.resolve(Seq("t", "a"), caseInsensitiveResolution).isEmpty) + } + + test("SPARK-59146: pipe SET cleanup preserves empty hidden output overrides") { + case class MetadataLeaf( + override val output: Seq[Attribute], + override val metadataOutput: Seq[Attribute]) extends LeafNode + + val attribute = AttributeReference("a", IntegerType)().withQualifier(Seq("t")) + val metadata = MetadataAttribute("_metadata", StringType) + val child = MetadataLeaf(Seq(attribute), Seq(metadata)) + val project = Project(Seq(attribute), PipeSetInput(child)) + val retained = attribute.markAsQualifiedAccessOnly().markAsPipeSetRetained() + project.setTagValue(Project.hiddenOutputTag, Seq(retained)) + + val cleanedProject = EliminateResolvedPipeSetInputs(project).asInstanceOf[Project] + + assert(!cleanedProject.exists(_.isInstanceOf[PipeSetInput])) + assert(cleanedProject.getTagValue(Project.hiddenOutputTag).contains(Nil)) + assert(cleanedProject.metadataOutput.isEmpty) + } + + test("SPARK-59146: pipe SET metadata lookup is linear in assignment count") { + case class CountingMetadataLeaf(override val output: Seq[Attribute]) extends LeafNode { + var metadataOutputCalls: Int = 0 + + override def metadataOutput: Seq[Attribute] = { + metadataOutputCalls += 1 + Nil + } + } + + val attribute = AttributeReference("a", IntegerType)().withQualifier(Seq("t")) + val leaf = CountingMetadataLeaf(Seq(attribute)) + val assignments = (1 to 10).foldLeft[LogicalPlan](leaf) { case (child, _) => + val setInput = PipeSetInput(child) + Project(setInput.output, setInput) + } + + assignments.metadataOutput + assert(leaf.metadataOutputCalls === 1) + } + + test("SPARK-59146: distinct-like operators discard pipe SET retained metadata") { + case class MetadataLeaf( + override val output: Seq[Attribute], + override val metadataOutput: Seq[Attribute]) extends LeafNode + + val attribute = AttributeReference("a", IntegerType)().withQualifier(Seq("t")) + val metadata = MetadataAttribute("_metadata", StringType) + val setInput = PipeSetInput(MetadataLeaf(Seq(attribute), Seq(metadata))) + val plans = Seq[LogicalPlan]( + Distinct(setInput), + Deduplicate(setInput.output, setInput), + DeduplicateWithinWatermark(setInput.output, setInput) + ) + + plans.foreach { plan => + assert(plan.metadataOutput.map(_.exprId) === Seq(metadata.exprId), plan.toString) + } + } } /** diff --git a/sql/core/src/test/resources/sql-tests/analyzer-results/pipe-operators.sql.out b/sql/core/src/test/resources/sql-tests/analyzer-results/pipe-operators.sql.out index 94fbe261ca29c..e5af579cba71a 100644 --- a/sql/core/src/test/resources/sql-tests/analyzer-results/pipe-operators.sql.out +++ b/sql/core/src/test/resources/sql-tests/analyzer-results/pipe-operators.sql.out @@ -992,6 +992,227 @@ Project [a#x, a#x, z2#x] +- LocalRelation [a#x] +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select t.a +-- !query analysis +Project [a#x] ++- Project [(a#x + 1) AS a#x, b#x, a#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x] + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select a, t.a, t.b +-- !query analysis +Project [a#x, a#x, b#x] ++- Project [(a#x + 1) AS a#x, b#x, a#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x] + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> as u +-- !query analysis +SubqueryAlias u ++- Project [(a#x + 1) AS a#x, b#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x] + + +-- !query +select 1 as x, named_struct('x', 2) as col +|> as col +|> set x = x + 1 +|> select x, col.x +-- !query analysis +Project [x#x, x#x] ++- Project [(x#x + 1) AS x#x, col#x, x#x] + +- SubqueryAlias col + +- Project [1 AS x#x, named_struct(x, 2) AS col#x] + +- OneRowRelation + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select * +-- !query analysis +Project [a#x, b#x] ++- Project [(a#x + 1) AS a#x, b#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x] + + +-- !query +values (1, 2, 3) as t(a, b, c) +|> set b = 20 +|> select t.* +-- !query analysis +Project [a#x, b#x, c#x] ++- Project [a#x, 20 AS b#x, c#x, b#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x, c#x] + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1, b = b + 1 +|> select a, b, t.* +-- !query analysis +Project [a#x, b#x, a#x, b#x] ++- Project [a#x, (b#x + 1) AS b#x, a#x, b#x] + +- Project [(a#x + 1) AS a#x, b#x, a#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x] + + +-- !query +values (1, 2) as s(a, b) +|> select a, a, b +|> as t +|> set b = 3 +|> select t.* +-- !query analysis +Project [a#x, a#x, b#x] ++- Project [a#x, a#x, 3 AS b#x, b#x] + +- SubqueryAlias t + +- Project [a#x, a#x, b#x] + +- SubqueryAlias s + +- LocalRelation [a#x, b#x] + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> where t.a = 1 +|> select a, t.a +-- !query analysis +Project [a#x, a#x] ++- Filter (a#x = 1) + +- PipeOperator + +- Project [(a#x + 1) AS a#x, b#x, a#x] + +- SubqueryAlias t + +- LocalRelation [a#x, b#x] + + +-- !query +values (1), (2) as t(a) +|> set a = 0 +|> select distinct a +|> select t.a +-- !query analysis +org.apache.spark.sql.catalyst.ExtendedAnalysisException +{ + "errorClass" : "UNRESOLVED_COLUMN.WITH_SUGGESTION", + "sqlState" : "42703", + "messageParameters" : { + "objectName" : "`t`.`a`", + "proposal" : "`a`" + }, + "queryContext" : [ { + "objectType" : "", + "objectName" : "", + "startIndex" : 69, + "stopIndex" : 71, + "fragment" : "t.a" + } ] +} + + +-- !query +values (1), (2) as t(a) +|> set a = 0 +|> select distinct a, t.a +|> select t.a +|> order by t.a +-- !query analysis +Sort [a#x ASC NULLS FIRST], true ++- PipeOperator + +- Project [a#x] + +- Distinct + +- Project [a#x, a#x] + +- Project [0 AS a#x, a#x] + +- SubqueryAlias t + +- LocalRelation [a#x] + + +-- !query +values (0) as t(x) +|> set x = x + 1 +|> join lateral (select x + rand(0) as y) +|> select x, y >= x as y_at_least_x +-- !query analysis +[Analyzer test output redacted due to nondeterminism] + + +-- !query +values (1, 10) as lhs(k, l) +|> inner join values (1, 20) as rhs(k, r) using (k) +|> set k = k + 1 +|> select k, lhs.k, rhs.k, l, r +-- !query analysis +Project [k#x, k#x, k#x, l#x, r#x] ++- Project [(k#x + 1) AS k#x, l#x, r#x, k#x, k#x] + +- Project [k#x, l#x, r#x, k#x] + +- Join Inner, (k#x = k#x) + :- SubqueryAlias lhs + : +- LocalRelation [k#x, l#x] + +- SubqueryAlias rhs + +- LocalRelation [k#x, r#x] + + +-- !query +values (1) as t(a) +|> set a = a + 1 +|> aggregate count(t.*) +-- !query analysis +org.apache.spark.sql.AnalysisException +{ + "errorClass" : "INVALID_USAGE_OF_STAR_WITH_TABLE_IDENTIFIER_IN_COUNT", + "sqlState" : "42000", + "messageParameters" : { + "tableName" : "`t`" + } +} + + +-- !query +table t +|> set x = x + 1 +|> select x, _metadata.file_name is not null as has_file_name +|> order by x +-- !query analysis +Sort [x#x ASC NULLS FIRST], true ++- PipeOperator + +- Project [x#x, isnotnull(_metadata#x.file_name) AS has_file_name#x] + +- Project [(x#x + 1) AS x#x, y#x, _metadata#x] + +- SubqueryAlias spark_catalog.default.t + +- Relation spark_catalog.default.t[x#x,y#x,_metadata#x] csv + + +-- !query +table t +|> select *, _metadata +|> set _metadata = 1 +|> select x, _metadata, + t._metadata.file_name is not null as original_metadata_available +|> order by x +-- !query analysis +Sort [x#x ASC NULLS FIRST], true ++- PipeOperator + +- Project [x#x, _metadata#x, isnotnull(_metadata#x.file_name) AS original_metadata_available#x] + +- Project [x#x, y#x, 1 AS _metadata#x, _metadata#x] + +- Project [x#x, y#x, _metadata#x] + +- SubqueryAlias spark_catalog.default.t + +- Relation spark_catalog.default.t[x#x,y#x,_metadata#x] csv + + -- !query table t |> set z = 1 diff --git a/sql/core/src/test/resources/sql-tests/inputs/pipe-operators.sql b/sql/core/src/test/resources/sql-tests/inputs/pipe-operators.sql index 916b114f17f57..71725ba37366a 100644 --- a/sql/core/src/test/resources/sql-tests/inputs/pipe-operators.sql +++ b/sql/core/src/test/resources/sql-tests/inputs/pipe-operators.sql @@ -361,6 +361,99 @@ values (0), (1) lhs(a) |> limit 2 |> select lhs.a, rhs.a, z2; +-- A table alias refers to the original value of a column affected by SET. +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select t.a; + +-- Unqualified names expose assigned values while qualified names expose the original row. +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select a, t.a, t.b; + +-- A trailing alias does not expose the retained input row in the visible schema. +values (1, 10) as t(a, b) +|> set a = a + 1 +|> as u; + +-- A qualified column takes precedence over a visible struct field with the same multipart name. +select 1 as x, named_struct('x', 2) as col +|> as col +|> set x = x + 1 +|> select x, col.x; + +-- Retaining the original row does not change the visible SET schema or unqualified star. +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select *; + +-- A qualified star returns the original row in its original column order. +values (1, 2, 3) as t(a, b, c) +|> set b = 20 +|> select t.*; + +-- The alias continues to refer to the input row across sequential assignments. +values (1, 10) as t(a, b) +|> set a = a + 1, b = b + 1 +|> select a, b, t.*; + +-- Repeated source attributes retain every position in a qualified star. +values (1, 2) as s(a, b) +|> select a, a, b +|> as t +|> set b = 3 +|> select t.*; + +-- Qualified source values remain available through intervening pipe operators. +values (1, 10) as t(a, b) +|> set a = a + 1 +|> where t.a = 1 +|> select a, t.a; + +-- Qualified source values do not cross a DISTINCT boundary. +values (1), (2) as t(a) +|> set a = 0 +|> select distinct a +|> select t.a; + +-- An explicitly selected source value becomes part of the DISTINCT key. +values (1), (2) as t(a) +|> set a = 0 +|> select distinct a, t.a +|> select t.a +|> order by t.a; + +-- A one-row SET input remains eligible for a nondeterministic lateral subquery. +values (0) as t(x) +|> set x = x + 1 +|> join lateral (select x + rand(0) as y) +|> select x, y >= x as y_at_least_x; + +-- Both sides of a USING join retain qualified access to an assigned key. +values (1, 10) as lhs(k, l) +|> inner join values (1, 20) as rhs(k, r) using (k) +|> set k = k + 1 +|> select k, lhs.k, rhs.k, l, r; + +-- A qualified star in COUNT remains invalid when SET replaces every visible qualified column. +values (1) as t(a) +|> set a = a + 1 +|> aggregate count(t.*); + +-- An unselected metadata column remains available after SET. +table t +|> set x = x + 1 +|> select x, _metadata.file_name is not null as has_file_name +|> order by x; + +-- A selected metadata column remains qualified-accessible after SET replaces it. +table t +|> select *, _metadata +|> set _metadata = 1 +|> select x, _metadata, + t._metadata.file_name is not null as original_metadata_available +|> order by x; + -- SET operators: negative tests. --------------------------------- diff --git a/sql/core/src/test/resources/sql-tests/results/pipe-operators.sql.out b/sql/core/src/test/resources/sql-tests/results/pipe-operators.sql.out index 66cf72698a2d1..d0a0e0e61fb20 100644 --- a/sql/core/src/test/resources/sql-tests/results/pipe-operators.sql.out +++ b/sql/core/src/test/resources/sql-tests/results/pipe-operators.sql.out @@ -909,6 +909,204 @@ struct 1 1 4 +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select t.a +-- !query schema +struct +-- !query output +1 + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select a, t.a, t.b +-- !query schema +struct +-- !query output +2 1 10 + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> as u +-- !query schema +struct +-- !query output +2 10 + + +-- !query +select 1 as x, named_struct('x', 2) as col +|> as col +|> set x = x + 1 +|> select x, col.x +-- !query schema +struct +-- !query output +2 1 + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> select * +-- !query schema +struct +-- !query output +2 10 + + +-- !query +values (1, 2, 3) as t(a, b, c) +|> set b = 20 +|> select t.* +-- !query schema +struct +-- !query output +1 2 3 + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1, b = b + 1 +|> select a, b, t.* +-- !query schema +struct +-- !query output +2 11 1 10 + + +-- !query +values (1, 2) as s(a, b) +|> select a, a, b +|> as t +|> set b = 3 +|> select t.* +-- !query schema +struct +-- !query output +1 1 2 + + +-- !query +values (1, 10) as t(a, b) +|> set a = a + 1 +|> where t.a = 1 +|> select a, t.a +-- !query schema +struct +-- !query output +2 1 + + +-- !query +values (1), (2) as t(a) +|> set a = 0 +|> select distinct a +|> select t.a +-- !query schema +struct<> +-- !query output +org.apache.spark.sql.catalyst.ExtendedAnalysisException +{ + "errorClass" : "UNRESOLVED_COLUMN.WITH_SUGGESTION", + "sqlState" : "42703", + "messageParameters" : { + "objectName" : "`t`.`a`", + "proposal" : "`a`" + }, + "queryContext" : [ { + "objectType" : "", + "objectName" : "", + "startIndex" : 69, + "stopIndex" : 71, + "fragment" : "t.a" + } ] +} + + +-- !query +values (1), (2) as t(a) +|> set a = 0 +|> select distinct a, t.a +|> select t.a +|> order by t.a +-- !query schema +struct +-- !query output +1 +2 + + +-- !query +values (0) as t(x) +|> set x = x + 1 +|> join lateral (select x + rand(0) as y) +|> select x, y >= x as y_at_least_x +-- !query schema +struct +-- !query output +1 true + + +-- !query +values (1, 10) as lhs(k, l) +|> inner join values (1, 20) as rhs(k, r) using (k) +|> set k = k + 1 +|> select k, lhs.k, rhs.k, l, r +-- !query schema +struct +-- !query output +2 1 1 10 20 + + +-- !query +values (1) as t(a) +|> set a = a + 1 +|> aggregate count(t.*) +-- !query schema +struct<> +-- !query output +org.apache.spark.sql.AnalysisException +{ + "errorClass" : "INVALID_USAGE_OF_STAR_WITH_TABLE_IDENTIFIER_IN_COUNT", + "sqlState" : "42000", + "messageParameters" : { + "tableName" : "`t`" + } +} + + +-- !query +table t +|> set x = x + 1 +|> select x, _metadata.file_name is not null as has_file_name +|> order by x +-- !query schema +struct +-- !query output +1 true +2 true + + +-- !query +table t +|> select *, _metadata +|> set _metadata = 1 +|> select x, _metadata, + t._metadata.file_name is not null as original_metadata_available +|> order by x +-- !query schema +struct +-- !query output +0 1 true +1 1 true + + -- !query table t |> set z = 1 diff --git a/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverGuardSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverGuardSuite.scala index 5b93c14bf9627..e500a34f0434f 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverGuardSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverGuardSuite.scala @@ -17,15 +17,17 @@ package org.apache.spark.sql.analysis.resolver -import org.apache.spark.SparkException +import org.apache.spark.{SparkException, SparkThrowable} +import org.apache.spark.sql.Row import org.apache.spark.sql.catalyst.analysis.resolver.{ AnalyzerBridgeState, ExplicitlyUnsupportedResolverFeature, Resolver, ResolverGuard } -import org.apache.spark.sql.catalyst.expressions.Literal +import org.apache.spark.sql.catalyst.expressions.{Literal, PipeOperator, PipeSetInput} import org.apache.spark.sql.catalyst.plans.logical._ +import org.apache.spark.sql.catalyst.util._ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.SharedSparkSession @@ -109,6 +111,304 @@ class ResolverGuardSuite extends ResolverGuardSuiteBase { checkResolverGuard("SELECT table.* FROM VALUES(1) as table") } + test("SPARK-59146: pipe SET retains qualified source columns") { + checkResolverGuard( + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 |> SELECT a, t.a, t.b") + checkResolverGuard( + "VALUES (1, 2, 3) AS t(a, b, c) |> SET b = 20 |> SELECT t.*") + val repeatedSourceQuery = + "VALUES (1, 2) AS s(a, b) |> SELECT a, a, b " + + "|> AS t |> SET b = 3 |> SELECT t.*" + checkResolverGuard(repeatedSourceQuery) + val sequentialAssignmentsQuery = + "VALUES (1, 10) AS t(a, b) " + + "|> SET a = a + 1, b = b + 1 |> SELECT a, b, t.*" + checkResolverGuard(sequentialAssignmentsQuery) + checkResolverGuard( + "SELECT 1 AS x, NAMED_STRUCT('x', 2) AS col " + + "|> AS col |> SET x = x + 1 |> SELECT x, col.x") + + withSQLConf( + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED.key -> "true", + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false", + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "false") { + // Dataset collection analyzes a DeserializeToObject, which is not supported by the + // single-pass resolver, so execute the already-resolved physical plan directly. + val qualifiedColumnRows = sql( + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 |> SELECT a, t.a, t.b" + ).queryExecution.executedPlan.executeCollectPublic() + assert(qualifiedColumnRows.toSeq === Seq(Row(2, 1, 10))) + + val qualifiedStarRows = sql( + "VALUES (1, 2, 3) AS t(a, b, c) |> SET b = 20 |> SELECT t.*" + ).queryExecution.executedPlan.executeCollectPublic() + assert(qualifiedStarRows.toSeq === Seq(Row(1, 2, 3))) + + val repeatedSource = sql(repeatedSourceQuery) + assert(repeatedSource.schema.fieldNames === Array("a", "a", "b")) + val repeatedSourceRows = repeatedSource.queryExecution.executedPlan.executeCollectPublic() + assert(repeatedSourceRows.toSeq === Seq(Row(1, 1, 2))) + + val sequentialAssignments = sql(sequentialAssignmentsQuery) + val sequentialAssignmentRows = + sequentialAssignments.queryExecution.executedPlan.executeCollectPublic() + assert(sequentialAssignmentRows.toSeq === Seq(Row(2, 11, 1, 10))) + } + + val aliasQuery = "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 |> AS u" + checkResolverGuard(aliasQuery) + withSQLConf( + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "true", + SQLConf.ANALYZER_DUAL_RUN_SAMPLE_RATE.key -> "1.0", + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false") { + assert(sql(aliasQuery).schema.fieldNames === Array("a", "b")) + } + } + + test("SPARK-59146: retained columns follow outer join nullability") { + val outerJoinQueries = Seq( + "VALUES (1) AS t(a) |> SET a = a + 1 " + + "|> RIGHT OUTER JOIN VALUES (3) AS u(c) ON a = c |> SELECT t.a" -> + Set(Row(null)), + "VALUES (1) AS t(a) |> SET a = a + 1 " + + "|> RIGHT OUTER JOIN VALUES (3) AS u(a) USING (a) |> SELECT t.a" -> + Set(Row(null)), + "VALUES (1) AS t(a) |> SET a = a + 1 " + + "|> FULL OUTER JOIN VALUES (3) AS u(c) ON a = c |> SELECT t.a" -> + Set(Row(1), Row(null)), + "VALUES (1) AS t(a) |> SET a = a + 1 " + + "|> FULL OUTER JOIN VALUES (3) AS u(a) USING (a) |> SELECT t.a" -> + Set(Row(1), Row(null))) + outerJoinQueries.foreach { case (query, _) => checkResolverGuard(query) } + + Seq(false, true).foreach { singlePassEnabled => + withSQLConf( + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED.key -> singlePassEnabled.toString, + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false", + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "false") { + outerJoinQueries.foreach { case (query, expectedRows) => + val result = sql(query) + assert(result.schema.fieldNames === Array("a")) + assert(result.schema("a").nullable) + val rows = result.queryExecution.executedPlan.executeCollectPublic() + assert(rows.toSet === expectedRows) + } + } + } + } + + test("SPARK-59146: missing attributes cross pipe SET input in dual-run analysis") { + withSQLConf( + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "true", + SQLConf.ANALYZER_DUAL_RUN_SAMPLE_RATE.key -> "1.0", + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false") { + val dataFrame = sql( + "SELECT x FROM VALUES (1, 2), (3, 0) AS t(x, y) |> SET x = x + 1") + assert(dataFrame.orderBy("y").collect().map(_.getInt(0)) === Array(4, 2)) + } + } + + test("SPARK-59146: retained columns do not change terminal operator output") { + val terminalFilterQuery = + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 |> WHERE t.a = 1" + val terminalFilter = sql(terminalFilterQuery) + assert(terminalFilter.schema.fieldNames === Array("a", "b")) + assert(terminalFilter.collect().toSeq === Seq(Row(2, 10))) + + // Pipe WHERE is unsupported in single-pass, so isolate its supported Filter subtree. + val parsedFilter = spark.sessionState.sqlParser.parsePlan(terminalFilterQuery) + .asInstanceOf[Filter] + val filterWithoutPipeBoundary = parsedFilter.copy( + child = parsedFilter.child.asInstanceOf[PipeOperator].child) + val resolver = new Resolver( + catalogManager = spark.sessionState.catalogManager, + extensions = spark.sessionState.analyzer.singlePassResolverExtensions, + metadataResolverExtensions = + spark.sessionState.analyzer.singlePassMetadataResolverExtensions) + val resolvedFilter = resolver.lookupMetadataAndResolve(filterWithoutPipeBoundary) + assert(resolvedFilter.schema.fieldNames === Array("a", "b")) + val resolvedFilterRows = + spark.sessionState.executePlan(resolvedFilter).executedPlan.executeCollectPublic() + assert(resolvedFilterRows.toSeq === Seq(Row(2, 10))) + + val terminalOrderByQuery = + "VALUES (2, 20), (1, 10) AS t(a, b) |> SET a = -a ORDER BY t.a" + val terminalSortByQuery = + "VALUES (2, 20), (1, 10) AS t(a, b) |> SET a = -a SORT BY t.a" + Seq(terminalOrderByQuery, terminalSortByQuery).foreach(query => checkResolverGuard(query)) + val terminalJoinQuery = + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 " + + "|> INNER JOIN VALUES (1, 20) AS u(c, d) ON t.a = u.c" + val terminalSemiJoinQuery = + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 " + + "|> LEFT SEMI JOIN VALUES (2) AS u(c) ON a = u.c" + val terminalBoundaryCases = Seq( + (terminalJoinQuery, Array("a", "b", "c", "d"), Seq(Row(2, 10, 1, 20))), + (s"$terminalJoinQuery |> AS v", Array("a", "b", "c", "d"), Seq(Row(2, 10, 1, 20))), + (s"$terminalJoinQuery LIMIT 1", Array("a", "b", "c", "d"), Seq(Row(2, 10, 1, 20))), + (s"$terminalJoinQuery OFFSET 0", Array("a", "b", "c", "d"), Seq(Row(2, 10, 1, 20))), + (s"$terminalJoinQuery ORDER BY a", Array("a", "b", "c", "d"), Seq(Row(2, 10, 1, 20))), + ( + s"$terminalJoinQuery |> TABLESAMPLE (100 PERCENT)", + Array("a", "b", "c", "d"), + Seq(Row(2, 10, 1, 20)) + ), + ( + s"WITH x AS (SELECT 1) $terminalJoinQuery", + Array("a", "b", "c", "d"), + Seq(Row(2, 10, 1, 20)) + ), + ( + s"WITH x AS ($terminalJoinQuery) SELECT * FROM x", + Array("a", "b", "c", "d"), + Seq(Row(2, 10, 1, 20)) + ), + ( + "WITH x AS (VALUES (1, 10) AS t(a, b) |> SET a = a + 1) SELECT * FROM x", + Array("a", "b"), + Seq(Row(2, 10)) + ), + ( + s"SELECT * FROM ($terminalJoinQuery LIMIT 1) AS v", + Array("a", "b", "c", "d"), + Seq(Row(2, 10, 1, 20)) + ), + (s"$terminalSemiJoinQuery ORDER BY t.a", Array("a", "b"), Seq(Row(2, 10))), + ( + s"$terminalJoinQuery |> SELECT a, t.a, c", + Array("a", "a", "c"), + Seq(Row(2, 1, 1)) + ) + ) + terminalBoundaryCases.foreach { case (query, _, _) => + withClue(s"Query: $query") { + checkResolverGuard(query) + } + } + val limitSelectQuery = + s"$terminalJoinQuery |> LIMIT 1 |> SELECT a, t.a, c" + def limitSelectPlan: LogicalPlan = { + spark.sessionState.sqlParser.parsePlan(limitSelectQuery).transformDown { + case pipeOperator: PipeOperator => pipeOperator.child + } + } + checkResolverGuard(limitSelectPlan, unsupportedReason = None) + withSQLConf( + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "true", + SQLConf.ANALYZER_DUAL_RUN_SAMPLE_RATE.key -> "1.0", + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false") { + terminalBoundaryCases.foreach { case (query, _, _) => + withClue(s"Query: $query") { + assert(sql(query).queryExecution.analyzed.resolved) + } + } + val limitSelectResult = org.apache.spark.sql.classic.Dataset.ofRows(spark, limitSelectPlan) + assert(limitSelectResult.queryExecution.analyzed.resolved) + } + val aggregateOrderByQuery = + "SELECT a, 1, sum(b) FROM VALUES (1, 2) AS v(a, b) GROUP BY 1, 2 ORDER BY 1" + withSQLConf( + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "true", + SQLConf.ANALYZER_DUAL_RUN_SAMPLE_RATE.key -> "1.0", + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false") { + assert(sql(aggregateOrderByQuery).collect().toSeq === Seq(Row(1, 1, 2))) + } + Seq(false, true).foreach { singlePassEnabled => + withSQLConf( + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED.key -> singlePassEnabled.toString, + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false", + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "false") { + val orderedResult = sql(terminalOrderByQuery) + assert(orderedResult.schema.fieldNames === Array("a", "b")) + val orderedRows = orderedResult.queryExecution.executedPlan.executeCollectPublic() + assert(orderedRows.toSeq === Seq(Row(-1, 10), Row(-2, 20))) + + val sortedResult = sql(terminalSortByQuery) + assert(sortedResult.schema.fieldNames === Array("a", "b")) + val sortedRows = sortedResult.queryExecution.executedPlan.executeCollectPublic() + assert(sortedRows.toSet === Set(Row(-1, 10), Row(-2, 20))) + + terminalBoundaryCases.foreach { case (query, expectedFields, expectedRows) => + val joinResult = sql(query) + assert(joinResult.schema.fieldNames === expectedFields) + val joinRows = joinResult.queryExecution.executedPlan.executeCollectPublic() + assert(joinRows.toSeq === expectedRows) + } + + val limitSelectResult = org.apache.spark.sql.classic.Dataset.ofRows(spark, limitSelectPlan) + assert(limitSelectResult.schema.fieldNames === Array("a", "a", "c")) + assert(limitSelectResult.collect().toSeq === Seq(Row(2, 1, 1))) + } + } + } + + test("SPARK-59146: pipe SET metadata does not cross Dataset boundaries") { + Seq(false, true).foreach { singlePassEnabled => + withSQLConf( + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED.key -> singlePassEnabled.toString, + SQLConf.ANALYZER_SINGLE_PASS_RESOLVER_ENABLED_TENTATIVELY.key -> "false", + SQLConf.ANALYZER_DUAL_RUN_LEGACY_AND_SINGLE_PASS_RESOLVER.key -> "false") { + val dataFrame = sql("VALUES (1, 10) AS t(a, b) |> SET a = a + 1") + assert(!dataFrame.queryExecution.analyzed.exists(_.isInstanceOf[PipeSetInput])) + + val error = intercept[SparkThrowable] { + dataFrame("t.a") + } + assert(error.getCondition === "UNRESOLVED_COLUMN.WITH_SUGGESTION") + + val joinedDataFrame = sql( + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 " + + "|> INNER JOIN VALUES (2, 20) AS u(a, c) USING (a)") + val analyzedJoin = joinedDataFrame.queryExecution.analyzed + assert(!analyzedJoin.exists(_.isInstanceOf[PipeSetInput])) + assert(!analyzedJoin.exists { + case project: Project => + project.getTagValue(Project.hiddenOutputTag).exists(_.exists(_.pipeSetRetained)) + case _ => false + }) + + val joinError = intercept[SparkThrowable] { + joinedDataFrame.select("t.a").queryExecution.analyzed + } + assert(joinError.getCondition === "UNRESOLVED_COLUMN.WITH_SUGGESTION") + assert(analyzedJoin.metadataOutput.exists { attribute => + attribute.name == "a" && attribute.qualifier == Seq("u") && !attribute.pipeSetRetained + }) + } + } + } + + test("SPARK-59146: clean pipe SET input before finalizing analysis-only commands") { + val parsedPlan = spark.sessionState.sqlParser.parsePlan( + "CREATE TEMPORARY VIEW pipe_set_view AS " + + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 " + + "|> INNER JOIN VALUES (2, 20) AS u(a, c) USING (a)") + val analyzedCommand = spark.sessionState.analyzer.execute(parsedPlan) + + assert(analyzedCommand.children.isEmpty) + assert(analyzedCommand.innerChildren.length === 1) + val analyzedViewPlan = analyzedCommand.innerChildren.head.asInstanceOf[LogicalPlan] + assert(!analyzedViewPlan.exists(_.isInstanceOf[PipeSetInput])) + assert(!analyzedViewPlan.exists { + case project: Project => + project.getTagValue(Project.hiddenOutputTag).exists(_.exists(_.pipeSetRetained)) + case _ => false + }) + } + + test("SPARK-59146: dropDuplicates discards hidden pipe SET values") { + val dataFrame = sql("VALUES (1), (2) AS t(a) |> SET a = 0") + val error = intercept[SparkThrowable] { + dataFrame.dropDuplicates().select("t.a").queryExecution.analyzed + } + assert(error.getCondition === "UNRESOLVED_COLUMN.WITH_SUGGESTION") + + val selectedSource = sql( + "VALUES (1), (2) AS t(a) |> SET a = 0 |> SELECT a, t.a AS source_a") + assert(selectedSource.dropDuplicates().select("source_a").collect().map(_.getInt(0)).sorted === + Array(1, 2)) + } + test("Binary arithmetic") { checkResolverGuard("SELECT col1+col2 FROM VALUES(1,2)") checkResolverGuard("SELECT 1 + 2.3 / 2 - 3 DIV 2 + 3.0 * 10.0") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverSuite.scala index 7b628357bc538..cdd32a565d375 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/analysis/resolver/ResolverSuite.scala @@ -24,9 +24,11 @@ import org.apache.spark.sql.catalyst.analysis.resolver.{ Resolver, ResolverExtension } -import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference} +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, PipeSetInput} import org.apache.spark.sql.catalyst.plans.NormalizePlan import org.apache.spark.sql.catalyst.plans.logical.{LeafNode, LogicalPlan, Project} +import org.apache.spark.sql.catalyst.rules.Rule +import org.apache.spark.sql.catalyst.util._ import org.apache.spark.sql.test.SharedSparkSession import org.apache.spark.sql.types.IntegerType @@ -107,6 +109,40 @@ class ResolverSuite extends SharedSparkSession { ) } + test("SPARK-59146: pipe SET cleanup follows extended rewrite rules") { + var extendedRuleSawPipeSetInput = false + val extendedRewriteRule = new Rule[LogicalPlan] { + override def apply(plan: LogicalPlan): LogicalPlan = { + if (plan.exists(_.isInstanceOf[PipeSetInput])) { + extendedRuleSawPipeSetInput = true + } + plan.transformUp { + case pipeSetInput: PipeSetInput => pipeSetInput.child + } + } + } + val resolver = new Resolver( + catalogManager = spark.sessionState.catalogManager, + sharedRelationCache = spark.sharedState.relationCache, + extendedRewriteRules = Seq(extendedRewriteRule)) + + val result = resolver.lookupMetadataAndResolve(spark.sessionState.sqlParser.parsePlan( + "VALUES (1, 10) AS t(a, b) |> SET a = a + 1 " + + "|> INNER JOIN VALUES (2, 20) AS u(a, c) USING (a) " + + "|> SELECT a")) + val hiddenOutput = result.collect { + case project: Project => + project.getTagValue(Project.hiddenOutputTag).toSeq.flatten + }.flatten + + assert(extendedRuleSawPipeSetInput) + assert(!result.exists(_.isInstanceOf[PipeSetInput])) + assert(!hiddenOutput.exists(_.pipeSetRetained)) + assert(hiddenOutput.exists { attribute => + attribute.name == "a" && attribute.qualifier == Seq("u") + }) + } + private def createResolver(extensions: Seq[ResolverExtension] = Seq.empty): Resolver = { new Resolver( spark.sessionState.catalogManager,