diff --git a/sql3/planner/compilealtertable.go b/sql3/planner/compilealtertable.go index b3c215179..c52b35128 100644 --- a/sql3/planner/compilealtertable.go +++ b/sql3/planner/compilealtertable.go @@ -92,7 +92,7 @@ func (p *ExecutionPlanner) analyzeAlterTableStatement(stmt *parser.AlterTableSta return sql3.NewErrUnknownType(col.Type.Name.NamePos.Line, col.Type.Name.NamePos.Column, typeName) } - if strings.ToLower(columnName) == "_id" { + if strings.ToLower(columnName) == string(dax.PrimaryKeyFieldName) { //not allowed to add an _id column after the fact return sql3.NewErrTableIDColumnAlter(col.Name.NamePos.Line, col.Name.NamePos.Column) } diff --git a/sql3/planner/compilebulkinsert.go b/sql3/planner/compilebulkinsert.go index b5cb7b1ea..3f42fb4f8 100644 --- a/sql3/planner/compilebulkinsert.go +++ b/sql3/planner/compilebulkinsert.go @@ -383,7 +383,7 @@ func (p *ExecutionPlanner) analyzeBulkInsertStatement(ctx context.Context, stmt } columnNameMap[colName] = struct{}{} - if strings.EqualFold(cm.Name, "_id") { + if strings.EqualFold(cm.Name, string(dax.PrimaryKeyFieldName)) { foundID = true } } diff --git a/sql3/planner/compilecreatetable.go b/sql3/planner/compilecreatetable.go index 4eb3b859c..0712afc8e 100644 --- a/sql3/planner/compilecreatetable.go +++ b/sql3/planner/compilecreatetable.go @@ -54,7 +54,7 @@ func (p *ExecutionPlanner) compileCreateTableStatement(ctx context.Context, stmt columnName := strings.ToLower(parser.IdentName(col.Name)) typeName := parser.IdentName(col.Type.Name) - if strings.ToLower(columnName) == "_id" { + if strings.ToLower(columnName) == string(dax.PrimaryKeyFieldName) { if strings.EqualFold(typeName, dax.BaseTypeString) { isKeyed = true } @@ -262,7 +262,7 @@ func (p *ExecutionPlanner) analyzeCreateTableStatement(stmt *parser.CreateTableS return sql3.NewErrUnknownType(col.Type.Name.NamePos.Line, col.Type.Name.NamePos.Column, typeName) } - if strings.ToLower(columnName) == "_id" { + if strings.ToLower(columnName) == string(dax.PrimaryKeyFieldName) { //check the type if !(strings.EqualFold(typeName, dax.BaseTypeID) || strings.EqualFold(typeName, dax.BaseTypeString)) { return sql3.NewErrTableIDColumnType(col.Type.Name.NamePos.Line, col.Type.Name.NamePos.Column) @@ -279,7 +279,7 @@ func (p *ExecutionPlanner) analyzeCreateTableStatement(stmt *parser.CreateTableS return err } } - _, ok := checkedColumns["_id"] + _, ok := checkedColumns[string(dax.PrimaryKeyFieldName)] if !ok { return sql3.NewErrTableMustHaveIDColumn(stmt.Create.Line, stmt.Create.Column) } diff --git a/sql3/planner/compileinsert.go b/sql3/planner/compileinsert.go index 86c21c511..1a0a9e51f 100644 --- a/sql3/planner/compileinsert.go +++ b/sql3/planner/compileinsert.go @@ -33,7 +33,7 @@ func (p *ExecutionPlanner) compileInsertStatement(ctx context.Context, stmt *par for _, columnIdent := range stmt.Columns { colName := strings.ToLower(parser.IdentName(columnIdent)) - if strings.EqualFold(colName, "_id") { + if strings.EqualFold(colName, string(dax.PrimaryKeyFieldName)) { targetColumns = append(targetColumns, newQualifiedRefPlanExpression(tableName, colName, 0, parser.NewDataTypeID())) continue } @@ -111,8 +111,8 @@ func (p *ExecutionPlanner) analyzeInsertStatement(ctx context.Context, stmt *par colName := strings.ToLower(parser.IdentName(columnIdent)) var typeName parser.ExprDataType - if strings.EqualFold(colName, "_id") { - columnNameMap["_id"] = struct{}{} + if strings.EqualFold(colName, string(dax.PrimaryKeyFieldName)) { + columnNameMap[string(dax.PrimaryKeyFieldName)] = struct{}{} // Determine, from the existing table, whether the _id is of // type ID or STRING. @@ -151,7 +151,7 @@ func (p *ExecutionPlanner) analyzeInsertStatement(ctx context.Context, stmt *par } // Ensure we have an _id column. - if _, ok := columnNameMap["_id"]; !ok { + if _, ok := columnNameMap[string(dax.PrimaryKeyFieldName)]; !ok { return sql3.NewErrInsertMustHaveIDColumn(stmt.ColumnsLparen.Line, stmt.ColumnsLparen.Column) } diff --git a/sql3/planner/compileselect.go b/sql3/planner/compileselect.go index a40723523..9a03bda0f 100644 --- a/sql3/planner/compileselect.go +++ b/sql3/planner/compileselect.go @@ -83,12 +83,12 @@ func (p *ExecutionPlanner) compileSelectStatement(stmt *parser.SelectStatement, for _, agg := range aggregates { InspectExpression(agg, func(expr types.PlanExpression) bool { switch ex := expr.(type) { - case *sumPlanExpression, *countPlanExpression, *countDistinctPlanExpression, - *avgPlanExpression, *minPlanExpression, *maxPlanExpression, - *percentilePlanExpression: + case types.Aggregable: ch := ex.Children() - // first arg is always the ref - aggregateAndGroupByExprs = append(aggregateAndGroupByExprs, ch[0]) + // first arg is always the ref, except for count(*) + if len(ch) > 0 { + aggregateAndGroupByExprs = append(aggregateAndGroupByExprs, ch[0]) + } return false } return true @@ -100,9 +100,7 @@ func (p *ExecutionPlanner) compileSelectStatement(stmt *parser.SelectStatement, havingReferences := make([]*qualifiedRefPlanExpression, 0) InspectExpression(having, func(expr types.PlanExpression) bool { switch ex := expr.(type) { - case *sumPlanExpression, *countPlanExpression, *countDistinctPlanExpression, - *avgPlanExpression, *minPlanExpression, *maxPlanExpression, - *percentilePlanExpression: + case types.Aggregable: return false case *qualifiedRefPlanExpression: havingReferences = append(havingReferences, ex) @@ -139,9 +137,7 @@ func (p *ExecutionPlanner) compileSelectStatement(stmt *parser.SelectStatement, for _, expr := range projections { InspectExpression(expr, func(expr types.PlanExpression) bool { switch ex := expr.(type) { - case *sumPlanExpression, *countPlanExpression, *countDistinctPlanExpression, - *avgPlanExpression, *minPlanExpression, *maxPlanExpression, - *percentilePlanExpression: + case types.Aggregable: //return false for these, because thats as far down we want to inspect return false case *qualifiedRefPlanExpression: @@ -237,9 +233,7 @@ func (p *ExecutionPlanner) gatherExprAggregates(expr types.PlanExpression, aggre result := aggregates InspectExpression(expr, func(expr types.PlanExpression) bool { switch ex := expr.(type) { - case *sumPlanExpression, *countPlanExpression, *countDistinctPlanExpression, - *avgPlanExpression, *minPlanExpression, *maxPlanExpression, - *percentilePlanExpression: + case types.Aggregable: found := false for _, ag := range result { //compare based on string representation @@ -249,7 +243,7 @@ func (p *ExecutionPlanner) gatherExprAggregates(expr types.PlanExpression, aggre } } if !found { - result = append(result, ex) + result = append(result, ex.(types.PlanExpression)) } // return false because thats as far down we want to inspect return false diff --git a/sql3/planner/compileshow.go b/sql3/planner/compileshow.go index 35a28b0a3..5fd14bbd0 100644 --- a/sql3/planner/compileshow.go +++ b/sql3/planner/compileshow.go @@ -23,7 +23,7 @@ func (p *ExecutionPlanner) compileShowDatabasesStatement(ctx context.Context, st columns := []types.PlanExpression{ &qualifiedRefPlanExpression{ tableName: "fb_databases", - columnName: "_id", + columnName: string(dax.PrimaryKeyFieldName), columnIndex: 0, dataType: parser.NewDataTypeString(), }, @@ -82,7 +82,7 @@ func (p *ExecutionPlanner) compileShowTablesStatement(ctx context.Context, stmt columns := []types.PlanExpression{ &qualifiedRefPlanExpression{ tableName: "fb_tables", - columnName: "_id", + columnName: string(dax.PrimaryKeyFieldName), columnIndex: 0, dataType: parser.NewDataTypeString(), }, @@ -151,7 +151,7 @@ func (p *ExecutionPlanner) compileShowColumnsStatement(ctx context.Context, stmt columns := []types.PlanExpression{&qualifiedRefPlanExpression{ tableName: "fb_table_columns", - columnName: "_id", + columnName: string(dax.PrimaryKeyFieldName), columnIndex: 0, dataType: parser.NewDataTypeString(), }, &qualifiedRefPlanExpression{ diff --git a/sql3/planner/expression.go b/sql3/planner/expression.go index 0476c7089..7ad6314de 100644 --- a/sql3/planner/expression.go +++ b/sql3/planner/expression.go @@ -2787,7 +2787,11 @@ func (p *ExecutionPlanner) compileCallExpr(expr *parser.Call) (_ types.PlanExpre if expr.Distinct.IsValid() { agg = newCountDistinctPlanExpression(args[0], expr.ResultDataType) } else { - agg = newCountPlanExpression(args[0], expr.ResultDataType) + if expr.Star.IsValid() { + agg = newCountStarPlanExpression(expr.ResultDataType) + } else { + agg = newCountPlanExpression(args[0], expr.ResultDataType) + } } return agg, nil diff --git a/sql3/planner/expressionagg.go b/sql3/planner/expressionagg.go index 26ac2551f..e1098c1f8 100644 --- a/sql3/planner/expressionagg.go +++ b/sql3/planner/expressionagg.go @@ -76,6 +76,67 @@ func (c *aggregateCountDistinct) Eval(ctx context.Context) (interface{}, error) return int64(len(c.valueSeen)), nil } +// countStarPlanExpression handles COUNT(*) +type countStarPlanExpression struct { + arg types.PlanExpression + returnDataType parser.ExprDataType +} + +var _ types.Aggregable = (*countStarPlanExpression)(nil) + +func newCountStarPlanExpression(returnDataType parser.ExprDataType) *countStarPlanExpression { + return &countStarPlanExpression{ + returnDataType: returnDataType, + } +} + +func (n *countStarPlanExpression) Evaluate(currentRow []interface{}) (interface{}, error) { + if n.arg != nil { + arg, ok := n.arg.(*qualifiedRefPlanExpression) + if !ok { + return nil, sql3.NewErrInternalf("unexpected aggregate function arg type '%T'", n.arg) + } + return currentRow[arg.columnIndex], nil + } + return int64(1), nil +} + +func (n *countStarPlanExpression) NewBuffer() (types.AggregationBuffer, error) { + return NewAggCountBuffer(n), nil +} + +func (n *countStarPlanExpression) FirstChildExpr() types.PlanExpression { + return n.arg +} + +func (n *countStarPlanExpression) Type() parser.ExprDataType { + return n.returnDataType +} + +func (n *countStarPlanExpression) String() string { + return "count(*)" +} + +func (n *countStarPlanExpression) Plan() map[string]interface{} { + result := make(map[string]interface{}) + result["_expr"] = fmt.Sprintf("%T", n) + result["description"] = n.String() + result["dataType"] = n.Type().TypeDescription() + return result +} + +func (n *countStarPlanExpression) Children() []types.PlanExpression { + return []types.PlanExpression{} +} + +func (n *countStarPlanExpression) WithChildren(children ...types.PlanExpression) (types.PlanExpression, error) { + if len(children) != 1 { + return nil, sql3.NewErrInternalf("unexpected number of children '%d'", len(children)) + } + n.arg = children[0] + return n, nil +} + // countPlanExpression handles COUNT() type countPlanExpression struct { arg types.PlanExpression diff --git a/sql3/planner/expressionanalyzercall.go b/sql3/planner/expressionanalyzercall.go index 02d8b281c..85deb038f 100644 --- a/sql3/planner/expressionanalyzercall.go +++ b/sql3/planner/expressionanalyzercall.go @@ -6,6 +6,7 @@ import ( "context" "strings" + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/sql3" "github.com/featurebasedb/featurebase/v3/sql3/parser" ) @@ -22,26 +23,16 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars } switch strings.ToUpper(call.Name.Name) { case "COUNT": - //check to see if we have a star, if we do turn it into a qualified ref to _id - if call.Star.IsValid() && len(call.Args) == 0 { - newArg := &parser.Ident{ - NamePos: call.Star, - Name: "_id", + if len(call.Args) > 0 && !call.Star.IsValid() { + // one argument only + if len(call.Args) != 1 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) } - arg, err := p.analyzeExpression(ctx, newArg, scope) - if err != nil { - return nil, err + //make sure it's a qualified ref + _, ok := call.Args[0].(*parser.QualifiedRef) + if !ok { + return nil, sql3.NewErrExpectedColumnReference(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } - call.Args = append(call.Args, arg) - } - // one argument only - if len(call.Args) != 1 { - return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) - } - //make sure it's a qualified ref - _, ok := call.Args[0].(*parser.QualifiedRef) - if !ok { - return nil, sql3.NewErrExpectedColumnReference(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } //COUNT always returns int call.ResultDataType = parser.NewDataTypeInt() @@ -58,7 +49,7 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars // if it is a ref, we shouldn't do a sum on the _id ref, ok := call.Args[0].(*parser.QualifiedRef) - if ok && strings.EqualFold(ref.Column.Name, "_id") { + if ok && strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Name.Name) } @@ -82,7 +73,7 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars // if it is a ref, we shouldn't do a avg on the _id ref, ok := call.Args[0].(*parser.QualifiedRef) - if ok && strings.EqualFold(ref.Column.Name, "_id") { + if ok && strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Name.Name) } @@ -110,7 +101,7 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars } //can't do a percentile on _id - if strings.EqualFold(ref.Column.Name, "_id") { + if strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Name.Name) } @@ -146,7 +137,7 @@ func (p *ExecutionPlanner) analyzeCallExpression(ctx context.Context, call *pars // if it is a ref, we shouldn't do a min/max on the _id ref, ok := call.Args[0].(*parser.QualifiedRef) - if ok && strings.EqualFold(ref.Column.Name, "_id") { + if ok && strings.EqualFold(ref.Column.Name, string(dax.PrimaryKeyFieldName)) { return nil, sql3.NewErrIdColumnNotValidForAggregateFunction(call.Args[0].Pos().Line, call.Args[0].Pos().Column, call.Name.Name) } diff --git a/sql3/planner/expressionpql.go b/sql3/planner/expressionpql.go index edb32f25e..e0f7ea290 100644 --- a/sql3/planner/expressionpql.go +++ b/sql3/planner/expressionpql.go @@ -7,6 +7,7 @@ import ( "strconv" "strings" + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/pql" "github.com/featurebasedb/featurebase/v3/sql3" "github.com/featurebasedb/featurebase/v3/sql3/parser" @@ -114,7 +115,7 @@ func (p *ExecutionPlanner) generatePQLCallFromExpr(ctx context.Context, expr typ } // if it is the _id column, we can use ConstRow with a list - if strings.EqualFold(lhs.columnName, "_id") { + if strings.EqualFold(lhs.columnName, string(dax.PrimaryKeyFieldName)) { values := make([]interface{}, len(list.exprs)) for i, m := range list.exprs { pqlValue, err := planExprToValue(m) @@ -204,7 +205,7 @@ func (p *ExecutionPlanner) generatePQLCallFromBinaryExpr(ctx context.Context, ex }, nil case *parser.DataTypeID: - if strings.EqualFold(lhs.columnName, "_id") { + if strings.EqualFold(lhs.columnName, string(dax.PrimaryKeyFieldName)) { return &pql.Call{ Name: "ConstRow", Args: map[string]interface{}{ @@ -221,7 +222,7 @@ func (p *ExecutionPlanner) generatePQLCallFromBinaryExpr(ctx context.Context, ex }, nil case *parser.DataTypeString: - if strings.EqualFold(lhs.columnName, "_id") { + if strings.EqualFold(lhs.columnName, string(dax.PrimaryKeyFieldName)) { return &pql.Call{ Name: "ConstRow", Args: map[string]interface{}{ @@ -436,8 +437,8 @@ func (p *ExecutionPlanner) generatePQLCallFromBinaryExpr(ctx context.Context, ex } switch typ := expr.lhs.Type().(type) { case *parser.DataTypeID: - if strings.EqualFold(lhs.columnName, "_id") { - return nil, sql3.NewErrInvalidColumnInFilterExpression(0, 0, "_id", "is/is not null") + if strings.EqualFold(lhs.columnName, string(dax.PrimaryKeyFieldName)) { + return nil, sql3.NewErrInvalidColumnInFilterExpression(0, 0, string(dax.PrimaryKeyFieldName), "is/is not null") } return nil, sql3.NewErrInvalidTypeInFilterExpression(0, 0, typ.TypeDescription(), "is/is not null") diff --git a/sql3/planner/expressiontypes.go b/sql3/planner/expressiontypes.go index e7342170d..cab4add69 100644 --- a/sql3/planner/expressiontypes.go +++ b/sql3/planner/expressiontypes.go @@ -23,7 +23,7 @@ func fieldSQLDataType(f *pilosa.FieldInfo) parser.ExprDataType { // a FieldTypeID. Another thing to be updated is the "_id" value itself; in // the dax package there is a constant called `PrimaryKeyFieldName` which // would be used here instead. - if f.Name == "_id" { + if f.Name == string(dax.PrimaryKeyFieldName) { switch f.Options.Type { case "id": return parser.NewDataTypeID() diff --git a/sql3/planner/inbuiltfunctionstable.go b/sql3/planner/inbuiltfunctionstable.go index f4527df7e..0968e659f 100644 --- a/sql3/planner/inbuiltfunctionstable.go +++ b/sql3/planner/inbuiltfunctionstable.go @@ -1,6 +1,7 @@ package planner import ( + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/sql3" "github.com/featurebasedb/featurebase/v3/sql3/parser" ) @@ -50,7 +51,7 @@ func (p *ExecutionPlanner) analyzeFunctionSubtable(call *parser.Call, scope pars } call.ResultDataType = parser.NewDataTypeSubtable([]*parser.SubtableColumn{ { - Name: "_id", + Name: string(dax.PrimaryKeyFieldName), DataType: parser.NewDataTypeID(), }, { diff --git a/sql3/planner/opcreatetable.go b/sql3/planner/opcreatetable.go index 0409d5176..fe9c877f3 100644 --- a/sql3/planner/opcreatetable.go +++ b/sql3/planner/opcreatetable.go @@ -106,7 +106,7 @@ func (i *createTableRowIter) Next(ctx context.Context) (types.Row, error) { idType = dax.BaseTypeString } fields = append(fields, &dax.Field{ - Name: "_id", + Name: dax.PrimaryKeyFieldName, Type: idType, }) diff --git a/sql3/planner/opfeaturebasecolumns.go b/sql3/planner/opfeaturebasecolumns.go index 444e5137c..8fd1e38a3 100644 --- a/sql3/planner/opfeaturebasecolumns.go +++ b/sql3/planner/opfeaturebasecolumns.go @@ -49,7 +49,7 @@ func (p *PlanOpFeatureBaseColumns) Schema() types.Schema { return types.Schema{ &types.PlannerColumn{ RelationName: "fb_table_columns", - ColumnName: "_id", + ColumnName: string(dax.PrimaryKeyFieldName), Type: parser.NewDataTypeString(), }, &types.PlannerColumn{ diff --git a/sql3/planner/opfeaturebasedatabases.go b/sql3/planner/opfeaturebasedatabases.go index 7f3ec5645..ae6f13860 100644 --- a/sql3/planner/opfeaturebasedatabases.go +++ b/sql3/planner/opfeaturebasedatabases.go @@ -51,7 +51,7 @@ func (p *PlanOpFeatureBaseDatabases) Schema() types.Schema { return types.Schema{ &types.PlannerColumn{ RelationName: "fb_databases", - ColumnName: "_id", + ColumnName: string(dax.PrimaryKeyFieldName), Type: parser.NewDataTypeString(), }, &types.PlannerColumn{ diff --git a/sql3/planner/opfeaturebasetables.go b/sql3/planner/opfeaturebasetables.go index 2eae4c2cc..0e439f0f0 100644 --- a/sql3/planner/opfeaturebasetables.go +++ b/sql3/planner/opfeaturebasetables.go @@ -9,6 +9,7 @@ import ( "time" pilosa "github.com/featurebasedb/featurebase/v3" + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/sql3/parser" "github.com/featurebasedb/featurebase/v3/sql3/planner/types" ) @@ -52,7 +53,7 @@ func (p *PlanOpFeatureBaseTables) Schema() types.Schema { return types.Schema{ &types.PlannerColumn{ RelationName: "fb_tables", - ColumnName: "_id", + ColumnName: string(dax.PrimaryKeyFieldName), Type: parser.NewDataTypeString(), }, &types.PlannerColumn{ diff --git a/sql3/planner/opinsert.go b/sql3/planner/opinsert.go index 53a319398..b5f7dee22 100644 --- a/sql3/planner/opinsert.go +++ b/sql3/planner/opinsert.go @@ -116,7 +116,7 @@ func (i *insertRowIter) Next(ctx context.Context) (types.Row, error) { posVals[j] = j - 1 continue } - if strings.EqualFold(i.targetColumns[j].columnName, "_id") { + if strings.EqualFold(i.targetColumns[j].columnName, string(dax.PrimaryKeyFieldName)) { posID = j foundPosID = true } diff --git a/sql3/planner/oppqlaggregate.go b/sql3/planner/oppqlaggregate.go index 64c3698f6..09169d6b6 100644 --- a/sql3/planner/oppqlaggregate.go +++ b/sql3/planner/oppqlaggregate.go @@ -42,7 +42,7 @@ func (p *PlanOpPQLAggregate) Plan() map[string]interface{} { if p.filter != nil { result["filter"] = p.filter.Plan() } - result["aggregate"] = p.aggregate.FirstChildExpr().Plan() + result["aggregate"] = p.aggregate.String() return result } @@ -64,7 +64,7 @@ func (p *PlanOpPQLAggregate) Schema() types.Schema { s := &types.PlannerColumn{ ColumnName: "", RelationName: "", - Type: p.aggregate.FirstChildExpr().Type(), + Type: p.aggregate.Type(), } result[0] = s return result @@ -135,7 +135,7 @@ func (i *pqlAggregateRowIter) Next(ctx context.Context) (types.Row, error) { call = &pql.Call{Name: "Count", Children: []*pql.Call{cond}} - case *countPlanExpression: + case *countPlanExpression, *countStarPlanExpression: if cond == nil { // COUNT() should ignore null values // if the data type of the expression supports an existence bitmap for @@ -156,7 +156,7 @@ func (i *pqlAggregateRowIter) Next(ctx context.Context) (types.Row, error) { case *avgPlanExpression: if cond == nil { - // COUNT() should ignore null values + // SUM() should ignore null values // if the data type of the expression supports an existence bitmap for // the underlying FeatureBase data type use it to eliminate nulls from the aggregate switch expr.dataType.(type) { diff --git a/sql3/planner/oppqldelete.go b/sql3/planner/oppqldelete.go index b4a6e40d8..051a98581 100644 --- a/sql3/planner/oppqldelete.go +++ b/sql3/planner/oppqldelete.go @@ -102,7 +102,7 @@ func (p *PlanOpPQLConstRowDelete) Expressions() []types.PlanExpression { tableName: p.tableName, columnIndex: 0, dataType: colType, - columnName: "_id", + columnName: string(dax.PrimaryKeyFieldName), }, } } diff --git a/sql3/planner/oppqldistinctscan.go b/sql3/planner/oppqldistinctscan.go index a85a0afee..f5f5004da 100644 --- a/sql3/planner/oppqldistinctscan.go +++ b/sql3/planner/oppqldistinctscan.go @@ -28,7 +28,7 @@ type PlanOpPQLDistinctScan struct { } func NewPlanOpPQLDistinctScan(p *ExecutionPlanner, tableName string, column string) (*PlanOpPQLDistinctScan, error) { - if strings.EqualFold("_id", column) { + if strings.EqualFold(string(dax.PrimaryKeyFieldName), column) { return nil, sql3.NewErrInternalf("non _id column required") } return &PlanOpPQLDistinctScan{ diff --git a/sql3/planner/oppqlgroupby.go b/sql3/planner/oppqlgroupby.go index 41686cc40..5687a9e73 100644 --- a/sql3/planner/oppqlgroupby.go +++ b/sql3/planner/oppqlgroupby.go @@ -44,7 +44,7 @@ func (p *PlanOpPQLGroupBy) Plan() map[string]interface{} { if p.filter != nil { result["filter"] = p.filter.Plan() } - result["aggregate"] = p.aggregate.FirstChildExpr().Plan() + result["aggregate"] = p.aggregate.String() ps := make([]interface{}, 0) for _, e := range p.groupByExprs { ps = append(ps, e.Plan()) @@ -82,7 +82,7 @@ func (p *PlanOpPQLGroupBy) Schema() types.Schema { s := &types.PlannerColumn{ ColumnName: p.aggregate.String(), RelationName: "", - Type: p.aggregate.FirstChildExpr().Type(), + Type: p.aggregate.Type(), } result[len(p.groupByExprs)] = s @@ -148,7 +148,7 @@ func (i *pqlGroupByRowIter) Next(ctx context.Context) (types.Row, error) { return nil, sql3.NewErrInternalf("unexpected expression type in group by list '%T'", c) } //don't ask for the _id field - if ref.Name() != "_id" { + if ref.Name() != string(dax.PrimaryKeyFieldName) { call.Children = append(call.Children, &pql.Call{ Name: "Rows", @@ -165,7 +165,7 @@ func (i *pqlGroupByRowIter) Next(ctx context.Context) (types.Row, error) { } switch i.aggregate.(type) { - case *countPlanExpression: + case *countPlanExpression, *countStarPlanExpression: //nop case *countDistinctPlanExpression: @@ -245,7 +245,7 @@ func (i *pqlGroupByRowIter) Next(ctx context.Context) (types.Row, error) { //now populate the aggregate value aggIdx := len(i.groupByColumns) switch i.aggregate.(type) { - case *countPlanExpression: + case *countPlanExpression, *countStarPlanExpression: row[aggIdx] = int64(group.Count) case *countDistinctPlanExpression, *sumPlanExpression: diff --git a/sql3/planner/oppqltablescan.go b/sql3/planner/oppqltablescan.go index 178a6b36c..431cf30c4 100644 --- a/sql3/planner/oppqltablescan.go +++ b/sql3/planner/oppqltablescan.go @@ -113,6 +113,19 @@ func (p *PlanOpPQLTableScan) WithChildren(children ...types.PlanOperator) (types return nil, nil } +func (p *PlanOpPQLTableScan) PrimaryKeyType() (parser.ExprDataType, error) { + tname := dax.TableName(p.tableName) + table, err := p.planner.schemaAPI.TableByName(context.Background(), tname) + if err != nil { + return nil, err + } + + if table.StringKeys() { + return parser.NewDataTypeString(), nil + } + return parser.NewDataTypeID(), nil +} + type targetColumn struct { columnIdx int srcColumnIdx int @@ -198,7 +211,7 @@ func (i *tableScanRowIter) Next(ctx context.Context) (types.Row, error) { for _, c := range i.columns { // skip the _id field - if strings.EqualFold(c, "_id") { + if strings.EqualFold(c, string(dax.PrimaryKeyFieldName)) { continue } @@ -250,7 +263,7 @@ func (i *tableScanRowIter) Next(ctx context.Context) (types.Row, error) { mappedColIdx := mappedColumn.columnIdx mappedSrcColIdx := mappedColumn.srcColumnIdx - if strings.EqualFold(c, "_id") { + if strings.EqualFold(c, string(dax.PrimaryKeyFieldName)) { if result.Column.Keyed { row[mappedColIdx] = result.Column.Key } else { diff --git a/sql3/planner/planoptimizer.go b/sql3/planner/planoptimizer.go index 6a08d81ad..e5c1761ef 100644 --- a/sql3/planner/planoptimizer.go +++ b/sql3/planner/planoptimizer.go @@ -8,6 +8,7 @@ import ( "reflect" "strings" + "github.com/featurebasedb/featurebase/v3/dax" "github.com/featurebasedb/featurebase/v3/sql3" "github.com/featurebasedb/featurebase/v3/sql3/parser" "github.com/featurebasedb/featurebase/v3/sql3/planner/types" @@ -636,22 +637,34 @@ func tryToReplaceGroupByWithPQLAggregate(ctx context.Context, a *ExecutionPlanne return thisNode, true, nil } + pkType, err := table.PrimaryKeyType() + if err != nil { + return thisNode, true, err + } // we can push down to pql if: // 1. the expression we are aggregating on is a qualifiedRef // 2. it is a bsi type // we always push down to pql if it's a ref and it's the _id column - for _, agg := range thisNode.Aggregates { - aggregable, ok := agg.(types.Aggregable) - if !ok { - return n, false, sql3.NewErrInternalf("unexpected aggregate function arg type '%T'", agg) - } - switch ref := aggregable.FirstChildExpr().(type) { - case *qualifiedRefPlanExpression: - if !strings.EqualFold(ref.columnName, "_id") && !typeIsBSI(ref.Type()) { + for i, agg := range thisNode.Aggregates { + switch aggregable := agg.(type) { + case *countStarPlanExpression: + // it's a count(*) on a pql table scan, so add the arg + newChildren := []types.PlanExpression{newQualifiedRefPlanExpression(table.tableName, string(dax.PrimaryKeyFieldName), 0, pkType)} + newAgg, err := aggregable.WithChildren(newChildren...) + if err != nil { + return n, true, err + } + thisNode.Aggregates[i] = newAgg + + case types.Aggregable: + switch ref := aggregable.FirstChildExpr().(type) { + case *qualifiedRefPlanExpression: + if !strings.EqualFold(ref.columnName, string(dax.PrimaryKeyFieldName)) && !typeIsBSI(ref.Type()) { + return thisNode, true, nil + } + default: return thisNode, true, nil } - default: - return thisNode, true, nil } } @@ -740,7 +753,7 @@ func tryToReplaceDistinctWithPQLDistinct(ctx context.Context, a *ExecutionPlanne } // make sure it's not the _id column - if strings.EqualFold(thisNode.columns[0], "_id") { + if strings.EqualFold(thisNode.columns[0], string(dax.PrimaryKeyFieldName)) { return thisNode, true, nil } @@ -809,15 +822,31 @@ func tryToReplaceGroupByWithPQLGroupBy(ctx context.Context, a *ExecutionPlanner, table := tables[0] //only do this if we have group by expressions if len(n.GroupByExprs) > 0 { + pkType, err := table.PrimaryKeyType() + if err != nil { + return n, true, err + } + //use a multi group by if more than 1 aggregate if len(n.Aggregates) > 1 { ops := make([]*PlanOpPQLGroupBy, 0) for _, agg := range n.Aggregates { - aggregable, ok := agg.(types.Aggregable) if !ok { return n, false, sql3.NewErrInternalf("unexpected aggregate function arg type '%T'", agg) } + + // if it's a count(*) on a pql table scan, so add the arg + star, ok := agg.(*countStarPlanExpression) + if ok { + newChildren := []types.PlanExpression{newQualifiedRefPlanExpression(table.tableName, string(dax.PrimaryKeyFieldName), 0, pkType)} + newAgg, err := star.WithChildren(newChildren...) + if err != nil { + return n, true, err + } + aggregable = newAgg.(types.Aggregable) + } + ops = append(ops, NewPlanOpPQLGroupBy(a, table.tableName, n.GroupByExprs, table.filter, aggregable)) } newOp := NewPlanOpPQLMultiGroupBy(a, ops, n.GroupByExprs) @@ -829,6 +858,17 @@ func tryToReplaceGroupByWithPQLGroupBy(ctx context.Context, a *ExecutionPlanner, if !ok { return n, false, sql3.NewErrInternalf("unexpected aggregate function arg type '%T'", n.Aggregates[0]) } + + // if it's a count(*) on a pql table scan, so add the arg + star, ok := aggregable.(*countStarPlanExpression) + if ok { + newChildren := []types.PlanExpression{newQualifiedRefPlanExpression(table.tableName, string(dax.PrimaryKeyFieldName), 0, pkType)} + newAgg, err := star.WithChildren(newChildren...) + if err != nil { + return n, true, err + } + aggregable = newAgg.(types.Aggregable) + } newOp := NewPlanOpPQLGroupBy(a, table.tableName, n.GroupByExprs, table.filter, aggregable) return newOp, false, nil } diff --git a/sql3/planner/systemobjects.go b/sql3/planner/systemobjects.go index 653a712ca..8fedda5bd 100644 --- a/sql3/planner/systemobjects.go +++ b/sql3/planner/systemobjects.go @@ -128,7 +128,7 @@ func (p *ExecutionPlanner) getViewByName(ctx context.Context, name string) (*vie tableName: "fb_views", columns: cols, predicate: newBinOpPlanExpression( - newQualifiedRefPlanExpression("fb_views", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_views", string(dax.PrimaryKeyFieldName), 0, parser.NewDataTypeString()), parser.EQ, newStringLiteralPlanExpression(name), parser.NewDataTypeBool(), @@ -163,7 +163,7 @@ func (p *ExecutionPlanner) insertView(ctx context.Context, view *viewSystemObjec planner: p, tableName: "fb_views", targetColumns: []*qualifiedRefPlanExpression{ - newQualifiedRefPlanExpression("fb_views", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_views", string(dax.PrimaryKeyFieldName), 0, parser.NewDataTypeString()), newQualifiedRefPlanExpression("fb_views", "name", 0, parser.NewDataTypeString()), newQualifiedRefPlanExpression("fb_views", "statement", 0, parser.NewDataTypeString()), newQualifiedRefPlanExpression("fb_views", "owner", 0, parser.NewDataTypeString()), @@ -202,7 +202,7 @@ func (p *ExecutionPlanner) updateView(ctx context.Context, view *viewSystemObjec planner: p, tableName: "fb_views", targetColumns: []*qualifiedRefPlanExpression{ - newQualifiedRefPlanExpression("fb_views", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_views", string(dax.PrimaryKeyFieldName), 0, parser.NewDataTypeString()), newQualifiedRefPlanExpression("fb_views", "statement", 0, parser.NewDataTypeString()), newQualifiedRefPlanExpression("fb_views", "updated_by", 0, parser.NewDataTypeString()), newQualifiedRefPlanExpression("fb_views", "updated_at", 0, parser.NewDataTypeTimestamp()), @@ -233,7 +233,7 @@ func (p *ExecutionPlanner) deleteView(ctx context.Context, viewName string) erro planner: p, tableName: "fb_views", filter: newBinOpPlanExpression( - newQualifiedRefPlanExpression("fb_views", "_id", 0, parser.NewDataTypeString()), + newQualifiedRefPlanExpression("fb_views", string(dax.PrimaryKeyFieldName), 0, parser.NewDataTypeString()), parser.EQ, newStringLiteralPlanExpression(viewName), parser.NewDataTypeBool(), diff --git a/sql3/test/defs/defs.go b/sql3/test/defs/defs.go index ebec2a002..b01a4bad7 100644 --- a/sql3/test/defs/defs.go +++ b/sql3/test/defs/defs.go @@ -10,6 +10,7 @@ import ( "github.com/PaesslerAG/gval" "github.com/PaesslerAG/jsonpath" + "github.com/featurebasedb/featurebase/v3/errors" ) // TableTests is the list of tests which get run by TestSQL_Execute in @@ -221,17 +222,17 @@ func operatorPresentAtPath(jplan []byte, path string, operator string) error { v := interface{}(nil) err := json.Unmarshal(jplan, &v) if err != nil { - return err + return errors.Wrap(err, fmt.Sprintf("expected '%s' to be present", operator)) } builder := gval.Full(jsonpath.PlaceholderExtension()) expr, err := builder.NewEvaluable(path) if err != nil { - return err + return errors.Wrap(err, fmt.Sprintf("expected '%s' to be present", operator)) } eval, err := expr(context.Background(), v) if err != nil { - return err + return errors.Wrap(err, fmt.Sprintf("expected '%s' to be present", operator)) } s, ok := eval.(string) if ok && strings.EqualFold(s, operator) {