Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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._`.
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 =>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
)
)
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -417,7 +417,7 @@ class NameScope(
*/
def isStarQualifiedByTable(unresolvedStar: UnresolvedStarBase): Boolean = {
unresolvedStar.isQualifiedByTable(
childOperatorOutput = output,
childOperatorOutput = output ++ hiddenOutput.filter(_.qualifiedAccessOnly),
resolver = nameComparator
)
}
Expand Down Expand Up @@ -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
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

/**
Expand Down Expand Up @@ -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,
Expand All @@ -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)
}
Expand All @@ -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]]).
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 =>
Expand Down Expand Up @@ -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)
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
}
}
}
Expand Down Expand Up @@ -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 =>
Expand Down Expand Up @@ -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
Expand All @@ -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()
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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)
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 =>
Expand Down Expand Up @@ -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 =>
Expand Down
Loading