From c52c2ae8f7d8de1e5fa657c9a5d4405a8c5a83ce Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Wed, 26 Jun 2019 23:01:51 +0200 Subject: [PATCH] Add AST node validation --- compose/internal/repository/ql/ast_nodes.go | 66 +++++++++++++++++-- .../internal/repository/ql/ast_nodes_test.go | 32 +++++++++ compose/internal/repository/ql/ast_parser.go | 8 ++- compose/internal/repository/ql/squirrel.go | 2 +- 4 files changed, 97 insertions(+), 11 deletions(-) create mode 100644 compose/internal/repository/ql/ast_nodes_test.go diff --git a/compose/internal/repository/ql/ast_nodes.go b/compose/internal/repository/ql/ast_nodes.go index b467d337c..a1d47aac8 100644 --- a/compose/internal/repository/ql/ast_nodes.go +++ b/compose/internal/repository/ql/ast_nodes.go @@ -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 +} diff --git a/compose/internal/repository/ql/ast_nodes_test.go b/compose/internal/repository/ql/ast_nodes_test.go new file mode 100644 index 000000000..39f68878f --- /dev/null +++ b/compose/internal/repository/ql/ast_nodes_test.go @@ -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()) + } + } +} diff --git a/compose/internal/repository/ql/ast_parser.go b/compose/internal/repository/ql/ast_parser.go index 50310ae3b..81b499cd7 100644 --- a/compose/internal/repository/ql/ast_parser.go +++ b/compose/internal/repository/ql/ast_parser.go @@ -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 } diff --git a/compose/internal/repository/ql/squirrel.go b/compose/internal/repository/ql/squirrel.go index 0503b3dc5..3d7c5a82c 100644 --- a/compose/internal/repository/ql/squirrel.go +++ b/compose/internal/repository/ql/squirrel.go @@ -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()