Implement basic SQL COUNT(*) query

This commit is contained in:
Ben Johnson 2021-08-12 15:43:35 -06:00
parent dac6234d0d
commit 9b8dc3d7e6
14 changed files with 13160 additions and 0 deletions

1
go.mod
View file

@ -15,6 +15,7 @@ require (
github.com/fsnotify/fsnotify v1.4.9 // indirect
github.com/glycerine/goconvey v0.0.0-20190410193231-58a59202ab31 // indirect
github.com/glycerine/idem v0.0.0-20190127113923-7a8083893311
github.com/go-test/deep v1.0.7
github.com/gogo/protobuf v1.3.2
github.com/golang/protobuf v1.3.3
github.com/google/go-cmp v0.5.5

2
go.sum
View file

@ -97,6 +97,8 @@ github.com/go-logfmt/logfmt v0.4.0/go.mod h1:3RMwSq7FuexP4Kalkev3ejPJsZTpXXBr9+V
github.com/go-ole/go-ole v1.2.4 h1:nNBDSCOigTSiarFpYE9J/KtEA1IOW4CNeqT9TQDqCxI=
github.com/go-ole/go-ole v1.2.4/go.mod h1:XCwSNxSkXRo4vlyPy93sltvi/qJq0jqQhjqQNIwKuxM=
github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY=
github.com/go-test/deep v1.0.7 h1:/VSMRlnY/JSyqxQUzQLKVMAskpY/NZKFA5j2P+0pP2M=
github.com/go-test/deep v1.0.7/go.mod h1:QV8Hv/iy04NyLBxAdO9njL0iVPN1S4d/A3NVv1V36o8=
github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ=
github.com/gogo/protobuf v1.2.1/go.mod h1:hp+jE20tsWTFYpLwKvXlhS1hjn+gTNwPg2I6zVXpSg4=
github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q=

302
planner.go Normal file
View file

@ -0,0 +1,302 @@
// 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 pilosa
import (
"context"
"database/sql"
"fmt"
"io"
"strings"
"github.com/molecula/featurebase/v2/pql"
"github.com/molecula/featurebase/v2/sql2"
)
type Planner struct {
executor *executor
}
func NewPlanner(executor *executor) *Planner {
return &Planner{executor: executor}
}
func (p *Planner) PlanStatement(ctx context.Context, stmt sql2.Statement) (*Stmt, error) {
node, err := p.planStatement(ctx, stmt)
if err != nil {
return nil, err
}
return &Stmt{node: node}, nil
}
func (p *Planner) planStatement(ctx context.Context, stmt sql2.Statement) (StmtNode, error) {
switch stmt := stmt.(type) {
case *sql2.SelectStatement:
return p.planSelectStatement(ctx, stmt)
default:
return nil, fmt.Errorf("cannot plan statement: %T", stmt)
}
}
func (p *Planner) planSelectStatement(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) {
if stmt.IsAggregate() {
return p.planAggregateSelectStatement(ctx, stmt)
}
return p.planNonAggregateSelectStatement(ctx, stmt)
}
func (p *Planner) planAggregateSelectStatement(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) {
// 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 aggregate query")
case *sql2.ParenSource:
return nil, fmt.Errorf("cannot use parenthesized source in aggregate query")
case *sql2.QualifiedTableName:
source = src
case *sql2.SelectStatement:
return nil, fmt.Errorf("cannot use sub-select in aggregate query")
default:
return nil, fmt.Errorf("unexpected source type in aggregate query: %T", source)
}
// TODO: Support multiple aggregate calls.
if len(stmt.Columns) > 1 {
return nil, fmt.Errorf("only one call allowed in aggregate query")
}
// Extract aggregate call.
col := stmt.Columns[0]
var call *sql2.Call
switch expr := col.Expr.(type) {
case *sql2.Call:
call = expr
default:
return nil, fmt.Errorf("unsupported expression in aggregate query: %T", expr)
}
callName := strings.ToUpper(sql2.IdentName(call.Name))
switch callName {
case "COUNT":
return NewCountNode(p.executor, sql2.IdentName(source.Name)), nil
default:
return nil, fmt.Errorf("unsupported call in aggregate query: %T", callName)
}
// TODO: Support HAVING
}
func (p *Planner) planNonAggregateSelectStatement(ctx context.Context, stmt *sql2.SelectStatement) (_ StmtNode, err error) {
panic("TODO: Implement non-aggregate SELECT")
}
type Stmt struct {
node StmtNode
}
func (stmt *Stmt) Close() error { return nil }
func (stmt *Stmt) QueryRowContext(ctx context.Context, args ...interface{}) *StmtRow {
rows, err := stmt.QueryContext(ctx, args...)
if err != nil {
return &StmtRow{err: err}
}
return &StmtRow{rows: rows}
}
func (stmt *Stmt) QueryContext(ctx context.Context, args ...interface{}) (*StmtRows, error) {
// TODO: Handle bind arguments.
rows := &StmtRows{
ctx: ctx,
node: stmt.node,
}
// Initialize the node.
if err := rows.node.First(ctx); err != nil {
return nil, fmt.Errorf("Query: initialize statement: %w", err)
}
return rows, nil
}
type StmtRows struct {
ctx context.Context
node StmtNode
err error
}
func (rs *StmtRows) Close() error {
return nil
}
func (rs *StmtRows) Err() error {
if rs.err != nil && rs.err != sql.ErrNoRows {
return rs.err
}
return nil
}
func (rs *StmtRows) Next() bool {
if rs.err != nil {
return false
}
if rs.err = rs.node.Next(rs.ctx); rs.err != nil {
return false
}
return true
}
func (rs *StmtRows) Scan(dst ...interface{}) error {
if rs.err != nil {
return rs.err
}
// Check len(dest) against node row length.
row := rs.node.Row()
if len(dst) != len(row) {
return fmt.Errorf("Scan(): expected %d values, received %d values", len(dst), len(row))
}
// Copy values from row to destination pointers.
for i := range dst {
// Handle null values.
// TODO: Handle double pointers.
if row[i] == nil {
switch p := dst[i].(type) {
case *int:
*p = 0
case *int64:
*p = 0
case *uint:
*p = 0
case *uint64:
*p = 0
default:
return fmt.Errorf("cannot scan NULL value into %T destination at index %d", p, i)
}
continue
}
// Copy row value to scan destination.
switch v := row[i].(type) {
case int64:
switch p := dst[i].(type) {
case *int:
*p = int(v)
case *int64:
*p = v
case *uint:
*p = uint(v)
case *uint64:
*p = uint64(v)
default:
return fmt.Errorf("cannot scan %T value into %T destination at index %d", v, p, i)
}
default:
return fmt.Errorf("unexpected %T value at index %d", v, i)
}
}
return nil
}
type StmtRow struct {
err error
rows *StmtRows
}
func (r *StmtRow) Scan(dest ...interface{}) error {
if r.err != nil {
return r.err
}
defer r.rows.Close()
if !r.rows.Next() {
if err := r.rows.Err(); err != nil {
return err
}
return sql.ErrNoRows
}
if err := r.rows.Scan(dest...); err != nil {
return err
}
return r.rows.Close()
}
func (r *StmtRow) Err() error {
return r.err
}
type StmtNode interface {
// Initializes the node to its start.
First(ctx context.Context) error
// Moves the node to the next available row. Returns sql.ErrNoRows if done.
Next(ctx context.Context) error
// Returns the current row in the node.
Row() []interface{}
// Returns column definitions for the node.
// Columns() []*Column
// Returns a reference to the value register for a named column.
// Lookup(table, column string) (interface{}, error)
}
var _ StmtNode = (*CountNode)(nil)
// CountNode executes a COUNT(*) against a FeatureBase index and returns a single row.
type CountNode struct {
executor *executor
indexName string
row []interface{}
}
func NewCountNode(executor *executor, indexName string) *CountNode {
return &CountNode{
executor: executor,
indexName: indexName,
}
}
func (n *CountNode) First(ctx context.Context) error {
n.row = nil
return nil
}
func (n *CountNode) Next(ctx context.Context) error {
if n.row != nil {
return io.EOF
}
result, err := n.executor.Execute(ctx, n.indexName, &pql.Query{
Calls: []*pql.Call{
{Name: "Count", Children: []*pql.Call{{Name: "All"}}},
},
}, nil, nil)
if err != nil {
return err
}
n.row = []interface{}{int64(result.Results[0].(uint64))}
return nil
}
func (n *CountNode) Row() []interface{} { return n.row }

