From 98e7ade5915eeb73f2b901d583f789a4bb1d4adf Mon Sep 17 00:00:00 2001 From: Ben Johnson Date: Thu, 7 Oct 2021 14:20:07 -0600 Subject: [PATCH] CORE-809: Aggregate COUNT() with INNER JOIN --- planner.go | 263 +++++++++++++++++++++++++++++++++++++++--------- planner_test.go | 83 +++++++++++++++ sql2/ast.go | 130 ++++++++++++++++++++++++ 3 files changed, 426 insertions(+), 50 deletions(-) diff --git a/planner.go b/planner.go index 2a418b6c3..7942b32b4 100644 --- a/planner.go +++ b/planner.go @@ -62,6 +62,11 @@ func (p *Planner) planSelectStatement(ctx context.Context, stmt *sql2.SelectStat } func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) { + // Handle specific case of a two-table INNER JOIN with a COUNT(). + if _, ok := stmt.Source.(*sql2.JoinClause); ok { + return p.planAggregateCountJoin(ctx, stmt) + } + indexName, err := statementTableName(stmt) if err != nil { return nil, err @@ -116,7 +121,13 @@ func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.S switch callName { case "COUNT": if len(groupByCols) == 0 { - return NewCountNode(p.executor, indexName, columns[0], cond), nil + if cond == nil { + cond = &pql.Call{Name: "All"} + } + return NewCountNode(p.executor, indexName, columns[0], &pql.Call{ + Name: "Count", + Children: []*pql.Call{cond}, + }), nil } var aggregate *pql.Call @@ -163,6 +174,156 @@ func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.S // TODO: Support HAVING } +func (p *Planner) planAggregateCountJoin(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) { + // Ensure we have an INNER JOIN. + join := stmt.Source.(*sql2.JoinClause) // caller checked + if !join.Operator.Inner.IsValid() { + return nil, fmt.Errorf("only inner joins are currently supported") + } + + // Determine the two tables we are joining. + tbl0, ok := join.X.(*sql2.QualifiedTableName) + if !ok { + return nil, fmt.Errorf("left side of join must be a table") + } + tbl1, ok := join.Y.(*sql2.QualifiedTableName) + if !ok { + return nil, fmt.Errorf("left side of join must be a table") + } + + // Ensure INNER JOIN has an "ON" constraint. + if join.Constraint == nil { + return nil, fmt.Errorf("joins must have an ON constraint") + } + cons, ok := join.Constraint.(*sql2.OnConstraint) + if !ok { + return nil, fmt.Errorf("joins only support an ON constraint") + } + + // Determine the joined columns. + cx, ok := cons.X.(*sql2.BinaryExpr) + if !ok { + return nil, fmt.Errorf("join must use a binary expression") + } else if cx.Op != sql2.EQ { + return nil, fmt.Errorf("join must use an equality expression") + } + + // Extract join columns & validate that they reference known tables and join on "_id". + x, ok := cx.X.(*sql2.QualifiedRef) + if !ok { + return nil, fmt.Errorf("left-hand side of join expression must be a table-qualified column") + } else if x.Table.Name != tbl0.TableName() && x.Table.Name != tbl1.TableName() { + return nil, fmt.Errorf("no such table: %q", x.Table.Name) + } + + y, ok := cx.Y.(*sql2.QualifiedRef) + if !ok { + return nil, fmt.Errorf("right-hand side of join expression must be a table-qualified column") + } else if y.Table.Name != tbl0.TableName() && y.Table.Name != tbl1.TableName() { + return nil, fmt.Errorf("no such table: %q", y.Table.Name) + } + + if x.Column.Name != "_id" && y.Column.Name != "_id" { + return nil, fmt.Errorf("must join table on _id column") + } else if x.Column.Name == "_id" && y.Column.Name == "_id" { + return nil, fmt.Errorf("cannot join _id field of two tables") + } + + // Move ID column to LHS. + if x.Column.Name != "_id" { + x, y = y, x + } + + // Move parent table to LHS. + if x.Table.Name != tbl0.TableName() { + tbl0, tbl1 = tbl1, tbl0 + } + + // Ensure column expression is a single COUNT. + if len(stmt.Columns) != 1 { + return nil, fmt.Errorf("only COUNT() is supported on joined tables") + } + expr, ok := stmt.Columns[0].Expr.(*sql2.Call) + if !ok || strings.ToUpper(expr.Name.Name) != "COUNT" { + return nil, fmt.Errorf("only COUNT() is supported on joined tables") + } + + // Extract WHERE clause and separate by parent/child tables. + var cond0, cond1 sql2.Expr + for _, cond := range sql2.SplitExprTree(stmt.WhereExpr) { + tblName, ok := sql2.ExprTableName(cond) + if !ok { + return nil, fmt.Errorf("cannot filter across multiple tables in an expression") + } else if tblName == "" { + return nil, fmt.Errorf("expression must reference a table name") + } else if tblName != tbl0.TableName() && tblName != tbl1.TableName() { + return nil, fmt.Errorf("no such table: %q", tblName) + } + + // Match to parent table. + if tblName == tbl0.TableName() { + if cond0 == nil { + cond0 = cond + } else { + cond0 = &sql2.BinaryExpr{X: cond0, Op: sql2.AND, Y: cond} + } + continue + } + + // Match to child table. + if cond1 == nil { + cond1 = cond + } + cond1 = &sql2.BinaryExpr{X: cond1, Op: sql2.AND, Y: cond} + } + + // Convert conditions to PQL. + pqlCond0, err := p.planExprPQL(ctx, stmt, cond0) + if err != nil { + return nil, err + } else if pqlCond0 == nil { + pqlCond0 = &pql.Call{Name: "All"} + } + + pqlCond1, err := p.planExprPQL(ctx, stmt, cond1) + if err != nil { + return nil, err + } else if pqlCond1 == nil { + pqlCond1 = &pql.Call{ + Name: "Row", + Args: map[string]interface{}{y.Column.Name: &pql.Condition{ + Op: pql.NEQ, + }}, + } + } + + return NewCountNode(p.executor, tbl0.Name.Name, + &StmtColumn{ + Name: stmt.Columns[0].Name(), + Type: sql2.DataTypeInt, + }, + &pql.Call{ + Name: "Count", + Children: []*pql.Call{{ + Name: "Intersect", + Children: []*pql.Call{ + pqlCond0, + { + Name: "Distinct", + Children: []*pql.Call{ + pqlCond1, + }, + Args: map[string]interface{}{ + "index": tbl1.Name.Name, + "field": y.Column.Name, + }, + }, + }, + }}, + }, + ), nil +} + func (p *Planner) planNonAggregateSelectStatement(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) { indexName, err := statementTableName(stmt) if err != nil { @@ -389,6 +550,53 @@ func (p *Planner) checkStatement(stmt sql2.Statement) error { } func (p *Planner) checkSelectStatement(stmt *sql2.SelectStatement) error { + if err := p.expandSelectStatementWildcards(stmt); err != nil { + return err + } + + // Type check expressions in statement. + for _, col := range stmt.Columns { + if err := p.checkExpr(&col.Expr, stmt); err != nil { + return err + } + } + + if err := p.checkExpr(&stmt.WhereExpr, stmt); err != nil { + return err + } + + for i := range stmt.GroupByExprs { + if err := p.checkExpr(&stmt.GroupByExprs[i], stmt); err != nil { + return err + } + } + + if err := p.checkExpr(&stmt.HavingExpr, stmt); err != nil { + return err + } + + for _, term := range stmt.OrderingTerms { + if err := p.checkExpr(&term.X, stmt); err != nil { + return err + } + } + + if err := p.checkExpr(&stmt.LimitExpr, stmt); err != nil { + return err + } + + if err := p.checkExpr(&stmt.OffsetExpr, stmt); err != nil { + return err + } + + return nil +} + +func (p *Planner) expandSelectStatementWildcards(stmt *sql2.SelectStatement) error { + if !stmt.HasWildcard() { + return nil + } + indexName, err := statementTableName(stmt) if err != nil { return err @@ -441,44 +649,8 @@ func (p *Planner) checkSelectStatement(stmt *sql2.SelectStatement) error { } stmt.Columns = columns - // Type check expressions in statement. - for _, col := range stmt.Columns { - if err := p.checkExpr(&col.Expr, stmt); err != nil { - return err - } - } - - if err := p.checkExpr(&stmt.WhereExpr, stmt); err != nil { - return err - } - - for i := range stmt.GroupByExprs { - if err := p.checkExpr(&stmt.GroupByExprs[i], stmt); err != nil { - return err - } - } - - if err := p.checkExpr(&stmt.HavingExpr, stmt); err != nil { - return err - } - - for _, term := range stmt.OrderingTerms { - if err := p.checkExpr(&term.X, stmt); err != nil { - return err - } - } - - if err := p.checkExpr(&stmt.LimitExpr, stmt); err != nil { - return err - } - - if err := p.checkExpr(&stmt.OffsetExpr, stmt); err != nil { - return err - } - return nil } - func (p *Planner) checkExpr(expr *sql2.Expr, stmt sql2.Statement) error { if e, err := sql2.Walk(&sqlExprTypeChecker{ holder: p.executor.Holder, @@ -959,20 +1131,17 @@ type CountNode struct { executor *executor indexName string column *StmtColumn - cond *pql.Call // conditional + call *pql.Call row []interface{} } -func NewCountNode(executor *executor, indexName string, column *StmtColumn, cond *pql.Call) *CountNode { - if cond == nil { - cond = &pql.Call{Name: "All"} - } +func NewCountNode(executor *executor, indexName string, column *StmtColumn, call *pql.Call) *CountNode { return &CountNode{ executor: executor, indexName: indexName, column: column, - cond: cond, + call: call, } } @@ -990,13 +1159,7 @@ func (n *CountNode) Next(ctx context.Context) error { return sql.ErrNoRows } - q := &pql.Query{ - Calls: []*pql.Call{ - {Name: "Count", Children: []*pql.Call{n.cond}}, - }, - } - - result, err := n.executor.Execute(ctx, n.indexName, q, nil, nil) + result, err := n.executor.Execute(ctx, n.indexName, &pql.Query{Calls: []*pql.Call{n.call}}, nil, nil) if err != nil { return err } diff --git a/planner_test.go b/planner_test.go index 68d7505e4..8694b724c 100644 --- a/planner_test.go +++ b/planner_test.go @@ -428,6 +428,89 @@ func TestPlanner_GroupBy(t *testing.T) { }) } +func TestPlanner_InnerJoin(t *testing.T) { + c := test.MustRunCluster(t, 1) + defer c.Close() + + i0, err := c.GetHolder(0).CreateIndex("i0", pilosa.IndexOptions{TrackExistence: true}) + if err != nil { + t.Fatal(err) + } + defer i0.Close() + + if _, err := i0.CreateField("a", pilosa.OptFieldTypeInt(0, 1000)); err != nil { + t.Fatal(err) + } + + i1, err := c.GetHolder(0).CreateIndex("i1", pilosa.IndexOptions{TrackExistence: true}) + if err != nil { + t.Fatal(err) + } + defer i1.Close() + + if _, err := i1.CreateField("parentid", pilosa.OptFieldTypeInt(0, 1000)); err != nil { + t.Fatal(err) + } else if _, err := i1.CreateField("x", pilosa.OptFieldTypeInt(0, 1000)); err != nil { + t.Fatal(err) + } + + // Populate with data. + if _, err := c.GetNode(0).API.Query(context.Background(), &pilosa.QueryRequest{ + Index: "i0", + Query: ` + Set(1, a=10) + Set(2, a=20) + Set(3, a=30) + `}); err != nil { + t.Fatal(err) + } + + if _, err := c.GetNode(0).API.Query(context.Background(), &pilosa.QueryRequest{ + Index: "i1", + Query: ` + Set(1, parentid=1) + Set(1, x=100) + + Set(2, parentid=1) + Set(2, x=200) + + Set(3, parentid=2) + Set(3, x=300) + `}); err != nil { + t.Fatal(err) + } + + t.Run("Count", func(t *testing.T) { + results, columns := mustQueryRows(t, c.GetNode(0).Server, `SELECT COUNT(*) FROM i0 INNER JOIN i1 ON i0._id = i1.parentid`) + if diff := cmp.Diff([][]interface{}{ + {int64(2)}, + }, results); diff != "" { + t.Fatal(diff) + } + + if diff := cmp.Diff([]*pilosa.StmtColumn{ + {Name: "count", Type: "INT"}, + }, columns); diff != "" { + t.Fatal(diff) + } + }) + + t.Run("CountWithParentCondition", func(t *testing.T) { + results, columns := mustQueryRows(t, c.GetNode(0).Server, `SELECT COUNT(*) FROM i0 INNER JOIN i1 ON i0._id = i1.parentid WHERE i0.a = 10`) + if diff := cmp.Diff([][]interface{}{ + {int64(1)}, + }, results); diff != "" { + t.Fatal(diff) + } + + if diff := cmp.Diff([]*pilosa.StmtColumn{ + {Name: "count", Type: "INT"}, + }, columns); diff != "" { + t.Fatal(diff) + } + }) +} + func mustQueryRows(tb testing.TB, svr *pilosa.Server, q string) (results [][]interface{}, columns []*pilosa.StmtColumn) { tb.Helper() diff --git a/sql2/ast.go b/sql2/ast.go index 929d1f945..9d2e1a4db 100644 --- a/sql2/ast.go +++ b/sql2/ast.go @@ -351,6 +351,119 @@ func ExprString(expr Expr) string { return expr.String() } +// ExprTableName returns the name of the table referenced in an expression. +// Returns ok as false if more than one table referenced. Returns a blank string +// if no tables are referenced. +func ExprTableName(expr Expr) (table string, ok bool) { + switch expr := expr.(type) { + case *BindExpr, *BlobLit, *BoolLit, *Ident, *NullLit, *NumberLit, *StringLit: + return "", true + + case *BinaryExpr: + x, ok := ExprTableName(expr.X) + if !ok { + return "", false + } + + y, ok := ExprTableName(expr.Y) + if !ok { + return "", false + } + + if x == "" { + return y, true + } else if y == "" { + return x, true + } else if x == y { + return x, true + } + return "", false + + case *Call: + for _, arg := range expr.Args { + tbl, ok := ExprTableName(arg) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + } + return table, true + + case *CaseExpr: + tbl, ok := ExprTableName(expr.Operand) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + + tbl, ok = ExprTableName(expr.ElseExpr) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + + for _, blk := range expr.Blocks { + tbl, ok := ExprTableName(blk.Condition) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + + tbl, ok = ExprTableName(blk.Body) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + } + return table, true + + case *CastExpr: + return ExprTableName(expr.X) + + case *Exists: + return "", false // TODO + + case *ExprList: + for _, e := range expr.Exprs { + tbl, ok := ExprTableName(e) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + } + return table, true + + case *ParenExpr: + return ExprTableName(expr.X) + + case *QualifiedRef: + return expr.Table.Name, true + + case *Raise: + return "", true + + case *Range: + tbl, ok := ExprTableName(expr.X) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + + tbl, ok = ExprTableName(expr.Y) + if !ok || (table != "" && tbl != table) { + return "", false + } + table = tbl + return table, true + + case *UnaryExpr: + return ExprTableName(expr.X) + + default: + return "", false + } +} + // SplitExprTree splits apart expr so it is a list of all AND joined expressions. // For example, the expression "A AND B AND (C OR (D AND E))" would be split into // a list of "A", "B", "C OR (D AND E)". @@ -3001,6 +3114,23 @@ func (s *SelectStatement) IsAggregate() bool { return false } +// HasWildcard returns true any result column contains a wildcard (STAR). +func (s *SelectStatement) HasWildcard() bool { + for _, col := range s.Columns { + // Unqualified wildcard. + if col.Star.IsValid() { + return true + } + + // Table-qualified wildcard. + if ref, ok := col.Expr.(*QualifiedRef); ok && ref.Star.IsValid() { + return true + } + } + + return false +} + // String returns the string representation of the statement. func (s *SelectStatement) String() string { var buf bytes.Buffer