diff --git a/parser/ast.go b/parser/ast.go index e06260e..a2c214b 100644 --- a/parser/ast.go +++ b/parser/ast.go @@ -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) } @@ -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) } @@ -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) } @@ -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 @@ -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) } @@ -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) } @@ -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 { @@ -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() } @@ -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) } @@ -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 @@ -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 } @@ -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 } @@ -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 diff --git a/parser/traversal_drift_test.go b/parser/traversal_drift_test.go index 3fc75a0..c4e958d 100644 --- a/parser/traversal_drift_test.go +++ b/parser/traversal_drift_test.go @@ -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 `.` 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") +} diff --git a/parser/walk.go b/parser/walk.go index 077bee2..696f9e6 100644 --- a/parser/walk.go +++ b/parser/walk.go @@ -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 @@ -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 }