From a93223f573547fff2d9adc7e2ee6937aae759146 Mon Sep 17 00:00:00 2001 From: git-hulk Date: Sun, 5 Jul 2026 11:05:25 +0800 Subject: [PATCH] Fix Accept/Walk field-level drift and InsertStmt End panic The two traversal engines each encode every node's children by hand, and the existing drift test only asserts that both engines know the same node TYPES - it cannot see a child FIELD one engine traverses and the other forgot. A new static test compares, per node type, the fields referenced by Accept against the fields referenced by Walk's case, and it immediately found ten drifted fields: Accept was missing (Walk visited them): - AlterTableDropPartition.Settings - CreateDatabase.Name, CreateDatabase.Comment - CreateTable.Comment - IntervalFrom.Interval - ProjectionOrderByClause.Columns - SelectQuery.DistinctOn - UUID.Value Walk was missing (Accept visited them): - InsertStmt.Values (INSERT ... VALUES rows were never walked) - SelectQuery.Except - WhenClause.Else Also: - InsertStmt.End() panicked on `INSERT INTO t FORMAT CSV`: with the data arriving out of band there are no Values and no SelectExpr, and the method indexed Values[len-1] unguarded. It now falls back to Format, ColumnNames, then Table. - InsertStmt.Accept visited Format before Table, violating source order; it now matches Walk (Table, ColumnNames, Format, Values, SelectExpr). - DestinationClause.Accept did not visit its own TableSchema; CreateMaterializedView.Accept compensated by reaching into its child, so any other DestinationClause holder silently skipped the schema and the visit happened outside the destination's Enter/Leave bracket. The clause now owns its child, and its End() includes the schema. No golden fixture changed. Co-Authored-By: Claude Fable 5 --- parser/ast.go | 77 +++++++++++++--- parser/traversal_drift_test.go | 163 +++++++++++++++++++++++++++++++++ parser/walk.go | 11 +++ 3 files changed, 240 insertions(+), 11 deletions(-) diff --git a/parser/ast.go b/parser/ast.go index da902fb..3c8a1fb 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) } @@ -5141,6 +5175,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 @@ -5358,6 +5397,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 } @@ -5981,17 +6025,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 } @@ -6000,6 +6050,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 7e78294..dd33e8c 100644 --- a/parser/walk.go +++ b/parser/walk.go @@ -69,6 +69,9 @@ func Walk(node Expr, fn WalkFunc) bool { if !Walk(n.UnionDistinct, fn) { return false } + if !Walk(n.Except, fn) { + return false + } if !Walk(n.Format, fn) { return false } @@ -151,6 +154,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 @@ -320,6 +326,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 }