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
}