Implement non-aggregate SELECT query

This commit is contained in:
Ben Johnson 2021-08-30 10:04:32 -06:00
parent 1586861ed2
commit 92646cbdf7
4 changed files with 429 additions and 12 deletions

View file

@ -0,0 +1,60 @@
// Copyright 2021 Pilosa Corp.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package main
import (
"context"
"encoding/json"
"flag"
"fmt"
"os"
"strings"
"github.com/molecula/featurebase/v2/sql2"
)
func main() {
if err := run(context.Background(), os.Args[1:]); err == flag.ErrHelp {
os.Exit(1)
} else if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
}
func run(ctx context.Context, args []string) error {
fs := flag.NewFlagSet("featurebase-parse-sql", flag.ContinueOnError)
if err := fs.Parse(args); err != nil {
return err
}
q := fs.Arg(0)
if q == "" {
return fmt.Errorf("query required")
}
stmt, err := sql2.NewParser(strings.NewReader(q)).ParseStatement()
if err != nil {
return err
}
buf, err := json.MarshalIndent(stmt, "", " ")
if err != nil {
return err
}
fmt.Println(string(buf))
return nil
}

View file

@ -18,7 +18,6 @@ import (
"context"
"database/sql"
"fmt"
"io"
"strconv"
"strings"
@ -73,6 +72,7 @@ func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.S
default:
return nil, fmt.Errorf("unexpected source type in aggregate query: %T", source)
}
indexName := sql2.IdentName(source.Name)
// Convert WHERE clause.
cond, err := p.planExprPQL(ctx, stmt, stmt.WhereExpr)
@ -98,7 +98,7 @@ func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.S
callName := strings.ToUpper(sql2.IdentName(call.Name))
switch callName {
case "COUNT":
return NewCountNode(p.executor, sql2.IdentName(source.Name), cond), nil
return NewCountNode(p.executor, indexName, col.Name(), cond), nil
default:
return nil, fmt.Errorf("unsupported call in aggregate query: %T", callName)
}
@ -107,7 +107,86 @@ func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.S
}
func (p *Planner) planNonAggregateSelectStatement(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) {
return nil, fmt.Errorf("cannot plan non-aggregate SELECT query")
// Extract table name from source.
var source *sql2.QualifiedTableName
switch src := stmt.Source.(type) {
case *sql2.JoinClause:
return nil, fmt.Errorf("cannot use JOIN in non-aggregate query")
case *sql2.ParenSource:
return nil, fmt.Errorf("cannot use parenthesized source in non-aggregate query")
case *sql2.QualifiedTableName:
source = src
case *sql2.SelectStatement:
return nil, fmt.Errorf("cannot use sub-select in non-aggregate query")
default:
return nil, fmt.Errorf("unexpected source type in non-aggregate query: %T", source)
}
indexName := sql2.IdentName(source.Name)
// Lookup index.
idx := p.executor.Holder.Index(indexName)
if idx == nil {
return nil, newNotFoundError(ErrIndexNotFound, indexName)
}
// Convert WHERE clause.
cond, err := p.planExprPQL(ctx, stmt, stmt.WhereExpr)
if err != nil {
return nil, err
}
// Build column list.
var columnNames, columnAliases []string
for _, col := range stmt.Columns {
// Unqualified wildcard.
if col.Star.IsValid() {
columnNames = append(columnNames, "_id")
columnAliases = append(columnAliases, "_id")
for _, field := range idx.Fields() {
if field.Name() == "_exists" {
continue
}
columnNames = append(columnNames, field.Name())
columnAliases = append(columnAliases, field.Name())
}
continue
}
// Handle expressions and qualified references.
switch expr := col.Expr.(type) {
case *sql2.Ident:
columnNames = append(columnNames, expr.Name)
columnAliases = append(columnAliases, col.Name())
case *sql2.QualifiedRef:
if tbl := sql2.IdentName(expr.Table); tbl != "" && tbl != source.TableName() {
return nil, fmt.Errorf("no such table: %q", tbl)
}
if expr.Star.IsValid() {
columnNames = append(columnNames, "_id")
columnAliases = append(columnAliases, "_id")
for _, field := range idx.Fields() {
if field.Name() == "_exists" {
continue
}
columnNames = append(columnNames, field.Name())
columnAliases = append(columnAliases, field.Name())
}
} else {
columnNames = append(columnNames, sql2.IdentName(expr.Column))
columnAliases = append(columnAliases, sql2.IdentName(expr.Column))
}
default:
return nil, fmt.Errorf("unsupported column expression: %T", expr)
}
}
return NewExtractNode(p.executor, indexName, columnNames, columnAliases, cond), nil
}
// planExprPQL returns a PQL call tree for a given expression.
@ -334,6 +413,10 @@ func (rs *StmtRows) Err() error {
return nil
}
func (rs *StmtRows) Columns() []string {
return rs.node.Columns()
}
func (rs *StmtRows) Next() bool {
if rs.err != nil {
return false
@ -370,6 +453,8 @@ func (rs *StmtRows) Scan(dst ...interface{}) error {
*p = 0
case *uint64:
*p = 0
case *interface{}:
*p = nil
default:
return fmt.Errorf("cannot scan NULL value into %T destination at index %d", p, i)
}
@ -388,6 +473,8 @@ func (rs *StmtRows) Scan(dst ...interface{}) error {
*p = uint(v)
case *uint64:
*p = uint64(v)
case *interface{}:
*p = v
default:
return fmt.Errorf("cannot scan %T value into %T destination at index %d", v, p, i)
}
@ -438,34 +525,154 @@ type StmtNode interface {
Row() []interface{}
// Returns column definitions for the node.
// Columns() []*Column
Columns() []string
// Returns a reference to the value register for a named column.
// Lookup(table, column string) (interface{}, error)
}
var _ StmtNode = (*ExtractNode)(nil)
// ExtractNode executes an Extract() query against a FeatureBase index.
type ExtractNode struct {
executor *executor
indexName string
columns []string
aliases []string
cond *pql.Call
result []ExtractedTableColumn
row []interface{}
}
func NewExtractNode(executor *executor, indexName string, columns, aliases []string, cond *pql.Call) *ExtractNode {
if cond == nil {
cond = &pql.Call{Name: "All"}
}
// Ensure ID column is always the first column.
if len(columns) > 0 && columns[0] != "_id" {
columns = append([]string{"_id"}, columns...)
aliases = append([]string{"_id"}, aliases...)
}
// TODO: Move "_id" column to the first position if it is specified later on in column list.
return &ExtractNode{
executor: executor,
indexName: indexName,
columns: columns, // source column names
aliases: aliases, // external column alias
cond: cond,
row: make([]interface{}, len(columns)),
}
}
func (n *ExtractNode) Columns() []string {
return n.aliases
}
func (n *ExtractNode) First(ctx context.Context) error {
n.result = nil
return nil
}
func (n *ExtractNode) Next(ctx context.Context) error {
if err := n.init(ctx); err != nil {
return err
}
if len(n.result) == 0 {
for i := range n.row {
n.row[i] = nil
}
return sql.ErrNoRows
}
// Copy ID value to current row.
result := n.result[0]
if result.Column.Keyed {
n.row[0] = result.Column.Key
} else {
n.row[0] = int64(result.Column.ID)
}
// Copy values to current row.
for i, v := range result.Rows {
n.row[i+1] = v
}
n.result = n.result[1:]
return nil
}
func (n *ExtractNode) init(ctx context.Context) error {
if n.result != nil {
return nil
}
// Generate PQL query with all specified rows.
// Skip first column as it is the ID column.
call := &pql.Call{Name: "Extract", Children: []*pql.Call{n.cond}}
for _, column := range n.columns[1:] {
call.Children = append(call.Children,
&pql.Call{
Name: "Rows",
Args: map[string]interface{}{"field": column},
},
)
}
// Execute Extract() against cluster.
result, err := n.executor.Execute(ctx, n.indexName, &pql.Query{Calls: []*pql.Call{call}}, nil, nil)
if err != nil {
return err
} else if result.Err != nil {
return result.Err
} else if len(result.Results) != 1 {
return fmt.Errorf("expected single result table from Extract(), got %d results", len(result.Results))
}
// Extract out the column/row data from resultset.
tbl, ok := result.Results[0].(ExtractedTable)
if !ok {
return fmt.Errorf("unexpected Extract() result type: %T", result.Results[0])
}
n.result = tbl.Columns
return nil
}
func (n *ExtractNode) Row() []interface{} { return n.row }
var _ StmtNode = (*CountNode)(nil)
// CountNode executes a COUNT(*) against a FeatureBase index and returns a single row.
type CountNode struct {
executor *executor
indexName string
cond *pql.Call // conditional
executor *executor
indexName string
columnName string
cond *pql.Call // conditional
row []interface{}
}
func NewCountNode(executor *executor, indexName string, cond *pql.Call) *CountNode {
func NewCountNode(executor *executor, indexName string, columnName string, cond *pql.Call) *CountNode {
if cond == nil {
cond = &pql.Call{Name: "All"}
}
return &CountNode{
executor: executor,
indexName: indexName,
cond: cond,
executor: executor,
indexName: indexName,
columnName: columnName,
cond: cond,
}
}
func (n *CountNode) Columns() []string {
return []string{n.columnName}
}
func (n *CountNode) First(ctx context.Context) error {
n.row = nil
return nil
@ -473,7 +680,7 @@ func (n *CountNode) First(ctx context.Context) error {
func (n *CountNode) Next(ctx context.Context) error {
if n.row != nil {
return io.EOF
return sql.ErrNoRows
}
q := &pql.Query{

View file

@ -16,8 +16,10 @@ package pilosa_test
import (
"context"
"strings"
"testing"
"github.com/google/go-cmp/cmp"
"github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/test"
)
@ -166,3 +168,134 @@ func TestPlanner_Count(t *testing.T) {
})
})
}
func TestPlanner_Select(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)
} else if _, err := i0.CreateField("b", 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("x", pilosa.OptFieldTypeInt(0, 1000)); err != nil {
t.Fatal(err)
} else if _, err := i1.CreateField("y", 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(1, b=100)
Set(2, a=20)
Set(2, b=200)
`}); err != nil {
t.Fatal(err)
}
t.Run("UnqualifiedColumns", func(t *testing.T) {
results := mustQueryRows(t, c.GetNode(0).Server, `SELECT _id, a, b FROM i0`)
if diff := cmp.Diff(results, [][]interface{}{
{int64(1), int64(10), int64(100)},
{int64(2), int64(20), int64(200)},
}); diff != "" {
t.Fatal(diff)
}
})
t.Run("QualifiedColumns", func(t *testing.T) {
results := mustQueryRows(t, c.GetNode(0).Server, `SELECT i0._id, i0.a, i0.b FROM i0`)
if diff := cmp.Diff(results, [][]interface{}{
{int64(1), int64(10), int64(100)},
{int64(2), int64(20), int64(200)},
}); diff != "" {
t.Fatal(diff)
}
})
t.Run("UnqualifiedStar", func(t *testing.T) {
results := mustQueryRows(t, c.GetNode(0).Server, `SELECT * FROM i0`)
if diff := cmp.Diff(results, [][]interface{}{
{int64(1), int64(10), int64(100)},
{int64(2), int64(20), int64(200)},
}); diff != "" {
t.Fatal(diff)
}
})
t.Run("QualifiedStar", func(t *testing.T) {
results := mustQueryRows(t, c.GetNode(0).Server, `SELECT i0.* FROM i0`)
if diff := cmp.Diff(results, [][]interface{}{
{int64(1), int64(10), int64(100)},
{int64(2), int64(20), int64(200)},
}); diff != "" {
t.Fatal(diff)
}
})
t.Run("ErrFieldNotFound", func(t *testing.T) {
stmt, err := c.GetNode(0).Server.PlanSQL(context.Background(), `SELECT xyz FROM i0`)
if err != nil {
t.Fatal(err)
}
defer stmt.Close()
var xyz interface{}
if err := stmt.QueryRowContext(context.Background()).Scan(&xyz); err == nil || !strings.Contains(err.Error(), `xyz: field not found`) {
t.Fatalf("unexpected error: %v", err)
}
})
}
func mustQueryRows(tb testing.TB, svr *pilosa.Server, q string) [][]interface{} {
tb.Helper()
stmt, err := svr.PlanSQL(context.Background(), q)
if err != nil {
tb.Fatal(err)
}
defer stmt.Close()
rows, err := stmt.QueryContext(context.Background())
if err != nil {
tb.Fatal(err)
}
results := make([][]interface{}, 0)
for rows.Next() {
result := make([]interface{}, len(rows.Columns()))
// Create list of scan destination pointers.
dsts := make([]interface{}, len(result))
for i := range result {
dsts[i] = &result[i]
}
if err := rows.Scan(dsts...); err != nil {
tb.Fatal(err)
}
results = append(results, result)
}
if err := rows.Err(); err != nil {
tb.Fatal(err)
}
return results
}

View file

@ -3053,6 +3053,23 @@ type ResultColumn struct {
Alias *Ident // alias name
}
// Name returns the column name. Uses the alias, if specified.
// Otherwise returns a generated name.
func (c *ResultColumn) Name() string {
if c.Alias != nil {
return IdentName(c.Alias)
}
switch expr := c.Expr.(type) {
case *Ident:
return IdentName(expr)
case *QualifiedRef:
return expr.String()
default:
return ""
}
}
// IsAggregate returns true if column contains an aggregate function expression.
func (c *ResultColumn) IsAggregate() bool {
if c.Star.IsValid() {