Skip to content
Merged
Show file tree
Hide file tree
Changes from 14 commits
Commits
Show all changes
16 commits
Select commit Hold shift + click to select a range
24cce80
fix(plan): preserve empty correlated aggregate projections
VioletQwQ-0 Jul 31, 2026
9f33798
Merge branch 'main' into codex/issue-25959-empty-correlated-agg
mergify[bot] Jul 31, 2026
57e72e0
Merge branch 'main' into codex/issue-25959-empty-correlated-agg
mergify[bot] Jul 31, 2026
64f729e
Merge branch 'main' into codex/issue-25959-empty-correlated-agg
mergify[bot] Jul 31, 2026
92de3c0
Merge branch 'main' into codex/issue-25959-empty-correlated-agg
mergify[bot] Jul 31, 2026
bb3db5f
test: update correlated aggregate plan expectations
VioletQwQ-0 Jul 31, 2026
345f5d8
Merge remote-tracking branch 'upstream/main' into codex/issue-25959-e…
VioletQwQ-0 Jul 31, 2026
c1af1fd
fix(plan): finalize empty correlated aggregates after join
VioletQwQ-0 Aug 2, 2026
1b74e52
Merge remote-tracking branch 'upstream/main' into codex/issue-25959-e…
VioletQwQ-0 Aug 2, 2026
0db612e
Merge remote-tracking branch 'upstream/main' into codex/issue-25959-e…
VioletQwQ-0 Aug 3, 2026
30229cd
Merge remote-tracking branch 'upstream/main' into codex/issue-25959-e…
VioletQwQ-0 Aug 3, 2026
cd7396c
test: update associative plans for scalar aggregate rewrite
VioletQwQ-0 Aug 3, 2026
7b94a2d
Merge remote-tracking branch 'upstream/main' into codex/issue-25959-e…
VioletQwQ-0 Aug 3, 2026
fcefaa8
test: cover CTE correlated aggregate fallback
VioletQwQ-0 Aug 3, 2026
eb983b2
fix(plan): restore aggregate empty results after decorrelation
VioletQwQ-0 Aug 3, 2026
55203b0
Merge branch 'main' into codex/issue-25959-empty-correlated-agg
mergify[bot] Aug 3, 2026
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
154 changes: 149 additions & 5 deletions pkg/sql/plan/flatten_subquery.go
Original file line number Diff line number Diff line change
Expand Up @@ -158,11 +158,12 @@ func (builder *QueryBuilder) flattenSubquery(nodeID int32, subquery *plan.Subque

switch subquery.Typ {
case plan.SubqueryRef_SCALAR:
var rewrite bool
var rewriteCount bool

// Uncorrelated subquery
// Preserve the legacy COUNT fallback for plan shapes that cannot use the
// more precise empty-input projection reconstruction below.
if len(joinPreds) > 0 && builder.findAggrCount(subCtx.aggregates) {
rewrite = true
rewriteCount = true
}

if scalarExistential {
Expand Down Expand Up @@ -194,6 +195,12 @@ func (builder *QueryBuilder) flattenSubquery(nodeID int32, subquery *plan.Subque
joinType = plan.Node_LEFT
}

postJoinProjection, finalizeProjection, err :=
builder.prepareCorrelatedScalarAggregatePostJoinProjection(subID, subCtx, joinPreds)
if err != nil {
return nodeID, nil, err
}

nodeID = builder.appendNode(&plan.Node{
NodeType: plan.Node_JOIN,
Children: []int32{nodeID, subID},
Expand All @@ -211,7 +218,9 @@ func (builder *QueryBuilder) flattenSubquery(nodeID int32, subquery *plan.Subque
}

retExpr := scalarMatch
if retExpr == nil {
if finalizeProjection {
retExpr = postJoinProjection
} else if retExpr == nil {
retExpr = &plan.Expr{
Typ: subCtx.results[0].Typ,
Expr: &plan.Expr_Col{
Expand All @@ -232,7 +241,7 @@ func (builder *QueryBuilder) flattenSubquery(nodeID int32, subquery *plan.Subque
return 0, nil, err
}
}
if rewrite {
if !finalizeProjection && rewriteCount {
argsType := make([]types.Type, 1)
argsType[0] = makeTypeByPlan2Expr(retExpr)
fGet, err := function.GetFunctionByName(builder.GetContext(), "isnull", argsType)
Expand Down Expand Up @@ -597,6 +606,141 @@ func (builder *QueryBuilder) generateRowComparison(op string, child *plan.Expr,
}
}

// prepareCorrelatedScalarAggregatePostJoinProjection moves the scalar final
// expression above the LEFT JOIN used to decorrelate an implicit single-group
// aggregate. pullupThroughAgg groups the inner input by the correlation key, so
// a missing key produces no right row. Evaluating COALESCE, arithmetic, or CASE
// below the join would therefore skip the expression for that outer row.
//
// The right projection is rewritten to expose only raw aggregate outputs and
// the correlation keys that pullupThroughProj already appended. COUNT outputs
// are restored to zero after null extension; other supported aggregates keep
// NULL as their empty-input value. The saved final expression is then evaluated
// against those post-join values.
//
// This is intentionally limited to the ordinary PROJECT -> AGG shape. Wrappers
// that can remove or reorder the aggregate row (for example HAVING, DISTINCT,
// SORT, or LIMIT) keep the legacy path.
func (builder *QueryBuilder) prepareCorrelatedScalarAggregatePostJoinProjection(
subID int32,
subCtx *BindContext,
joinPreds []*plan.Expr,
) (*plan.Expr, bool, error) {
if !subCtx.hasSingleRow || len(subCtx.groups) != 0 || len(subCtx.aggregates) == 0 || len(joinPreds) == 0 {
return nil, false, nil
}

project := builder.qry.Nodes[subID]
if project.NodeType != plan.Node_PROJECT || len(project.Children) != 1 || len(project.BindingTags) != 1 ||
len(project.ProjectList) == 0 || project.Limit != nil || project.Offset != nil || project.RankOption != nil {
return nil, false, nil
}

agg := builder.qry.Nodes[project.Children[0]]
if agg.NodeType != plan.Node_AGG || len(agg.BindingTags) < 2 || agg.BindingTags[1] != subCtx.aggregateTag ||
len(agg.AggList) != len(subCtx.aggregates) {
return nil, false, nil
}

projectTag := project.BindingTags[0]
projectedAggregates := make([]*plan.Expr, len(agg.AggList))
rawAggregates := make([]*plan.Expr, len(agg.AggList))
firstAppendedPos := int32(len(project.ProjectList))
for i, aggregate := range agg.AggList {
fn := aggregate.GetF()
if fn == nil || fn.Func == nil {
return nil, false, nil
}

projectPos := int32(0)
if i > 0 {
projectPos = firstAppendedPos + int32(i-1)
}
rawAggregates[i] = GetColExpr(aggregate.Typ, subCtx.aggregateTag, int32(i))
projected := GetColExpr(aggregate.Typ, projectTag, projectPos)
projected.Typ.NotNullable = false

switch fn.Func.ObjName {
case "sum", "avg", "min", "max", "json_arrayagg":
projectedAggregates[i] = projected
case "count", "starcount":
var err error
projectedAggregates[i], err = builder.restoreEmptyCount(projected, aggregate.Typ)
if err != nil {
return nil, false, err
}
default:
return nil, false, nil
}
}

postJoinProjection, ok := replaceAggregateRefsForPostJoin(
DeepCopyExpr(project.ProjectList[0]), subCtx.aggregateTag, projectedAggregates)
if !ok {
return nil, false, nil
}
postJoinProjection, stillCorrelated := decreaseDepth(postJoinProjection)
if stillCorrelated {
return nil, false, nil
}

newProjectList := make([]*plan.Expr, len(project.ProjectList), len(project.ProjectList)+len(rawAggregates)-1)
copy(newProjectList, project.ProjectList)
newProjectList[0] = rawAggregates[0]
newProjectList = append(newProjectList, rawAggregates[1:]...)
project.ProjectList = newProjectList
return postJoinProjection, true, nil
}

func (builder *QueryBuilder) restoreEmptyCount(countExpr *plan.Expr, aggregateType plan.Type) (*plan.Expr, error) {
isNullExpr, err := BindFuncExprImplByPlanExpr(builder.GetContext(), "isnull", []*plan.Expr{countExpr})
if err != nil {
return nil, err
}
zeroExpr := makePlan2Int64ConstExprWithType(0)
zeroExpr.Typ = aggregateType
zeroExpr.Typ.NotNullable = true
return BindFuncExprImplByPlanExpr(builder.GetContext(), "case", []*plan.Expr{
isNullExpr,
zeroExpr,
DeepCopyExpr(countExpr),
})
}

func replaceAggregateRefsForPostJoin(
expr *plan.Expr,
aggregateTag int32,
projectedAggregates []*plan.Expr,
) (*plan.Expr, bool) {
if expr == nil {
return nil, false
}

switch item := expr.Expr.(type) {
case *plan.Expr_Col:
if item.Col.RelPos != aggregateTag {
return nil, false
}
if item.Col.ColPos < 0 || int(item.Col.ColPos) >= len(projectedAggregates) {
return nil, false
}
return DeepCopyExpr(projectedAggregates[item.Col.ColPos]), true
case *plan.Expr_F:
for i, arg := range item.F.Args {
var ok bool
item.F.Args[i], ok = replaceAggregateRefsForPostJoin(arg, aggregateTag, projectedAggregates)
if !ok {
return nil, false
}
}
return expr, true
case *plan.Expr_List, *plan.Expr_W, *plan.Expr_Sub:
return nil, false
default:
return expr, true
}
}

func (builder *QueryBuilder) findAggrCount(aggrs []*plan.Expr) bool {
for _, aggr := range aggrs {
switch exprImpl := aggr.Expr.(type) {
Expand Down
Loading
Loading