77
planner_test.go Normal file
View file

@ -0,0 +1,77 @@
// 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 pilosa_test
import (
"context"
"strings"
"testing"
"github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/sql2"
"github.com/molecula/featurebase/v2/test"
)
func TestPlanner_Count(t *testing.T) {
c := test.MustRunCluster(t, 1)
defer c.Close()
index, err := c.GetHolder(0).CreateIndex("i", pilosa.IndexOptions{TrackExistence: true})
if err != nil {
t.Fatal(err)
}
defer index.Close()
if _, err := index.CreateField("f"); err != nil {
t.Fatal(err)
}
// Populate with data.
if _, err := c.GetNode(0).API.Query(context.Background(), &pilosa.QueryRequest{
Index: "i",
Query: `
Set(1, f=10)
Set(2, f=10)
Set(3, f=11)
`}); err != nil {
t.Fatal(err)
}
// Parse SQL into AST.
q := `SELECT COUNT(*) AS "count" FROM i`
st, err := sql2.NewParser(strings.NewReader(q)).ParseStatement()
if err != nil {
t.Fatal(err)
}
// Generate a prepared statement with the execution plan.
stmt, err := pilosa.NewPlanner(c.GetNode(0).Server.Executor()).PlanStatement(context.Background(), st)
if err != nil {
t.Fatal(err)
}
defer stmt.Close()
// Scan first row from result set.
var n int
if err := stmt.QueryRowContext(context.Background()).Scan(&n); err != nil {
t.Fatal(err)
} else if got, want := n, 3; got != want {
t.Fatalf("Scan()=%d, want %d", got, want)
}
if err := stmt.Close(); err != nil {
t.Fatal(err)
}
}

View file

