Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 66 additions & 11 deletions parser/ast.go
Original file line number Diff line number Diff line change
Expand Up @@ -332,6 +332,11 @@ func (a *AlterTableDropPartition) Accept(visitor ASTVisitor) error {
if err := a.Partition.Accept(visitor); err != nil {
return err
}
if a.Settings != nil {
if err := a.Settings.Accept(visitor); err != nil {
return err
}
}
return visitor.VisitAlterTableDropPartition(a)
}

Expand Down Expand Up @@ -527,6 +532,11 @@ func (p *ProjectionOrderByClause) End() Pos {
func (p *ProjectionOrderByClause) Accept(visitor ASTVisitor) error {
visitor.Enter(p)
defer visitor.Leave(p)
if p.Columns != nil {
if err := p.Columns.Accept(visitor); err != nil {
return err
}
}
return visitor.VisitProjectionOrderBy(p)
}

Expand Down Expand Up @@ -1238,6 +1248,11 @@ func (u *UUID) End() Pos {
func (u *UUID) Accept(visitor ASTVisitor) error {
visitor.Enter(u)
defer visitor.Leave(u)
if u.Value != nil {
if err := u.Value.Accept(visitor); err != nil {
return err
}
}
return visitor.VisitUUID(u)
}

Expand Down Expand Up @@ -1266,6 +1281,11 @@ func (c *CreateDatabase) Type() string {
func (c *CreateDatabase) Accept(visitor ASTVisitor) error {
visitor.Enter(c)
defer visitor.Leave(c)
if c.Name != nil {
if err := c.Name.Accept(visitor); err != nil {
return err
}
}
if c.OnCluster != nil {
if err := c.OnCluster.Accept(visitor); err != nil {
return err
Expand All @@ -1276,6 +1296,11 @@ func (c *CreateDatabase) Accept(visitor ASTVisitor) error {
return err
}
}
if c.Comment != nil {
if err := c.Comment.Accept(visitor); err != nil {
return err
}
}
return visitor.VisitCreateDatabase(c)
}

Expand Down Expand Up @@ -1343,6 +1368,11 @@ func (c *CreateTable) Accept(visitor ASTVisitor) error {
return err
}
}
if c.Comment != nil {
if err := c.Comment.Accept(visitor); err != nil {
return err
}
}
return visitor.VisitCreateTable(c)
}

Expand Down Expand Up @@ -1424,14 +1454,10 @@ func (c *CreateMaterializedView) Accept(visitor ASTVisitor) error {
}
}
if c.Destination != nil {
// Destination.Accept visits its own TableSchema
if err := c.Destination.Accept(visitor); err != nil {
return err
}
if c.Destination.TableSchema != nil {
if err := c.Destination.TableSchema.Accept(visitor); err != nil {
return err
}
}
}
if c.SubQuery != nil {
if err := c.SubQuery.Accept(visitor); err != nil {
Expand Down Expand Up @@ -1971,6 +1997,9 @@ func (d *DestinationClause) Pos() Pos {
}

func (d *DestinationClause) End() Pos {
if d.TableSchema != nil {
return d.TableSchema.End()
}
return d.TableIdentifier.End()
}

Expand All @@ -1980,6 +2009,11 @@ func (d *DestinationClause) Accept(visitor ASTVisitor) error {
if err := d.TableIdentifier.Accept(visitor); err != nil {
return err
}
if d.TableSchema != nil {
if err := d.TableSchema.Accept(visitor); err != nil {
return err
}
}
return visitor.VisitDestinationExpr(d)
}

Expand Down Expand Up @@ -5131,6 +5165,11 @@ func (s *SelectQuery) Accept(visitor ASTVisitor) error {
return err
}
}
if s.DistinctOn != nil {
if err := s.DistinctOn.Accept(visitor); err != nil {
return err
}
}
if s.Top != nil {
if err := s.Top.Accept(visitor); err != nil {
return err
Expand Down Expand Up @@ -5353,6 +5392,11 @@ func (i *IntervalFrom) End() Pos {
func (i *IntervalFrom) Accept(visitor ASTVisitor) error {
visitor.Enter(i)
defer visitor.Leave(i)
if i.Interval != nil {
if err := i.Interval.Accept(visitor); err != nil {
return err
}
}
if err := i.FromExpr.Accept(visitor); err != nil {
return err
}
Expand Down Expand Up @@ -5976,17 +6020,23 @@ func (i *InsertStmt) End() Pos {
if i.SelectExpr != nil {
return i.SelectExpr.End()
}
return i.Values[len(i.Values)-1].End()
if len(i.Values) > 0 {
return i.Values[len(i.Values)-1].End()
}
// `INSERT INTO t FORMAT CSV` carries neither VALUES nor a SELECT — the
// data arrives out of band after the statement
if i.Format != nil {
return i.Format.End()
}
if i.ColumnNames != nil {
return i.ColumnNames.End()
}
return i.Table.End()
}

func (i *InsertStmt) Accept(visitor ASTVisitor) error {
visitor.Enter(i)
defer visitor.Leave(i)
if i.Format != nil {
if err := i.Format.Accept(visitor); err != nil {
return err
}
}
if err := i.Table.Accept(visitor); err != nil {
return err
}
Expand All @@ -5995,6 +6045,11 @@ func (i *InsertStmt) Accept(visitor ASTVisitor) error {
return err
}
}
if i.Format != nil {
if err := i.Format.Accept(visitor); err != nil {
return err
}
}
for _, value := range i.Values {
if err := value.Accept(visitor); err != nil {
return err
Expand Down
163 changes: 163 additions & 0 deletions parser/traversal_drift_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -135,3 +135,166 @@ func TestVisitJoinTableExprWithSampleRatio(t *testing.T) {
require.True(t, visitedSample, "SampleClause was not visited")
require.True(t, visitedJoinTable, "JoinTableExpr was not visited")
}

// TestTraversalEnginesVisitSameFields statically asserts that, for every node
// type, the set of child fields referenced by its Accept method matches the
// set referenced by its case in Walk's type switch. The type-level test above
// cannot catch a child field that one engine traverses and the other forgot
// (e.g. Walk missing InsertStmt.Values while Accept visits it).
func TestTraversalEnginesVisitSameFields(t *testing.T) {
entries, err := os.ReadDir(".")
require.NoError(t, err)

fset := token.NewFileSet()
acceptFields := map[string]map[string]bool{}
walkFields := map[string]map[string]bool{}

// collectSelectors records every selector `<base>.<Field>` in node whose
// base is the identifier baseName, into out.
collectSelectors := func(node ast.Node, baseName string, out map[string]bool) {
ast.Inspect(node, func(n ast.Node) bool {
sel, ok := n.(*ast.SelectorExpr)
if !ok {
return true
}
if ident, ok := sel.X.(*ast.Ident); ok && ident.Name == baseName {
out[sel.Sel.Name] = true
}
return true
})
}

for _, entry := range entries {
name := entry.Name()
if !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") {
continue
}
file, err := goparser.ParseFile(fset, name, nil, 0)
require.NoError(t, err)
for _, decl := range file.Decls {
d, ok := decl.(*ast.FuncDecl)
if !ok {
continue
}
switch {
case d.Name.Name == "Accept" && d.Recv != nil && len(d.Recv.List) == 1:
star, ok := d.Recv.List[0].Type.(*ast.StarExpr)
if !ok {
continue
}
typeIdent, ok := star.X.(*ast.Ident)
if !ok {
continue
}
recvName := d.Recv.List[0].Names[0].Name
fields := map[string]bool{}
collectSelectors(d.Body, recvName, fields)
acceptFields[typeIdent.Name] = fields
case d.Name.Name == "Walk" && d.Recv == nil:
ast.Inspect(d.Body, func(n ast.Node) bool {
cc, ok := n.(*ast.CaseClause)
if !ok {
return true
}
fields := map[string]bool{}
for _, stmt := range cc.Body {
collectSelectors(stmt, "n", fields)
}
for _, expr := range cc.List {
star, ok := expr.(*ast.StarExpr)
if !ok {
continue
}
if ident, ok := star.X.(*ast.Ident); ok {
walkFields[ident.Name] = fields
}
}
return true
})
}
}
}

require.NotEmpty(t, acceptFields)
require.NotEmpty(t, walkFields)

var problems []string
for typeName, aFields := range acceptFields {
wFields, ok := walkFields[typeName]
if !ok {
continue // type-level coverage is asserted by the test above
}
for _, field := range diffSet(aFields, wFields) {
problems = append(problems,
typeName+"."+field+" is traversed by Accept but not by Walk")
}
for _, field := range diffSet(wFields, aFields) {
problems = append(problems,
typeName+"."+field+" is traversed by Walk but not by Accept")
}
}
sort.Strings(problems)
require.Empty(t, problems, "child-field traversal drift between Accept and Walk")
}

// TestInsertStmtEndWithoutValues guards that End() does not panic on
// `INSERT INTO t FORMAT CSV`, where the data arrives out of band and the
// statement carries neither VALUES nor a SELECT.
func TestInsertStmtEndWithoutValues(t *testing.T) {
sql := "INSERT INTO t FORMAT CSV"
stmts, err := NewParser(sql).ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)
insert := stmts[0].(*InsertStmt)
require.Empty(t, insert.Values)
require.Nil(t, insert.SelectExpr)
require.Equal(t, Pos(len(sql)), insert.End())
}

// TestWalkVisitsInsertValues guards that Walk descends into INSERT ... VALUES
// rows; the InsertStmt case used to skip the Values field entirely.
func TestWalkVisitsInsertValues(t *testing.T) {
stmts, err := NewParser("INSERT INTO t VALUES (1, 2), (3, 4)").ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)

values := 0
Walk(stmts[0], func(node Expr) bool {
if _, ok := node.(*AssignmentValues); ok {
values++
}
return true
})
require.Equal(t, 2, values, "both VALUES rows should be walked")
}

// TestDestinationTableSchemaVisitedOnce guards that a materialized view's TO
// destination column list is visited exactly once, inside the destination
// clause, by both traversal engines.
func TestDestinationTableSchemaVisitedOnce(t *testing.T) {
sql := "CREATE MATERIALIZED VIEW mv TO dest (id UInt64) AS SELECT id FROM src"
stmts, err := NewParser(sql).ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)

acceptVisits := 0
visitor := &DefaultASTVisitor{
Visit: func(expr Expr) error {
if _, ok := expr.(*TableSchemaClause); ok {
acceptVisits++
}
return nil
},
}
require.NoError(t, stmts[0].Accept(visitor))

walkVisits := 0
Walk(stmts[0], func(node Expr) bool {
if _, ok := node.(*TableSchemaClause); ok {
walkVisits++
}
return true
})
require.Equal(t, 1, acceptVisits, "Accept should visit the destination schema exactly once")
require.Equal(t, walkVisits, acceptVisits, "Accept and Walk disagree on schema visits")
}
8 changes: 8 additions & 0 deletions parser/walk.go
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,9 @@ func Walk(node Expr, fn WalkFunc) bool {
if !Walk(n.Then, fn) {
return false
}
if !Walk(n.Else, fn) {
return false
}
case *CaseExpr:
if !Walk(n.Expr, fn) {
return false
Expand Down Expand Up @@ -326,6 +329,11 @@ func Walk(node Expr, fn WalkFunc) bool {
if !Walk(n.Format, fn) {
return false
}
for _, value := range n.Values {
if !Walk(value, fn) {
return false
}
}
if !Walk(n.SelectExpr, fn) {
return false
}
Expand Down
Loading