From 0bd17d6185bd0b7ec723e913ad50faf618d21ae2 Mon Sep 17 00:00:00 2001 From: pokeeffe-molecula <85502298+pokeeffe-molecula@users.noreply.github.com> Date: Tue, 22 Nov 2022 19:21:47 -0600 Subject: [PATCH] fixed top(x) where top cannot be pushed down into pql query (#2311) * fixed top(x) where top cannot be pushed down into pql query * review feedback (cherry picked from commit d3f10be743bffbe8c44e63fea151fdbd31f223e6) --- etcd/embed.go | 3 + sql3/planner/compileselect.go | 5 +- sql3/planner/expressionagg.go | 56 ++++-- sql3/planner/oppqlgroupby.go | 2 +- sql3/planner/oppqlmultiaggregate.go | 2 +- sql3/planner/oppqlmultigroupby.go | 16 +- sql3/planner/opprojection.go | 11 ++ sql3/planner/optop.go | 8 +- sql3/planner/planoptimizer.go | 267 ++++++++++------------------ sql3/planner/planwalker.go | 136 +++----------- sql3/test/defs/defs.go | 2 + sql3/test/defs/defs_top.go | 55 ++++++ 12 files changed, 254 insertions(+), 309 deletions(-) create mode 100644 sql3/test/defs/defs_top.go diff --git a/etcd/embed.go b/etcd/embed.go index 29a08913f..72ad3e2c6 100644 --- a/etcd/embed.go +++ b/etcd/embed.go @@ -229,6 +229,9 @@ func NewEtcd(opt Options, logger logger.Logger, replicas int, version string) *E if e.options.HeartbeatTTL == 0 { e.options.HeartbeatTTL = 5 // seconds + // !! DEBUG + // e.options.HeartbeatTTL = 3600 // seconds + // !! DEBUG } return e } diff --git a/sql3/planner/compileselect.go b/sql3/planner/compileselect.go index 36cdf7484..0554531d9 100644 --- a/sql3/planner/compileselect.go +++ b/sql3/planner/compileselect.go @@ -77,7 +77,10 @@ func (p *ExecutionPlanner) compileSelectStatement(stmt *parser.SelectStatement, for _, expr := range projections { InspectExpression(expr, func(expr types.PlanExpression) bool { switch ex := expr.(type) { - case *sumPlanExpression: + case *sumPlanExpression, *countPlanExpression, *countDistinctPlanExpression, + *avgPlanExpression, *minPlanExpression, *maxPlanExpression, + *percentilePlanExpression: + //return false for these, because thats as far down we want to inspect return false case *qualifiedRefPlanExpression: nonAggregateReferences = append(nonAggregateReferences, ex) diff --git a/sql3/planner/expressionagg.go b/sql3/planner/expressionagg.go index d4bf9899d..744a6e2f0 100644 --- a/sql3/planner/expressionagg.go +++ b/sql3/planner/expressionagg.go @@ -133,11 +133,16 @@ func (n *countPlanExpression) Plan() map[string]interface{} { } func (n *countPlanExpression) Children() []types.PlanExpression { - return []types.PlanExpression{} + return []types.PlanExpression{ + n.arg, + } } func (n *countPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { - return n, nil + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newCountPlanExpression(children[0], n.returnDataType), nil } // countDistinctPlanExpression handles COUNT(DISTINCT) @@ -196,11 +201,16 @@ func (n *countDistinctPlanExpression) Plan() map[string]interface{} { } func (n *countDistinctPlanExpression) Children() []types.PlanExpression { - return nil + return []types.PlanExpression{ + n.arg, + } } func (n *countDistinctPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { - return n, nil + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newCountDistinctPlanExpression(children[0], n.returnDataType), nil } // aggregator for the SUM function @@ -436,11 +446,16 @@ func (n *avgPlanExpression) Plan() map[string]interface{} { } func (n *avgPlanExpression) Children() []types.PlanExpression { - return nil + return []types.PlanExpression{ + n.arg, + } } func (n *avgPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { - return n, nil + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newAvgPlanExpression(children[0], n.returnDataType), nil } // aggregator for MIN @@ -533,11 +548,16 @@ func (n *minPlanExpression) Plan() map[string]interface{} { } func (n *minPlanExpression) Children() []types.PlanExpression { - return nil + return []types.PlanExpression{ + n.arg, + } } func (n *minPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { - return n, nil + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newMinPlanExpression(children[0], n.returnDataType), nil } // aggregator for MAX @@ -630,11 +650,16 @@ func (n *maxPlanExpression) Plan() map[string]interface{} { } func (n *maxPlanExpression) Children() []types.PlanExpression { - return nil + return []types.PlanExpression{ + n.arg, + } } func (n *maxPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { - return n, nil + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newMaxPlanExpression(children[0], n.returnDataType), nil } // percentilePlanExpression handles PERCENTILE() @@ -693,15 +718,22 @@ func (n *percentilePlanExpression) Plan() map[string]interface{} { result["_expr"] = fmt.Sprintf("%T", n) result["dataType"] = n.Type().TypeName() result["arg"] = n.arg.Plan() + result["ntharg"] = n.nthArg.Plan() return result } func (n *percentilePlanExpression) Children() []types.PlanExpression { - return nil + return []types.PlanExpression{ + n.arg, + n.nthArg, + } } func (n *percentilePlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { - return n, nil + if len(children) != 2 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return newPercentilePlanExpression(children[0], children[1], n.returnDataType), nil } // aggregator for last diff --git a/sql3/planner/oppqlgroupby.go b/sql3/planner/oppqlgroupby.go index dcd7ed708..679727cdb 100644 --- a/sql3/planner/oppqlgroupby.go +++ b/sql3/planner/oppqlgroupby.go @@ -84,7 +84,7 @@ func (p *PlanOpPQLGroupBy) Schema() types.Schema { result[idx] = s } s := &types.PlannerColumn{ - ColumnName: "", + ColumnName: p.aggregate.String(), RelationName: "", Type: p.aggregate.AggExpression().Type(), } diff --git a/sql3/planner/oppqlmultiaggregate.go b/sql3/planner/oppqlmultiaggregate.go index f6a86e356..71d7c9ba1 100644 --- a/sql3/planner/oppqlmultiaggregate.go +++ b/sql3/planner/oppqlmultiaggregate.go @@ -57,7 +57,7 @@ func (p *PlanOpPQLMultiAggregate) Schema() types.Schema { result := make(types.Schema, len(p.operators)) for idx, aggOp := range p.operators { s := &types.PlannerColumn{ - ColumnName: "", + ColumnName: aggOp.aggregate.String(), RelationName: "", Type: aggOp.aggregate.AggExpression().Type(), } diff --git a/sql3/planner/oppqlmultigroupby.go b/sql3/planner/oppqlmultigroupby.go index e3198cd2f..47d71cec2 100644 --- a/sql3/planner/oppqlmultigroupby.go +++ b/sql3/planner/oppqlmultigroupby.go @@ -6,7 +6,8 @@ import ( "context" "fmt" - "github.com/featurebasedb/featurebase/v3/sql3/planner/types" + "github.com/molecula/featurebase/v3/sql3" + "github.com/molecula/featurebase/v3/sql3/planner/types" ) // PlanOpPQLMultiGroupBy plan operator handles executing multiple 'sibling' pql group by queries @@ -80,7 +81,7 @@ func (p *PlanOpPQLMultiGroupBy) Schema() types.Schema { offset := len(p.groupByExprs) for idx, aggOp := range p.operators { s := &types.PlannerColumn{ - ColumnName: "", + ColumnName: aggOp.aggregate.String(), RelationName: "", Type: aggOp.aggregate.AggExpression().Type(), } @@ -116,6 +117,17 @@ func (p *PlanOpPQLMultiGroupBy) WithChildren(children ...types.PlanOperator) (ty return nil, nil } +func (p *PlanOpPQLMultiGroupBy) Expressions() []types.PlanExpression { + return p.groupByExprs +} + +func (p *PlanOpPQLMultiGroupBy) WithUpdatedExpressions(exprs ...types.PlanExpression) (types.PlanOperator, error) { + if len(exprs) != len(p.groupByExprs) { + return nil, sql3.NewErrInternalf("unexpected number of exprs '%d'", len(exprs)) + } + return NewPlanOpPQLMultiGroupBy(p.planner, p.operators, exprs), nil +} + // pqlMultiGroupByRowIter is an iterator for the PlanOpPQLMultiGroupBy operator // it provides rows consisting of the group by columns in the order they // were specified and lastly the aggregates in the order they were specified diff --git a/sql3/planner/opprojection.go b/sql3/planner/opprojection.go index 2c0dbe3e6..5ed6e1908 100644 --- a/sql3/planner/opprojection.go +++ b/sql3/planner/opprojection.go @@ -95,6 +95,17 @@ func (p *PlanOpProjection) Warnings() []string { return w } +func (p *PlanOpProjection) Expressions() []types.PlanExpression { + return p.Projections +} + +func (p *PlanOpProjection) WithUpdatedExpressions(exprs ...types.PlanExpression) (types.PlanOperator, error) { + if len(exprs) != len(p.Projections) { + return nil, sql3.NewErrInternalf("unexpected number of exprs '%d'", len(exprs)) + } + return NewPlanOpProjection(exprs, p.ChildOp), nil +} + func ExpressionToColumn(e types.PlanExpression) *types.PlannerColumn { var name string if n, ok := e.(types.IdentifiableByName); ok { diff --git a/sql3/planner/optop.go b/sql3/planner/optop.go index 95ea0badb..b2fe539fe 100644 --- a/sql3/planner/optop.go +++ b/sql3/planner/optop.go @@ -6,7 +6,8 @@ import ( "context" "fmt" - "github.com/featurebasedb/featurebase/v3/sql3/planner/types" + "github.com/molecula/featurebase/v3/sql3" + "github.com/molecula/featurebase/v3/sql3/planner/types" ) // PlanOpTop implements the TOP operator @@ -39,7 +40,10 @@ func (p *PlanOpTop) Children() []types.PlanOperator { } func (p *PlanOpTop) WithChildren(children ...types.PlanOperator) (types.PlanOperator, error) { - return nil, nil + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + return NewPlanOpTop(p.expr, children[0]), nil } func (p *PlanOpTop) Plan() map[string]interface{} { diff --git a/sql3/planner/planoptimizer.go b/sql3/planner/planoptimizer.go index b5b3df538..537b0626c 100644 --- a/sql3/planner/planoptimizer.go +++ b/sql3/planner/planoptimizer.go @@ -78,7 +78,7 @@ func (p *ExecutionPlanner) optimizePlan(ctx context.Context, plan types.PlanOper // log.Println("================================================================================") // log.Println("plan ppst-optimzation") - // jplan = plan.Plan() + // jplan = result.Plan() // a, _ = json.MarshalIndent(jplan, "", " ") // log.Println(string(a)) // log.Println("--------------------------------------------------------------------------------") @@ -408,7 +408,7 @@ func pushdownFilters(ctx context.Context, a *ExecutionPlanner, n types.PlanOpera filtersByTable := getFiltersByRelation(n) filters := newFilterSet(n.Predicate, filtersByTable, tableAliases) - // first push down filters to any op that implements FilteredRelation + // first push down filters to any op that supports a filter node, sameA, err := pushdownFiltersForFilterableRelations(n, filters) if err != nil { return nil, true, err @@ -723,14 +723,14 @@ func areAggregablesEqual(lhs types.Aggregable, rhs types.Aggregable) bool { // fixes references for a projection op depending on child func fixProjectionReferences(ctx context.Context, a *ExecutionPlanner, n types.PlanOperator, scope *OptimizerScope) (types.PlanOperator, bool, error) { return TransformPlanOp(n, func(node types.PlanOperator) (types.PlanOperator, bool, error) { - switch n := node.(type) { + switch thisNode := node.(type) { case *PlanOpProjection: - switch childOp := n.ChildOp.(type) { + switch childOp := thisNode.ChildOp.(type) { - case *PlanOpGroupBy: - //PlanOpGroupBy's iterator returns group by exprs, then aggregates in the order they appear + case *PlanOpGroupBy, *PlanOpPQLGroupBy, *PlanOpPQLMultiAggregate, *PlanOpPQLMultiGroupBy: childSchema := childOp.Schema() - for idx, pj := range n.Projections { + + for idx, pj := range thisNode.Projections { expr, _, err := TransformExpr(pj, func(e types.PlanExpression) (types.PlanExpression, bool, error) { switch thisAggregate := e.(type) { case types.Aggregable: @@ -744,168 +744,57 @@ func fixProjectionReferences(ctx context.Context, a *ExecutionPlanner, n types.P } } return nil, true, sql3.NewErrColumnNotFound(0, 0, thisAggregate.String()) - default: - return e, true, nil - } - }) - if err != nil { - return n, true, err - } - n.Projections[idx] = expr - } - return n, false, nil - case *PlanOpPQLGroupBy: - // PlanOpGroupBy's iterator returns group by exprs, then the single aggregate in the order they appear - - // make a map of the names of the group by columns - groupByColumnsNameMap := make(map[string]int) - for gidx, gbe := range childOp.groupByExprs { - gbeRef, ok := gbe.(*qualifiedRefPlanExpression) - if !ok { - return nil, false, sql3.NewErrInternalf("unexpected group by expression type '%T'", gbe) - } - gbeRef.columnIndex = gidx - groupByColumnsNameMap[gbeRef.columnName] = gidx - } - //set the index for the aggregate to be the length of the group by list - aggregateIndex := len(childOp.groupByExprs) - - //loop projections: - //1. looking for the Aggregable and replace it with a qualifiedRefPlanExpression pointing to - // the offset in the child iterator - //2. looking for the qualified refs replace it with a qualifiedRefPlanExpression pointing to - // the offset in the child iterator - for idx, pj := range n.Projections { - expr, _, err := TransformExpr(pj, func(e types.PlanExpression) (types.PlanExpression, bool, error) { - switch thisExpr := e.(type) { - case types.Aggregable: - ae := newQualifiedRefPlanExpression(fmt.Sprintf("$PlanOpPQLGroupBy:%d", aggregateIndex), "", aggregateIndex, e.Type()) - return ae, false, nil case *qualifiedRefPlanExpression: - colIdx, ok := groupByColumnsNameMap[thisExpr.columnName] - if ok { - ae := newQualifiedRefPlanExpression(fmt.Sprintf("$PlanOpPQLGroupBy.%s:%d", thisExpr.columnName, colIdx), thisExpr.columnName, colIdx, e.Type()) - return ae, false, nil - } - return e, true, nil - default: - return e, true, nil - } - }) - if err != nil { - return n, true, err - } - n.Projections[idx] = expr - } - return n, false, nil - - case *PlanOpPQLMultiAggregate: - //PlanOpGroupBy's iterator returns aggregates in the order they appear - - // make a list of the aggregables - aggregableList := make([]types.Aggregable, 0) - for _, op := range childOp.operators { - aggregableList = append(aggregableList, op.aggregate) - } - - //loop projections: - //1. looking for the Aggregable and replace it with a qualifiedRefPlanExpression pointing to - // the offset in the child iterator - for idx, pj := range n.Projections { - expr, _, err := TransformExpr(pj, func(e types.PlanExpression) (types.PlanExpression, bool, error) { - switch thisExpr := e.(type) { - case types.Aggregable: - for idx, a := range aggregableList { - if areAggregablesEqual(thisExpr, a) { - ae := newQualifiedRefPlanExpression(fmt.Sprintf("$PlanOpPQLMultiAggregate:%d", idx), "", idx, e.Type()) - return ae, false, nil + for idx, sc := range childSchema { + if matchesSchema(thisAggregate, sc) { + if idx != thisAggregate.columnIndex { + // update the column index + return newQualifiedRefPlanExpression(thisAggregate.tableName, thisAggregate.columnName, idx, thisAggregate.dataType), false, nil + } + return thisAggregate, true, nil } } - return e, true, nil + return nil, true, sql3.NewErrColumnNotFound(0, 0, thisAggregate.String()) + default: return e, true, nil } - }) - if err != nil { - return n, true, err - } - n.Projections[idx] = expr - } - return n, false, nil - - case *PlanOpPQLMultiGroupBy: - //PlanOpPQLMultiGroupBy's iterator returns group by exprs, then the aggregates in the order they appear - - // make a map of the names of the group by columns and update the indexes - groupByColumnsNameMap := make(map[string]int) - for gidx, gbe := range childOp.groupByExprs { - gbeRef, ok := gbe.(*qualifiedRefPlanExpression) - if !ok { - return nil, false, sql3.NewErrInternalf("unexpected group by expression type '%T'", gbe) - } - gbeRef.columnIndex = gidx - groupByColumnsNameMap[gbeRef.columnName] = gidx - } - //set the index for the start of the aggregates to be the length of the group by list - aggregateStartIndex := len(childOp.groupByExprs) - - // make a list of the aggregables - aggregableList := make([]types.Aggregable, 0) - for _, op := range childOp.operators { - aggregableList = append(aggregableList, op.aggregate) - } - - //loop projections: - //1. looking for the Aggregable and replace it with a qualifiedRefPlanExpression pointing to - // the offset in the child iterator - //2. looking for the qualified refs replace it with a qualifiedRefPlanExpression pointing to - // the offset in the child iterator - for idx, pj := range n.Projections { - expr, _, err := TransformExpr(pj, func(e types.PlanExpression) (types.PlanExpression, bool, error) { - switch thisExpr := e.(type) { + }, func(parentExpr, childExpr types.PlanExpression) bool { + // if the parent is an aggregable, and the child is a qualified ref + // we will skip, because the qualified ref should have already been handled in + // fixFieldRefs + switch parentExpr.(type) { case types.Aggregable: - for idx, a := range aggregableList { - if areAggregablesEqual(thisExpr, a) { - ae := newQualifiedRefPlanExpression(fmt.Sprintf("$PlanOpPQLMultiGroupBy:%d", idx+aggregateStartIndex), "", idx+aggregateStartIndex, e.Type()) - return ae, false, nil - } + switch childExpr.(type) { + case *qualifiedRefPlanExpression: + return false } - return e, true, nil - case *qualifiedRefPlanExpression: - colIdx, ok := groupByColumnsNameMap[thisExpr.columnName] - if ok { - ae := newQualifiedRefPlanExpression(fmt.Sprintf("$PlanOpPQLMultiGroupBy.%s:%d", thisExpr.columnName, colIdx), thisExpr.columnName, colIdx, e.Type()) - return ae, false, nil - } - return e, true, nil - - default: - return e, true, nil } + return true }) if err != nil { - return n, true, err + return thisNode, true, err } - n.Projections[idx] = expr + thisNode.Projections[idx] = expr } - return n, false, nil + return thisNode, false, nil // everything else that can be a child of projection case *PlanOpRelAlias, *PlanOpFilter, *PlanOpPQLTableScan, *PlanOpNestedLoops: - exprs, same, err := fixFieldRefIndexesOnExpressions(ctx, scope, a, childOp.Schema(), n.Projections...) + exprs, same, err := fixFieldRefIndexesOnExpressions(ctx, scope, a, childOp.Schema(), thisNode.Projections...) if err != nil { - return n, true, err + return thisNode, true, err } - n.Projections = exprs - return n, same, err + thisNode.Projections = exprs + return thisNode, same, err default: - return n, true, nil + return thisNode, true, nil } default: - return n, true, nil + return thisNode, true, nil } }) } @@ -914,6 +803,7 @@ func fixFieldRefs(ctx context.Context, a *ExecutionPlanner, n types.PlanOperator return TransformPlanOp(n, func(node types.PlanOperator) (types.PlanOperator, bool, error) { switch thisNode := node.(type) { case *PlanOpFilter: + // fix references for the expressions referenced in the filter predicate expression schema := thisNode.Schema() expressions := thisNode.Expressions() fixed, same, err := fixFieldRefIndexesOnExpressions(ctx, scope, a, schema, expressions...) @@ -927,6 +817,7 @@ func fixFieldRefs(ctx context.Context, a *ExecutionPlanner, n types.PlanOperator return newNode, same, nil case *PlanOpNestedLoops: + // fix references for the expressions referenced in the join condition expression schema := thisNode.Schema() expressions := thisNode.Expressions() fixed, same, err := fixFieldRefIndexesOnExpressions(ctx, scope, a, schema, expressions...) @@ -940,6 +831,7 @@ func fixFieldRefs(ctx context.Context, a *ExecutionPlanner, n types.PlanOperator return newNode, same, nil case *PlanOpGroupBy: + // fix references for the expressions referenced in the aggregate functions or the group by clause schema := thisNode.ChildOp.Schema() aggregateExpressions := thisNode.Aggregates fixedAggregateExpressions, aggregateSame, err := fixFieldRefIndexesOnExpressions(ctx, scope, a, schema, aggregateExpressions...) @@ -957,6 +849,27 @@ func fixFieldRefs(ctx context.Context, a *ExecutionPlanner, n types.PlanOperator newNode.warnings = append(newNode.warnings, thisNode.warnings...) return newNode, aggregateSame && groupBySame, nil + case *PlanOpPQLMultiGroupBy: + + schema := thisNode.operators[0].Schema() + for idx, op := range thisNode.operators { + if idx > 0 { + opSchema := op.Schema() + last := opSchema[len(opSchema)-1] + schema = append(schema, last) + } + } + expressions := thisNode.Expressions() + fixed, same, err := fixFieldRefIndexesOnExpressions(ctx, scope, a, schema, expressions...) + if err != nil { + return nil, true, err + } + newNode, err := thisNode.WithUpdatedExpressions(fixed...) + if err != nil { + return nil, true, err + } + return newNode, same, nil + default: return node, true, nil } @@ -1045,36 +958,6 @@ func getNestedLoopOperators(ctx context.Context, a *ExecutionPlanner, n types.Pl return joins } -func fixFieldRefIndexes(ctx context.Context, scope *OptimizerScope, a *ExecutionPlanner, schema types.Schema, exp types.PlanExpression) (types.PlanExpression, bool, error) { - return TransformExpr(exp, func(e types.PlanExpression) (types.PlanExpression, bool, error) { - switch typedExpr := e.(type) { - case *qualifiedRefPlanExpression: - for i, col := range schema { - newIndex := i - if strings.EqualFold(typedExpr.Name(), col.ColumnName) { - if len(typedExpr.tableName) > 0 { // do we have a qualifier? - if typedExpr.tableName == col.RelationName || typedExpr.tableName == col.AliasName { - if newIndex != typedExpr.columnIndex { - // update the column index - return newQualifiedRefPlanExpression(typedExpr.tableName, typedExpr.columnName, newIndex, typedExpr.dataType), false, nil - } - return e, true, nil - } - } else { // no qualifier - if newIndex != typedExpr.columnIndex { - // update the column index - return newQualifiedRefPlanExpression(typedExpr.tableName, typedExpr.columnName, newIndex, typedExpr.dataType), false, nil - } - return e, true, nil - } - } - } - return nil, true, sql3.NewErrColumnNotFound(0, 0, typedExpr.Name()) - } - return e, true, nil - }) -} - // for a list of expressions and an operator schema, fix the references for any qualifiedRef expressions func fixFieldRefIndexesOnExpressions(ctx context.Context, scope *OptimizerScope, a *ExecutionPlanner, schema types.Schema, expressions ...types.PlanExpression) ([]types.PlanExpression, bool, error) { var result []types.PlanExpression @@ -1100,3 +983,37 @@ func fixFieldRefIndexesOnExpressions(ctx context.Context, scope *OptimizerScope, } return expressions, true, nil } + +func matchesSchema(qualifiedRef *qualifiedRefPlanExpression, col *types.PlannerColumn) bool { + if strings.EqualFold(qualifiedRef.Name(), col.ColumnName) { + if len(qualifiedRef.tableName) == 0 { // do we have a qualifier? + return true + } + if qualifiedRef.tableName == col.RelationName || qualifiedRef.tableName == col.AliasName { + return true + } + } + return false +} + +func fixFieldRefIndexes(ctx context.Context, scope *OptimizerScope, a *ExecutionPlanner, schema types.Schema, exp types.PlanExpression) (types.PlanExpression, bool, error) { + return TransformExpr(exp, func(e types.PlanExpression) (types.PlanExpression, bool, error) { + switch typedExpr := e.(type) { + case *qualifiedRefPlanExpression: + for i, col := range schema { + newIndex := i + if matchesSchema(typedExpr, col) { + if newIndex != typedExpr.columnIndex { + // update the column index + return newQualifiedRefPlanExpression(typedExpr.tableName, typedExpr.columnName, newIndex, typedExpr.dataType), false, nil + } + return e, true, nil + } + } + return nil, true, sql3.NewErrColumnNotFound(0, 0, typedExpr.Name()) + } + return e, true, nil + }, func(parentExpr, childExpr types.PlanExpression) bool { + return true + }) +} diff --git a/sql3/planner/planwalker.go b/sql3/planner/planwalker.go index 554021595..bc71b82cb 100644 --- a/sql3/planner/planwalker.go +++ b/sql3/planner/planwalker.go @@ -276,64 +276,13 @@ type ExprWithPlanOpFunc func(op types.PlanOperator, expr types.PlanExpression) ( // whether the expression was modified, and an error or nil. type ExprFunc func(expr types.PlanExpression) (types.PlanExpression, bool, error) -// TransformPlanOpExprsWithPlanOp applies a transformation function to all expressions on the given plan operator from the bottom up in the context of the plan operator -func TransformPlanOpExprsWithPlanOp(op types.PlanOperator, f ExprWithPlanOpFunc) (types.PlanOperator, bool, error) { - return TransformPlanOp(op, func(n types.PlanOperator) (types.PlanOperator, bool, error) { - return TransformSinglePlanOpExprsInPlanOpContext(n, f) - }) -} - -// TransformPlanOpExprs applies a transformation function to all expressions on the given plan operator from the bottom up -func TransformPlanOpExprs(op types.PlanOperator, f ExprFunc) (types.PlanOperator, bool, error) { - return TransformPlanOpExprsWithPlanOp(op, func(operator types.PlanOperator, expr types.PlanExpression) (types.PlanExpression, bool, error) { - return f(expr) - }) -} - -// TransformSinglePlanOpExprsInPlanOpContext applies a transformation function to all expressions on a given plan operator in the context of that plan operator -func TransformSinglePlanOpExprsInPlanOpContext(op types.PlanOperator, f ExprWithPlanOpFunc) (types.PlanOperator, bool, error) { - ne, ok := op.(types.ContainsExpressions) - if !ok { - return op, true, nil - } - - exprs := ne.Expressions() - if len(exprs) == 0 { - return op, true, nil - } - - var ( - newExprs []types.PlanExpression - err error - ) - - for i := range exprs { - e := exprs[i] - e, same, err := TransformExprWithPlanOp(op, e, f) - if err != nil { - return nil, true, err - } - if !same { - if newExprs == nil { - newExprs = make([]types.PlanExpression, len(exprs)) - copy(newExprs, exprs) - } - newExprs[i] = e - } - } - - if len(newExprs) > 0 { - op, err = ne.WithUpdatedExpressions(newExprs...) - if err != nil { - return nil, true, err - } - return op, false, nil - } - return op, true, nil -} +// ExprExprSelectorFunc is a function that can be used as a filter selector during +// expression transformation - it is called before calling a transformation function on +// a child expression; if it returns false, the child is skipped +type ExprSelectorFunc func(parentExpr, childExpr types.PlanExpression) bool // TransformSinglePlanOpExpressions applies a transformation function to all expressions on the given plan operator -func TransformSinglePlanOpExpressions(op types.PlanOperator, f ExprFunc) (types.PlanOperator, bool, error) { +func TransformSinglePlanOpExpressions(op types.PlanOperator, f ExprFunc, selector ExprSelectorFunc) (types.PlanOperator, bool, error) { e, ok := op.(types.ContainsExpressions) if !ok { return op, true, nil @@ -347,7 +296,7 @@ func TransformSinglePlanOpExpressions(op types.PlanOperator, f ExprFunc) (types. var newExprs []types.PlanExpression for i := range exprs { expr := exprs[i] - expr, same, err := TransformExpr(expr, f) + expr, same, err := TransformExpr(expr, f, selector) if err != nil { return nil, true, err } @@ -370,12 +319,12 @@ func TransformSinglePlanOpExpressions(op types.PlanOperator, f ExprFunc) (types. } // TransformExpr applies a depth first transformation function to an expression -func TransformExpr(expr types.PlanExpression, f ExprFunc) (types.PlanExpression, bool, error) { +func TransformExpr(expr types.PlanExpression, transformFunction ExprFunc, selector ExprSelectorFunc) (types.PlanExpression, bool, error) { thisExpr := expr children := expr.Children() if len(children) == 0 { - return f(thisExpr) + return transformFunction(thisExpr) } var ( @@ -384,17 +333,19 @@ func TransformExpr(expr types.PlanExpression, f ExprFunc) (types.PlanExpression, ) for i := 0; i < len(children); i++ { - c := children[i] - c, same, err := TransformExpr(c, f) - if err != nil { - return nil, true, err - } - if !same { - if newChildren == nil { - newChildren = make([]types.PlanExpression, len(children)) - copy(newChildren, children) + child := children[i] + if selector(expr, child) { + c, same, err := TransformExpr(child, transformFunction, selector) + if err != nil { + return nil, true, err + } + if !same { + if newChildren == nil { + newChildren = make([]types.PlanExpression, len(children)) + copy(newChildren, children) + } + newChildren[i] = c } - newChildren[i] = c } } @@ -407,54 +358,9 @@ func TransformExpr(expr types.PlanExpression, f ExprFunc) (types.PlanExpression, } } - resultExpr, sameExpr, err := f(thisExpr) + resultExpr, sameExpr, err := transformFunction(thisExpr) if err != nil { return nil, true, err } return resultExpr, sameChildren && sameExpr, nil } - -// TransformExprWithPlanOp applies a depth first transformation function to an expression in the context of a plan operator -func TransformExprWithPlanOp(n types.PlanOperator, e types.PlanExpression, f ExprWithPlanOpFunc) (types.PlanExpression, bool, error) { - thisExpr := e - - children := thisExpr.Children() - if len(children) == 0 { - return f(n, e) - } - - var ( - newChildren []types.PlanExpression - err error - ) - - for i := 0; i < len(children); i++ { - c := children[i] - c, same, err := TransformExprWithPlanOp(n, c, f) - if err != nil { - return nil, true, err - } - if !same { - if newChildren == nil { - newChildren = make([]types.PlanExpression, len(children)) - copy(newChildren, children) - } - newChildren[i] = c - } - } - - sameChilren := true - if len(newChildren) > 0 { - sameChilren = false - thisExpr, err = thisExpr.WithChildren(newChildren...) - if err != nil { - return nil, true, err - } - } - - resultExpr, sameExpr, err := f(n, thisExpr) - if err != nil { - return nil, true, err - } - return resultExpr, sameChilren && sameExpr, nil -} diff --git a/sql3/test/defs/defs.go b/sql3/test/defs/defs.go index c06d3a592..f1510e9d0 100644 --- a/sql3/test/defs/defs.go +++ b/sql3/test/defs/defs.go @@ -15,6 +15,8 @@ var TableTests []TableTest = []TableTest{ selectTests, + topTests, + setLiteralTests, setFunctionTests, setParameterTests, diff --git a/sql3/test/defs/defs_top.go b/sql3/test/defs/defs_top.go new file mode 100644 index 000000000..1be23cc49 --- /dev/null +++ b/sql3/test/defs/defs_top.go @@ -0,0 +1,55 @@ +package defs + +var topTests = TableTest{ + name: "top-tests", + Table: tbl( + "skills", + srcHdrs( + srcHdr("_id", fldTypeID), + srcHdr("bools", fldTypeStringSet), + srcHdr("bools-exist", fldTypeStringSet), + srcHdr("id1", fldTypeID), + srcHdr("skills", fldTypeStringSet), + srcHdr("titles", fldTypeStringSet), + ), + srcRows( + srcRow(int64(1), nil, []string{"available_for_hire"}, int64(288), []string{"Marketing Manager"}, []string{"OEM negotiations", "Alumni Relations"}), + srcRow(int64(2), nil, []string{"available_for_hire"}, int64(288), []string{"Software Engineer I"}, []string{"Chief Cook", "Bottle Washer"}), + ), + ), + SQLTests: []SQLTest{ + { + SQLs: sqls( + "select top(1) * from skills where setcontains(skills, 'Marketing Manager');", + ), + ExpHdrs: hdrs( + hdr("_id", fldTypeID), + hdr("bools", fldTypeStringSet), + hdr("bools-exist", fldTypeStringSet), + hdr("id1", fldTypeID), + hdr("skills", fldTypeStringSet), + hdr("titles", fldTypeStringSet), + ), + ExpRows: rows( + row(int64(1), nil, []string{"available_for_hire"}, int64(288), []string{"Marketing Manager"}, []string{"Alumni Relations", "OEM negotiations"}), + ), + Compare: CompareExactUnordered, + SortStringKeys: true, + }, + { + SQLs: sqls( + "select top(10) count(*), skills from skills group by skills;", + ), + ExpHdrs: hdrs( + hdr("", fldTypeInt), + hdr("skills", fldTypeStringSet), + ), + ExpRows: rows( + row(int64(1), []string{"Marketing Manager"}), + row(int64(1), []string{"Software Engineer I"}), + ), + Compare: CompareExactUnordered, + SortStringKeys: true, + }, + }, +}