@ -1346,6 +1346,11 @@ func (srv *Server) GetTransaction(ctx context.Context, id string, remote bool) (
return trns, nil
}
// Executor returns the executor attached to the server. For testing only.
func (s *Server) Executor() *executor {
return s.executor
}
// countOpenFiles on operating systems that support lsof.
func countOpenFiles() (int, error) {
switch runtime.GOOS {

3507
sql2/ast.go Normal file

File diff suppressed because it is too large Load diff

1161
sql2/ast_test.go Normal file

File diff suppressed because it is too large Load diff

2960
sql2/parser.go Normal file

File diff suppressed because it is too large Load diff

3311
sql2/parser_test.go Normal file

File diff suppressed because it is too large Load diff

350
sql2/scanner.go Normal file
View file

@ -0,0 +1,350 @@
// 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 sql2
import (
"bufio"
"bytes"
"io"
"unicode"
)
type Scanner struct {
r io.RuneReader
buf bytes.Buffer
ch rune
pos Pos
full bool
}
func NewScanner(r io.Reader) *Scanner {
return &Scanner{
r: bufio.NewReader(r),
pos: Pos{Offset: -1, Line: 1},
}
}
func (s *Scanner) Scan() (pos Pos, token Token, lit string) {
for {
if ch := s.peek(); ch == -1 {
return s.pos, EOF, ""
} else if unicode.IsSpace(ch) {
s.read()
continue
} else if isDigit(ch) || ch == '.' {
return s.scanNumber()
} else if ch == 'x' || ch == 'X' {
return s.scanBlob()
} else if isAlpha(ch) || ch == '_' {
return s.scanUnquotedIdent(s.pos, "")
} else if ch == '"' {
return s.scanQuotedIdent()
} else if ch == '\'' {
return s.scanString()
} else if ch == '?' || ch == ':' || ch == '@' || ch == '$' {
return s.scanBind()
}
switch ch, pos := s.read(); ch {
case ';':
return pos, SEMI, ";"
case '(':
return pos, LP, "("
case ')':
return pos, RP, ")"
case ',':
return pos, COMMA, ","
case '!':
if s.peek() == '=' {
s.read()
return pos, NE, "!="
}
return pos, BITNOT, "!"
case '=':
return pos, EQ, "="
case '<':
if s.peek() == '=' {
s.read()
return pos, LE, "<="
} else if s.peek() == '<' {
s.read()
return pos, LSHIFT, "<<"
}
return pos, LT, "<"
case '>':
if s.peek() == '=' {
s.read()
return pos, GE, ">="
} else if s.peek() == '>' {
s.read()
return pos, RSHIFT, ">>"
}
return pos, GT, ">"
case '&':
return pos, BITAND, "&"
case '|':
if s.peek() == '|' {
s.read()
return pos, CONCAT, "||"
}
return pos, BITOR, "|"
case '+':
return pos, PLUS, "+"
case '-':
return pos, MINUS, "-"
case '*':
return pos, STAR, "*"
case '/':
return pos, SLASH, "/"
case '%':
return pos, REM, "%"
default:
return pos, ILLEGAL, string(ch)
}
}
}
func (s *Scanner) scanUnquotedIdent(pos Pos, prefix string) (Pos, Token, string) {
assert(isUnquotedIdent(s.peek()))
s.buf.Reset()
s.buf.WriteString(prefix)
for ch, _ := s.read(); isUnquotedIdent(ch); ch, _ = s.read() {
s.buf.WriteRune(ch)
}
s.unread()
lit := s.buf.String()
tok := Lookup(lit)
return pos, tok, lit
}
func (s *Scanner) scanQuotedIdent() (Pos, Token, string) {
ch, pos := s.read()
assert(ch == '"')
s.buf.Reset()
for {
ch, _ := s.read()
if ch == -1 {
return pos, ILLEGAL, `"` + s.buf.String()
} else if ch == '"' {
if s.peek() == '"' { // escaped quote
s.read()
s.buf.WriteRune('"')
continue
}
return pos, QIDENT, s.buf.String()
}
s.buf.WriteRune(ch)
}
}
func (s *Scanner) scanString() (Pos, Token, string) {
ch, pos := s.read()
assert(ch == '\'')
s.buf.Reset()
for {
ch, _ := s.read()
if ch == -1 {
return pos, ILLEGAL, `'` + s.buf.String()
} else if ch == '\'' {
if s.peek() == '\'' { // escaped quote
s.read()
s.buf.WriteRune('\'')
continue
}
return pos, STRING, s.buf.String()
}
s.buf.WriteRune(ch)
}
}
func (s *Scanner) scanBind() (Pos, Token, string) {
start, pos := s.read()
s.buf.Reset()
s.buf.WriteRune(start)
// Question mark starts a numeric bind.
if start == '?' {
for isDigit(s.peek()) {
ch, _ := s.read()
s.buf.WriteRune(ch)
}
return pos, BIND, s.buf.String()
}
// All other characters start an alphanumeric bind.
assert(start == ':' || start == '@' || start == '$')
for isUnquotedIdent(s.peek()) {
ch, _ := s.read()
s.buf.WriteRune(ch)
}
return pos, BIND, s.buf.String()
}
func (s *Scanner) scanBlob() (Pos, Token, string) {
start, pos := s.read()
assert(start == 'x' || start == 'X')
// If the next character is not a quote, it's an IDENT.
if isUnquotedIdent(s.peek()) {
return s.scanUnquotedIdent(pos, string(start))
} else if s.peek() != '\'' {
return pos, IDENT, string(start)
}
ch, _ := s.read()
assert(ch == '\'')
s.buf.Reset()
for i := 0; ; i++ {
ch, _ := s.read()
if ch == '\'' {
return pos, BLOB, s.buf.String()
} else if ch == -1 {
return pos, ILLEGAL, string(start) + `'` + s.buf.String()
} else if !isHex(ch) {
return pos, ILLEGAL, string(start) + `'` + s.buf.String() + string(ch)
}
s.buf.WriteRune(ch)
}
}
func (s *Scanner) scanNumber() (Pos, Token, string) {
assert(isDigit(s.peek()) || s.peek() == '.')
pos := s.pos
tok := INTEGER
s.buf.Reset()
// Read whole number if starting with a digit.
if isDigit(s.peek()) {
for isDigit(s.peek()) {
ch, _ := s.read()
s.buf.WriteRune(ch)
}
}
// Read decimal and successive digits.
if s.peek() == '.' {
tok = FLOAT
ch, _ := s.read()
s.buf.WriteRune(ch)
for isDigit(s.peek()) {
ch, _ := s.read()
s.buf.WriteRune(ch)
}
}
// Read exponent with optional +/- sign.
if ch := s.peek(); ch == 'e' || ch == 'E' {
tok = FLOAT
ch, _ := s.read()
s.buf.WriteRune(ch)
if s.peek() == '+' || s.peek() == '-' {
ch, _ := s.read()
s.buf.WriteRune(ch)
if !isDigit(s.peek()) {
return pos, ILLEGAL, s.buf.String()
}
for isDigit(s.peek()) {
ch, _ := s.read()
s.buf.WriteRune(ch)
}
} else if isDigit(s.peek()) {
for isDigit(s.peek()) {
ch, _ := s.read()
s.buf.WriteRune(ch)
}
} else {
return pos, ILLEGAL, s.buf.String()
}
}
lit := s.buf.String()
if lit == "." {
return pos, DOT, lit
}
return pos, tok, lit
}
func (s *Scanner) read() (rune, Pos) {
if s.full {
s.full = false
return s.ch, s.pos
}
var err error
s.ch, _, err = s.r.ReadRune()
if err != nil {
s.ch = -1
return s.ch, s.pos
}
s.pos.Offset++
if s.ch == '\n' {
s.pos.Line++
s.pos.Column = 0
} else {
s.pos.Column++
}
return s.ch, s.pos
}
func (s *Scanner) peek() rune {
if !s.full {
s.read()
s.unread()
}
return s.ch
}
func (s *Scanner) unread() {
assert(!s.full)
s.full = true
}
func isDigit(ch rune) bool {
return ch >= '0' && ch <= '9'
}
func isAlpha(ch rune) bool {
return (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')
}
func isHex(ch rune) bool {
return isDigit(ch) || (ch >= 'a' && ch <= 'f') || (ch >= 'A' && ch <= 'F')
}
func isUnquotedIdent(ch rune) bool {
return isAlpha(ch) || isDigit(ch) || ch == '_'
}
// IsInteger returns true if s only contains digits.
func IsInteger(s string) bool {
for _, ch := range s {
if !isDigit(ch) {
return false
}
}
return s != ""
}

177
sql2/scanner_test.go Normal file
View file

@ -0,0 +1,177 @@
// 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 sql2_test
import (
"strings"
"testing"
sql "github.com/molecula/featurebase/v2/sql2"
)
func TestScanner_Scan(t *testing.T) {
t.Run("IDENT", func(t *testing.T) {
t.Run("Unquoted", func(t *testing.T) {
AssertScan(t, `foo_BAR123`, sql.IDENT, `foo_BAR123`)
})
t.Run("Quoted", func(t *testing.T) {
AssertScan(t, `"crazy ~!#*&# column name"" foo"`, sql.QIDENT, `crazy ~!#*&# column name" foo`)
})
t.Run("NoEndQuote", func(t *testing.T) {
AssertScan(t, `"unfinished`, sql.ILLEGAL, `"unfinished`)
})
t.Run("x", func(t *testing.T) {
AssertScan(t, `x`, sql.IDENT, `x`)
})
t.Run("StartingX", func(t *testing.T) {
AssertScan(t, `xyz`, sql.IDENT, `xyz`)
})
})
t.Run("KEYWORD", func(t *testing.T) {
AssertScan(t, `BEGIN`, sql.BEGIN, `BEGIN`)
})
t.Run("STRING", func(t *testing.T) {
t.Run("OK", func(t *testing.T) {
AssertScan(t, `'this is ''a'' string'`, sql.STRING, `this is 'a' string`)
})
t.Run("NoEndQuote", func(t *testing.T) {
AssertScan(t, `'unfinished`, sql.ILLEGAL, `'unfinished`)
})
})
t.Run("BLOB", func(t *testing.T) {
t.Run("LowerX", func(t *testing.T) {
AssertScan(t, `x'0123456789abcdef'`, sql.BLOB, `0123456789abcdef`)
})
t.Run("UpperX", func(t *testing.T) {
AssertScan(t, `X'0123456789ABCDEF'`, sql.BLOB, `0123456789ABCDEF`)
})
t.Run("NoEndQuote", func(t *testing.T) {
AssertScan(t, `x'0123`, sql.ILLEGAL, `x'0123`)
})
t.Run("BadHex", func(t *testing.T) {
AssertScan(t, `x'hello`, sql.ILLEGAL, `x'h`)
})
})
t.Run("INTEGER", func(t *testing.T) {
AssertScan(t, `123`, sql.INTEGER, `123`)
})
t.Run("FLOAT", func(t *testing.T) {
AssertScan(t, `123.456`, sql.FLOAT, `123.456`)
AssertScan(t, `.1`, sql.FLOAT, `.1`)
AssertScan(t, `123e456`, sql.FLOAT, `123e456`)
AssertScan(t, `123E456`, sql.FLOAT, `123E456`)
AssertScan(t, `123.456E78`, sql.FLOAT, `123.456E78`)
AssertScan(t, `123.E45`, sql.FLOAT, `123.E45`)
AssertScan(t, `123E+4`, sql.FLOAT, `123E+4`)
AssertScan(t, `123E-4`, sql.FLOAT, `123E-4`)
AssertScan(t, `123E`, sql.ILLEGAL, `123E`)
AssertScan(t, `123E+`, sql.ILLEGAL, `123E+`)
AssertScan(t, `123E-`, sql.ILLEGAL, `123E-`)
})
t.Run("BIND", func(t *testing.T) {
AssertScan(t, `?'`, sql.BIND, `?`)
AssertScan(t, `?123'`, sql.BIND, `?123`)
AssertScan(t, `:foo_bar123'`, sql.BIND, `:foo_bar123`)
AssertScan(t, `@bar'`, sql.BIND, `@bar`)
AssertScan(t, `$baz'`, sql.BIND, `$baz`)
})
t.Run("EOF", func(t *testing.T) {
AssertScan(t, " \n\t\r", sql.EOF, ``)
})
t.Run("SEMI", func(t *testing.T) {
AssertScan(t, ";", sql.SEMI, ";")
})
t.Run("LP", func(t *testing.T) {
AssertScan(t, "(", sql.LP, "(")
})
t.Run("RP", func(t *testing.T) {
AssertScan(t, ")", sql.RP, ")")
})
t.Run("COMMA", func(t *testing.T) {
AssertScan(t, ",", sql.COMMA, ",")
})
t.Run("NE", func(t *testing.T) {
AssertScan(t, "!=", sql.NE, "!=")
})
t.Run("BITNOT", func(t *testing.T) {
AssertScan(t, "!", sql.BITNOT, "!")
})
t.Run("EQ", func(t *testing.T) {
AssertScan(t, "=", sql.EQ, "=")
})
t.Run("LE", func(t *testing.T) {
AssertScan(t, "<=", sql.LE, "<=")
})
t.Run("LSHIFT", func(t *testing.T) {
AssertScan(t, "<<", sql.LSHIFT, "<<")
})
t.Run("LT", func(t *testing.T) {
AssertScan(t, "<", sql.LT, "<")
})
t.Run("GE", func(t *testing.T) {
AssertScan(t, ">=", sql.GE, ">=")
})
t.Run("RSHIFT", func(t *testing.T) {
AssertScan(t, ">>", sql.RSHIFT, ">>")
})
t.Run("GT", func(t *testing.T) {
AssertScan(t, ">", sql.GT, ">")
})
t.Run("BITAND", func(t *testing.T) {
AssertScan(t, "&", sql.BITAND, "&")
})
t.Run("CONCAT", func(t *testing.T) {
AssertScan(t, "||", sql.CONCAT, "||")
})
t.Run("BITOR", func(t *testing.T) {
AssertScan(t, "|", sql.BITOR, "|")
})
t.Run("PLUS", func(t *testing.T) {
AssertScan(t, "+", sql.PLUS, "+")
})
t.Run("MINUS", func(t *testing.T) {
AssertScan(t, "-", sql.MINUS, "-")
})
t.Run("STAR", func(t *testing.T) {
AssertScan(t, "*", sql.STAR, "*")
})
t.Run("SLASH", func(t *testing.T) {
AssertScan(t, "/", sql.SLASH, "/")
})
t.Run("REM", func(t *testing.T) {
AssertScan(t, "%", sql.REM, "%")
})
t.Run("DOT", func(t *testing.T) {
AssertScan(t, ".", sql.DOT, ".")
})
t.Run("ILLEGAL", func(t *testing.T) {
AssertScan(t, "^", sql.ILLEGAL, "^")
})
}
// AssertScan asserts the value of the first scan to s.
func AssertScan(tb testing.TB, s string, expectedTok sql.Token, expectedLit string) {
tb.Helper()
_, tok, lit := sql.NewScanner(strings.NewReader(s)).Scan()
if tok != expectedTok || lit != expectedLit {
tb.Fatalf("Scan(%q)=<%s,%s>, want <%s,%s>", s, tok, lit, expectedTok, expectedLit)
}
}

556
sql2/token.go Normal file
View file

@ -0,0 +1,556 @@
// 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 sql2
import (
"fmt"
"strconv"
"strings"
)
var keywords map[string]Token
func init() {
keywords = make(map[string]Token)
for i := keyword_beg + 1; i < keyword_end; i++ {
keywords[tokens[i]] = i
}
keywords[tokens[NULL]] = NULL
keywords[tokens[TRUE]] = TRUE
keywords[tokens[FALSE]] = FALSE
}
// Token is the set of lexical tokens of the Go programming language.
type Token int
// The list of tokens.
const (
// Special tokens
ILLEGAL Token = iota
EOF
COMMENT
SPACE
literal_beg
IDENT // IDENT
QIDENT // "IDENT"
STRING // 'string'
BLOB // ???
FLOAT // 123.45
INTEGER // 123
NULL // NULL
TRUE // true
FALSE // false
BIND //? or ?NNN or :VVV or @VVV or $VVV
literal_end
operator_beg
SEMI // ;
LP // (
RP // )
COMMA // ,
NE // !=
EQ // =
LE // <=
LT // <
GT // >
GE // >=
BITAND // &
BITOR // |
BITNOT // !
LSHIFT // <<
RSHIFT // >>
PLUS // +
MINUS // -
STAR // *
SLASH // /
REM // %
CONCAT // ||
DOT // .
operator_end
keyword_beg
ABORT
ACTION
ADD
AFTER
AGG_COLUMN
AGG_FUNCTION
ALL
ALTER
ANALYZE
AND
AS
ASC
ASTERISK
ATTACH
AUTOINCREMENT
BEFORE
BEGIN
BETWEEN
BY
CASCADE
CASE
CAST
CHECK
COLUMN
COLUMNKW
COMMIT
CONFLICT
CONSTRAINT
CREATE
CROSS
CTIME_KW
CURRENT
CURRENT_TIME
CURRENT_DATE
CURRENT_TIMESTAMP
DATABASE
DEFAULT
DEFERRABLE
DEFERRED
DELETE
DESC
DETACH
DISTINCT
DO
DROP
EACH
ELSE
END
ESCAPE
EXCEPT
EXCLUDE
EXCLUSIVE
EXISTS
EXPLAIN
FAIL
FILTER
FIRST
FOLLOWING
FOR
FOREIGN
FROM
FUNCTION
GLOB
GROUP
GROUPS
HAVING
IF
IF_NULL_ROW
IGNORE
IMMEDIATE
IN
INDEX
INDEXED
INITIALLY
INNER
INSERT
INSTEAD
INTERSECT
INTO
IS
ISNOT
ISNULL // TODO: REMOVE?
JOIN
KEY
LAST
LEFT
LIKE
LIMIT
MATCH
NATURAL
NO
NOT
NOTBETWEEN
NOTEXISTS
NOTGLOB
NOTHING
NOTIN
NOTLIKE
NOTMATCH
NOTNULL
NOTREGEXP
NULLS
OF
OFFSET
ON
OR
ORDER
OTHERS
OUTER
OVER
PARTITION
PLAN
PRAGMA
PRECEDING
PRIMARY
QUERY
RAISE
RANGE
RECURSIVE
REFERENCES
REGEXP
REGISTER
REINDEX
RELEASE
RENAME
REPLACE
RESTRICT
ROLLBACK
ROW
ROWS
SAVEPOINT
SELECT
SELECT_COLUMN
SET
SPAN
TABLE
TEMP
THEN
TIES
TO
TRANSACTION
TRIGGER
TRUTH
UNBOUNDED
UNION
UNIQUE
UPDATE
USING
VACUUM
VALUES
VARIABLE
VECTOR
VIEW
VIRTUAL
WHEN
WHERE
WINDOW
WITH
WITHOUT
keyword_end
ANY // ???
)
var tokens = [...]string{
ILLEGAL: "ILLEGAL",
EOF: "EOF",
COMMENT: "COMMENT",
SPACE: "SPACE",
IDENT: "IDENT",
QIDENT: "QIDENT",
STRING: "STRING",
BLOB: "BLOB",
FLOAT: "FLOAT",
INTEGER: "INTEGER",
NULL: "NULL",
TRUE: "TRUE",
FALSE: "FALSE",
BIND: "BIND",
SEMI: ";",
LP: "(",
RP: ")",
COMMA: ",",
NE: "!=",
EQ: "=",
LE: "<=",
LT: "<",
GT: ">",
GE: ">=",
BITAND: "&",
BITOR: "|",
BITNOT: "!",
LSHIFT: "<<",
RSHIFT: ">>",
PLUS: "+",
MINUS: "-",
STAR: "*",
SLASH: "/",
REM: "%",
CONCAT: "||",
DOT: ".",
ABORT: "ABORT",
ACTION: "ACTION",
ADD: "ADD",
AFTER: "AFTER",
AGG_COLUMN: "AGG_COLUMN",
AGG_FUNCTION: "AGG_FUNCTION",
ALL: "ALL",
ALTER: "ALTER",
ANALYZE: "ANALYZE",
AND: "AND",
AS: "AS",
ASC: "ASC",
ASTERISK: "ASTERISK",
ATTACH: "ATTACH",
AUTOINCREMENT: "AUTOINCREMENT",
BEFORE: "BEFORE",
BEGIN: "BEGIN",
BETWEEN: "BETWEEN",
BY: "BY",
CASCADE: "CASCADE",
CASE: "CASE",
CAST: "CAST",
CHECK: "CHECK",
COLUMN: "COLUMN",
COLUMNKW: "COLUMNKW",
COMMIT: "COMMIT",
CONFLICT: "CONFLICT",
CONSTRAINT: "CONSTRAINT",
CREATE: "CREATE",
CROSS: "CROSS",
CTIME_KW: "CTIME_KW",
CURRENT: "CURRENT",
CURRENT_TIME: "CURRENT_TIME",
CURRENT_DATE: "CURRENT_DATE",
CURRENT_TIMESTAMP: "CURRENT_TIMESTAMP",
DATABASE: "DATABASE",
DEFAULT: "DEFAULT",
DEFERRABLE: "DEFERRABLE",
DEFERRED: "DEFERRED",
DELETE: "DELETE",
DESC: "DESC",
DETACH: "DETACH",
DISTINCT: "DISTINCT",
DO: "DO",
DROP: "DROP",
EACH: "EACH",
ELSE: "ELSE",
END: "END",
ESCAPE: "ESCAPE",
EXCEPT: "EXCEPT",
EXCLUDE: "EXCLUDE",
EXCLUSIVE: "EXCLUSIVE",
EXISTS: "EXISTS",
EXPLAIN: "EXPLAIN",
FAIL: "FAIL",
FILTER: "FILTER",
FIRST: "FIRST",
FOLLOWING: "FOLLOWING",
FOR: "FOR",
FOREIGN: "FOREIGN",
FROM: "FROM",
FUNCTION: "FUNCTION",
GLOB: "GLOB",
GROUP: "GROUP",
GROUPS: "GROUPS",
HAVING: "HAVING",
IF: "IF",
IF_NULL_ROW: "IF_NULL_ROW",
IGNORE: "IGNORE",
IMMEDIATE: "IMMEDIATE",
IN: "IN",
INDEX: "INDEX",
INDEXED: "INDEXED",
INITIALLY: "INITIALLY",
INNER: "INNER",
INSERT: "INSERT",
INSTEAD: "INSTEAD",
INTERSECT: "INTERSECT",
INTO: "INTO",
IS: "IS",
ISNOT: "ISNOT",
ISNULL: "ISNULL",
JOIN: "JOIN",
KEY: "KEY",
LAST: "LAST",
LEFT: "LEFT",
LIKE: "LIKE",
LIMIT: "LIMIT",
MATCH: "MATCH",
NO: "NO",
NATURAL: "NATURAL",
NOT: "NOT",
NOTBETWEEN: "NOTBETWEEN",
NOTEXISTS: "NOTEXISTS",
NOTGLOB: "NOTGLOB",
NOTHING: "NOTHING",
NOTIN: "NOTIN",
NOTLIKE: "NOTLIKE",
NOTMATCH: "NOTMATCH",
NOTNULL: "NOTNULL",
NOTREGEXP: "NOTREGEXP",
NULLS: "NULLS",
OF: "OF",
OFFSET: "OFFSET",
ON: "ON",
OR: "OR",
ORDER: "ORDER",
OTHERS: "OTHERS",
OUTER: "OUTER",
OVER: "OVER",
PARTITION: "PARTITION",
PLAN: "PLAN",
PRAGMA: "PRAGMA",
PRECEDING: "PRECEDING",
PRIMARY: "PRIMARY",
QUERY: "QUERY",
RAISE: "RAISE",
RANGE: "RANGE",
RECURSIVE: "RECURSIVE",
REFERENCES: "REFERENCES",
REGEXP: "REGEXP",
REGISTER: "REGISTER",
REINDEX: "REINDEX",
RELEASE: "RELEASE",
RENAME: "RENAME",
REPLACE: "REPLACE",
RESTRICT: "RESTRICT",
ROLLBACK: "ROLLBACK",
ROW: "ROW",
ROWS: "ROWS",
SAVEPOINT: "SAVEPOINT",
SELECT: "SELECT",
SELECT_COLUMN: "SELECT_COLUMN",
SET: "SET",
SPAN: "SPAN",
TABLE: "TABLE",
TEMP: "TEMP",
THEN: "THEN",
TIES: "TIES",
TO: "TO",
TRANSACTION: "TRANSACTION",
TRIGGER: "TRIGGER",
TRUTH: "TRUTH",
UNBOUNDED: "UNBOUNDED",
UNION: "UNION",
UNIQUE: "UNIQUE",
UPDATE: "UPDATE",
USING: "USING",
VACUUM: "VACUUM",
VALUES: "VALUES",
VARIABLE: "VARIABLE",
VECTOR: "VECTOR",
VIEW: "VIEW",
VIRTUAL: "VIRTUAL",
WHEN: "WHEN",
WHERE: "WHERE",
WINDOW: "WINDOW",
WITH: "WITH",
WITHOUT: "WITHOUT",
}
func (tok Token) String() string {
s := ""
if 0 <= tok && tok < Token(len(tokens)) {
s = tokens[tok]
}
if s == "" {
s = "token(" + strconv.Itoa(int(tok)) + ")"
}
return s
}
func Lookup(ident string) Token {
if tok, ok := keywords[strings.ToUpper(ident)]; ok {
return tok
}
return IDENT
}
func (tok Token) IsLiteral() bool {
return tok > literal_beg && tok < literal_end
}
func (tok Token) IsOperator() bool {
return tok > operator_beg && tok < operator_end
}
func (tok Token) IsKeyword() bool {
return tok > keyword_beg && tok < keyword_end
}
func (tok Token) IsBinaryOp() bool {
switch tok {
case PLUS, MINUS, STAR, SLASH, REM, CONCAT, NOT, BETWEEN,
LSHIFT, RSHIFT, BITAND, BITOR, LT, LE, GT, GE, EQ, NE,
IS, IN, LIKE, GLOB, MATCH, REGEXP, AND, OR:
return true
default:
return false
}
}
func isIdentToken(tok Token) bool {
return tok == IDENT || tok == QIDENT
}
const (
LowestPrec = 0 // non-operators
UnaryPrec = 13
HighestPrec = 14
)
func (op Token) Precedence() int {
switch op {
case OR:
return 1
case AND:
return 2
case NOT:
return 3
case IS, MATCH, LIKE, GLOB, REGEXP, BETWEEN, IN, ISNULL, NOTNULL, NE, EQ:
return 4
case GT, LE, LT, GE:
return 5
case ESCAPE:
return 6
case BITAND, BITOR, LSHIFT, RSHIFT:
return 7
case PLUS, MINUS:
return 8
case STAR, SLASH, REM:
return 9
case CONCAT:
return 10
case BITNOT:
return 11
}
return LowestPrec
}
type Pos struct {
Offset int // offset, starting at 0
Line int // line number, starting at 1
Column int // column number, starting at 1 (byte count)
}
// String returns a string representation of the position.
func (p Pos) String() string {
if !p.IsValid() {
return "-"
}
s := fmt.Sprintf("%d", p.Line)
if p.Column != 0 {
s += fmt.Sprintf(":%d", p.Column)
}
return s
}
// IsValid returns true if p is non-zero.
func (p Pos) IsValid() bool {
return p != Pos{}
}
func assert(condition bool) {
if !condition {
panic("assert failed")
}
}

27
sql2/token_test.go Normal file
View file

@ -0,0 +1,27 @@
// 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 sql2_test
import (
"testing"
sql "github.com/molecula/featurebase/v2/sql2"
)
func TestPos_String(t *testing.T) {
if got, want := (sql.Pos{}).String(), `-`; got != want {
t.Fatalf("String()=%q, want %q", got, want)
}
}

724
sql2/walk.go Normal file
View file

@ -0,0 +1,724 @@
// 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 sql2
// A Visitor's Visit method is invoked for each node encountered by Walk.
// If the result visitor w is not nil, Walk visits each of the children
// of node with the visitor w, followed by a call of w.Visit(nil).
type Visitor interface {
Visit(node Node) (w Visitor, err error)
VisitEnd(node Node) error
}
// Walk traverses an AST in depth-first order: It starts by calling
// v.Visit(node); node must not be nil. If the visitor w returned by
// v.Visit(node) is not nil, Walk is invoked recursively with visitor
// w for each of the non-nil children of node, followed by a call of
// w.Visit(nil).
func Walk(v Visitor, node Node) error {
return walk(v, node)
}
func walk(v Visitor, node Node) (err error) {
// Visit the node itself
if v, err = v.Visit(node); err != nil {
return err
} else if v == nil {
return nil
}
// Visit node's children.
switch n := node.(type) {
case *Assignment:
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
if err := walkExpr(v, n.Expr); err != nil {
return err
}
case *ExplainStatement:
if n.Stmt != nil {
if err := walk(v, n.Stmt); err != nil {
return err
}
}
case *RollbackStatement:
if err := walkIdent(v, n.SavepointName); err != nil {
return err
}
case *SavepointStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *ReleaseStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *CreateTableStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkColumnDefinitionList(v, n.Columns); err != nil {
return err
}
if err := walkConstraintList(v, n.Constraints); err != nil {
return err
}
if n.Select != nil {
if err := walk(v, n.Select); err != nil {
return err
}
}
case *AlterTableStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdent(v, n.NewName); err != nil {
return err
}
if err := walkIdent(v, n.ColumnName); err != nil {
return err
}
if err := walkIdent(v, n.NewColumnName); err != nil {
return err
}
if n.ColumnDef != nil {
if err := walk(v, n.ColumnDef); err != nil {
return err
}
}
case *AnalyzeStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *CreateViewStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
if n.Select != nil {
if err := walk(v, n.Select); err != nil {
return err
}
}
case *DropTableStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *DropViewStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *DropIndexStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *DropTriggerStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *CreateIndexStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdent(v, n.Table); err != nil {
return err
}
if err := walkIndexedColumnList(v, n.Columns); err != nil {
return err
}
if err := walkExpr(v, n.WhereExpr); err != nil {
return err
}
case *CreateTriggerStatement:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdentList(v, n.UpdateOfColumns); err != nil {
return err
}
if err := walkIdent(v, n.Table); err != nil {
return err
}
if err := walkExpr(v, n.WhenExpr); err != nil {
return err
}
for _, x := range n.Body {
if err := walk(v, x); err != nil {
return err
}
}
case *SelectStatement:
if n.WithClause != nil {
if err := walk(v, n.WithClause); err != nil {
return err
}
}
for _, x := range n.ValueLists {
if err := walk(v, x); err != nil {
return err
}
}
for _, x := range n.Columns {
if err := walk(v, x); err != nil {
return err
}
}
if n.Source != nil {
if err := walk(v, n.Source); err != nil {
return err
}
}
if err := walkExpr(v, n.WhereExpr); err != nil {
return err
}
if err := walkExprList(v, n.GroupByExprs); err != nil {
return err
}
if err := walkExpr(v, n.HavingExpr); err != nil {
return err
}
for _, x := range n.Windows {
if err := walk(v, x); err != nil {
return err
}
}
if n.Compound != nil {
if err := walk(v, n.Compound); err != nil {
return err
}
}
for _, x := range n.OrderingTerms {
if err := walk(v, x); err != nil {
return err
}
}
if err := walkExpr(v, n.LimitExpr); err != nil {
return err
}
if err := walkExpr(v, n.OffsetExpr); err != nil {
return err
}
case *InsertStatement:
if n.WithClause != nil {
if err := walk(v, n.WithClause); err != nil {
return err
}
}
if err := walkIdent(v, n.Table); err != nil {
return err
}
if err := walkIdent(v, n.Alias); err != nil {
return err
}
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
for _, x := range n.ValueLists {
if err := walk(v, x); err != nil {
return err
}
}
if n.Select != nil {
if err := walk(v, n.Select); err != nil {
return err
}
}
if n.UpsertClause != nil {
if err := walk(v, n.UpsertClause); err != nil {
return err
}
}
case *UpdateStatement:
if n.WithClause != nil {
if err := walk(v, n.WithClause); err != nil {
return err
}
}
if n.Table != nil {
if err := walk(v, n.Table); err != nil {
return err
}
}
for _, x := range n.Assignments {
if err := walk(v, x); err != nil {
return err
}
}
if err := walkExpr(v, n.WhereExpr); err != nil {
return err
}
case *UpsertClause:
if err := walkIndexedColumnList(v, n.Columns); err != nil {
return err
}
if err := walkExpr(v, n.WhereExpr); err != nil {
return err
}
for _, x := range n.Assignments {
if err := walk(v, x); err != nil {
return err
}
}
if err := walkExpr(v, n.UpdateWhereExpr); err != nil {
return err
}
case *DeleteStatement:
if n.WithClause != nil {
if err := walk(v, n.WithClause); err != nil {
return err
}
}
if n.Table != nil {
if err := walk(v, n.Table); err != nil {
return err
}
}
if err := walkExpr(v, n.WhereExpr); err != nil {
return err
}
for _, x := range n.OrderingTerms {
if err := walk(v, x); err != nil {
return err
}
}
if err := walkExpr(v, n.LimitExpr); err != nil {
return err
}
if err := walkExpr(v, n.OffsetExpr); err != nil {
return err
}
case *PrimaryKeyConstraint:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
case *NotNullConstraint:
if err := walkIdent(v, n.Name); err != nil {
return err
}
case *UniqueConstraint:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
case *CheckConstraint:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkExpr(v, n.Expr); err != nil {
return err
}
case *DefaultConstraint:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkExpr(v, n.Expr); err != nil {
return err
}
case *ForeignKeyConstraint:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
if err := walkIdent(v, n.ForeignTable); err != nil {
return err
}
if err := walkIdentList(v, n.ForeignColumns); err != nil {
return err
}
for _, x := range n.Args {
if err := walk(v, x); err != nil {
return err
}
}
case *ParenExpr:
if err := walkExpr(v, n.X); err != nil {
return err
}
case *UnaryExpr:
if err := walkExpr(v, n.X); err != nil {
return err
}
case *BinaryExpr:
if err := walkExpr(v, n.X); err != nil {
return err
}
if err := walkExpr(v, n.Y); err != nil {
return err
}
case *CastExpr:
if err := walkExpr(v, n.X); err != nil {
return err
}
if n.Type != nil {
if err := walk(v, n.Type); err != nil {
return err
}
}
case *CaseBlock:
if err := walkExpr(v, n.Condition); err != nil {
return err
}
if err := walkExpr(v, n.Body); err != nil {
return err
}
case *CaseExpr:
if err := walkExpr(v, n.Operand); err != nil {
return err
}
for _, x := range n.Blocks {
if err := walk(v, x); err != nil {
return err
}
}
if err := walkExpr(v, n.ElseExpr); err != nil {
return err
}
case *ExprList:
if err := walkExprList(v, n.Exprs); err != nil {
return err
}
case *QualifiedRef:
if err := walkIdent(v, n.Table); err != nil {
return err
}
if err := walkIdent(v, n.Column); err != nil {
return err
}
case *Call:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkExprList(v, n.Args); err != nil {
return err
}
if n.Filter != nil {
if err := walk(v, n.Filter); err != nil {
return err
}
}
if n.Over != nil {
if err := walk(v, n.Over); err != nil {
return err
}
}
case *FilterClause:
if err := walkExpr(v, n.X); err != nil {
return err
}
case *OverClause:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if n.Definition != nil {
if err := walk(v, n.Definition); err != nil {
return err
}
}
case *OrderingTerm:
if err := walkExpr(v, n.X); err != nil {
return err
}
case *FrameSpec:
if err := walkExpr(v, n.X); err != nil {
return err
}
if err := walkExpr(v, n.Y); err != nil {
return err
}
case *Range:
if err := walkExpr(v, n.X); err != nil {
return err
}
if err := walkExpr(v, n.Y); err != nil {
return err
}
case *Raise:
if n.Error != nil {
if err := walk(v, n.Error); err != nil {
return err
}
}
case *Exists:
if n.Select != nil {
if err := walk(v, n.Select); err != nil {
return err
}
}
case *ParenSource:
if n.X != nil {
if err := walk(v, n.X); err != nil {
return err
}
}
if err := walkIdent(v, n.Alias); err != nil {
return err
}
case *QualifiedTableName:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if err := walkIdent(v, n.Alias); err != nil {
return err
}
if err := walkIdent(v, n.Index); err != nil {
return err
}
case *JoinClause:
if n.X != nil {
if err := walk(v, n.X); err != nil {
return err
}
}
if n.Operator != nil {
if err := walk(v, n.Operator); err != nil {
return err
}
}
if n.Y != nil {
if err := walk(v, n.Y); err != nil {
return err
}
}
if n.Constraint != nil {
if err := walk(v, n.Constraint); err != nil {
return err
}
}
case *OnConstraint:
if err := walkExpr(v, n.X); err != nil {
return err
}
case *UsingConstraint:
if err := walkIdentList(v, n.Columns); err != nil {
return err
}
case *ColumnDefinition:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if n.Type != nil {
if err := walk(v, n.Type); err != nil {
return err
}
}
if err := walkConstraintList(v, n.Constraints); err != nil {
return err
}
case *ResultColumn:
if err := walkExpr(v, n.Expr); err != nil {
return err
}
if err := walkIdent(v, n.Alias); err != nil {
return err
}
case *IndexedColumn:
if err := walkExpr(v, n.X); err != nil {
return err
}
case *Window:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if n.Definition != nil {
if err := walk(v, n.Definition); err != nil {
return err
}
}
case *WindowDefinition:
if err := walkIdent(v, n.Base); err != nil {
return err
}
if err := walkExprList(v, n.Partitions); err != nil {
return err
}
for _, x := range n.OrderingTerms {
if err := walk(v, x); err != nil {
return err
}
}
if n.Frame != nil {
if err := walk(v, n.Frame); err != nil {
return err
}
}
case *Type:
if err := walkIdent(v, n.Name); err != nil {
return err
}
if n.Precision != nil {
if err := walk(v, n.Precision); err != nil {
return err
}
}
if n.Scale != nil {
if err := walk(v, n.Scale); err != nil {
return err
}
}
}
// Revisit original node after its children have been processed.
return v.VisitEnd(node)
}
// VisitFunc represents a function type that implements Visitor.
// Only executes on node entry.
type VisitFunc func(Node) error
// Visit executes fn. Walk visits node children if fn returns true.
func (fn VisitFunc) Visit(node Node) (Visitor, error) {
if err := fn(node); err != nil {
return nil, err
}
return fn, nil
}
// VisitEnd is a no-op.
func (fn VisitFunc) VisitEnd(node Node) error { return nil }
// VisitEndFunc represents a function type that implements Visitor.
// Only executes on node exit.
type VisitEndFunc func(Node) error
// Visit is a no-op.
func (fn VisitEndFunc) Visit(node Node) (Visitor, error) { return fn, nil }
// VisitEnd executes fn.
func (fn VisitEndFunc) VisitEnd(node Node) error { return fn(node) }
func walkIdent(v Visitor, x *Ident) error {
if x != nil {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}
func walkIdentList(v Visitor, a []*Ident) error {
for _, x := range a {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}
func walkExpr(v Visitor, x Expr) error {
if x != nil {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}
func walkExprList(v Visitor, a []Expr) error {
for _, x := range a {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}
func walkConstraintList(v Visitor, a []Constraint) error {
for _, x := range a {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}
func walkIndexedColumnList(v Visitor, a []*IndexedColumn) error {
for _, x := range a {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}
func walkColumnDefinitionList(v Visitor, a []*ColumnDefinition) error {
for _, x := range a {
if err := walk(v, x); err != nil {
return err
}
}
return nil
}