Add AST node validation
This commit is contained in:
@@ -11,6 +11,8 @@ type (
|
||||
ASTNode interface {
|
||||
fmt.Stringer
|
||||
squirrel.Sqlizer
|
||||
|
||||
Validate() error
|
||||
}
|
||||
|
||||
ASTSet []ASTNode // Stream of comma delimited nodes
|
||||
@@ -55,20 +57,28 @@ type (
|
||||
}
|
||||
)
|
||||
|
||||
func (n String) String() string { return fmt.Sprintf("%q", n.Value) }
|
||||
func (n String) Validate() (err error) { return }
|
||||
func (n String) String() string { return fmt.Sprintf("%q", n.Value) }
|
||||
|
||||
func (n Number) String() string { return n.Value }
|
||||
func (n Number) Validate() (err error) { return }
|
||||
func (n Number) String() string { return n.Value }
|
||||
|
||||
func (n Operator) String() string { return n.Kind }
|
||||
func (n Operator) Validate() (err error) { return }
|
||||
func (n Operator) String() string { return n.Kind }
|
||||
|
||||
func (n Keyword) String() string { return n.Keyword }
|
||||
func (n Keyword) Validate() (err error) { return }
|
||||
func (n Keyword) String() string { return n.Keyword }
|
||||
|
||||
func (n Interval) String() string { return fmt.Sprintf("INTERVAL %s %s", n.Value, n.Unit) }
|
||||
func (n Interval) Validate() (err error) { return }
|
||||
func (n Interval) String() string { return fmt.Sprintf("INTERVAL %s %s", n.Value, n.Unit) }
|
||||
|
||||
func (n Function) String() string { return fmt.Sprintf("%s(%s)", n.Name, n.Arguments) }
|
||||
func (n Function) Validate() (err error) { return }
|
||||
func (n Function) String() string { return fmt.Sprintf("%s(%s)", n.Name, n.Arguments) }
|
||||
|
||||
func (n Ident) String() string { return n.Value }
|
||||
func (n Ident) Validate() (err error) { return }
|
||||
func (n Ident) String() string { return n.Value }
|
||||
|
||||
func (n Column) Validate() (err error) { return }
|
||||
func (n Column) String() (out string) {
|
||||
out = n.Expr.String()
|
||||
if n.Alias != "" {
|
||||
@@ -78,6 +88,29 @@ func (n Column) String() (out string) {
|
||||
return
|
||||
}
|
||||
|
||||
func (nn ASTNodes) Validate() (err error) {
|
||||
if err = validate(nn); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
l := len(nn)
|
||||
if l == 0 {
|
||||
return fmt.Errorf("empty set")
|
||||
}
|
||||
|
||||
if op, ok := nn[0].(Operator); ok {
|
||||
return fmt.Errorf("malformed expression, unexpected operator '%s' at first node", op)
|
||||
}
|
||||
|
||||
if l > 1 {
|
||||
if op, ok := nn[l-1].(Operator); ok {
|
||||
return fmt.Errorf("malformed expression, unexpected operator '%s' at last node", op)
|
||||
}
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func (nn ASTNodes) String() (out string) {
|
||||
for _, n := range nn {
|
||||
out = out + n.String()
|
||||
@@ -86,6 +119,10 @@ func (nn ASTNodes) String() (out string) {
|
||||
return
|
||||
}
|
||||
|
||||
func (nn ASTSet) Validate() (err error) {
|
||||
return validate(nn)
|
||||
}
|
||||
|
||||
func (nn ASTSet) String() (out string) {
|
||||
for i, n := range nn {
|
||||
if i > 0 {
|
||||
@@ -97,6 +134,7 @@ func (nn ASTSet) String() (out string) {
|
||||
return
|
||||
}
|
||||
|
||||
func (nn Columns) Validate() (err error) { return }
|
||||
func (nn Columns) String() (out string) {
|
||||
for i, n := range nn {
|
||||
if i > 0 {
|
||||
@@ -116,3 +154,17 @@ func (nn Columns) Strings() (out []string) {
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func validate(nn []ASTNode) (err error) {
|
||||
if len(nn) == 0 {
|
||||
return fmt.Errorf("empty node set")
|
||||
}
|
||||
|
||||
for _, n := range nn {
|
||||
if err = n.Validate(); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
package ql
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Ensure the parser can parse strings into Statement ASTs.
|
||||
func Test_Validators(t *testing.T) {
|
||||
var tests = []struct {
|
||||
tree ASTNode
|
||||
}{
|
||||
{
|
||||
tree: ASTNodes{
|
||||
Ident{Value: "foo"},
|
||||
Operator{Kind: "="},
|
||||
},
|
||||
},
|
||||
{
|
||||
tree: ASTNodes{
|
||||
Operator{Kind: "="},
|
||||
Ident{Value: "foo"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for i, test := range tests {
|
||||
if err := test.tree.Validate(); err == nil {
|
||||
t.Fatalf("expecting error, got nil:\n"+
|
||||
" test case: %d. %q", i, test.tree.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -64,7 +64,7 @@ func (p *Parser) ParseSet(s string) (ASTNode, error) {
|
||||
if set, err := p.parseSet(); err != nil {
|
||||
return nil, err
|
||||
} else {
|
||||
return set, nil
|
||||
return set, set.Validate()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -74,10 +74,11 @@ func (p *Parser) ParseExpression(s string) (ASTNode, error) {
|
||||
if set, err := p.parseExpr(p.nextToken()); err != nil {
|
||||
return nil, err
|
||||
} else if len(set) == 1 {
|
||||
return set[0], nil
|
||||
return set[0], set[0].Validate()
|
||||
} else {
|
||||
return set, nil
|
||||
return set, set.Validate()
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func (p *Parser) ParseColumns(s string) (columns Columns, err error) {
|
||||
@@ -104,6 +105,7 @@ next:
|
||||
goto next
|
||||
}
|
||||
|
||||
err = columns.Validate()
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -49,7 +49,7 @@ func (nn ASTSet) ToSql() (out string, args []interface{}, err error) {
|
||||
return out, args, err
|
||||
}
|
||||
|
||||
// ToSql returns column alias expression or outpit of underlaying expression's ToSql()
|
||||
// ToSql returns column alias expression or output of underlying expression's ToSql()
|
||||
func (n Column) ToSql() (string, []interface{}, error) {
|
||||
if n.Alias != "" {
|
||||
return squirrel.Alias(n.Expr, n.Alias).ToSql()
|
||||
|
||||
Reference in New Issue
Block a user