Skip to content

Commit fefbd4b

Browse files
committed
syncer(dm): add MariaDB AST DDL rewriter
1 parent 804a466 commit fefbd4b

6 files changed

Lines changed: 762 additions & 4 deletions

File tree

dm/pkg/ddl/rewriter/rewriter.go

Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,135 @@
1+
// Copyright 2026 PingCAP, Inc.
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// See the License for the specific language governing permissions and
12+
// limitations under the License.
13+
14+
package rewriter
15+
16+
import (
17+
"strings"
18+
19+
"github.com/pingcap/tidb/pkg/parser"
20+
"github.com/pingcap/tidb/pkg/parser/ast"
21+
"github.com/pingcap/tidb/pkg/parser/format"
22+
_ "github.com/pingcap/tidb/pkg/types/parser_driver" // register parser driver
23+
)
24+
25+
// Rule rewrites one AST node in place.
26+
type Rule interface {
27+
Name() string
28+
Apply(ast.Node) (bool, error)
29+
}
30+
31+
// Rewriter applies a fixed, ordered set of AST rules.
32+
type Rewriter struct {
33+
rules []Rule
34+
}
35+
36+
// NewRewriter creates a rewriter from explicit rules.
37+
func NewRewriter(rules ...Rule) *Rewriter {
38+
return &Rewriter{rules: append([]Rule(nil), rules...)}
39+
}
40+
41+
// NewDefaultRewriter creates the default MariaDB compatibility AST rewriter.
42+
func NewDefaultRewriter() *Rewriter {
43+
return NewRewriter(DefaultRules()...)
44+
}
45+
46+
// NewRewriterForFlavor returns a default rewriter only for MariaDB upstreams.
47+
func NewRewriterForFlavor(flavor string) *Rewriter {
48+
if !strings.EqualFold(flavor, "mariadb") {
49+
return nil
50+
}
51+
return NewDefaultRewriter()
52+
}
53+
54+
// RewriteStmt applies all rules to stmt in place.
55+
func (r *Rewriter) RewriteStmt(stmt ast.StmtNode) (bool, error) {
56+
if r == nil || len(r.rules) == 0 || stmt == nil {
57+
return false, nil
58+
}
59+
visitor := &rewriteVisitor{rules: r.rules}
60+
stmt.Accept(visitor)
61+
return visitor.changed, visitor.err
62+
}
63+
64+
// RewriteSQL parses, rewrites, and restores SQL. It is mainly intended for unit tests
65+
// and small call sites that do not already have a parsed AST.
66+
func (r *Rewriter) RewriteSQL(sql string) (string, bool, error) {
67+
p := parser.New()
68+
stmts, _, err := p.Parse(sql, "", "")
69+
if err != nil {
70+
return "", false, err
71+
}
72+
73+
changed := false
74+
for _, stmt := range stmts {
75+
stmtChanged, err := r.RewriteStmt(stmt)
76+
if err != nil {
77+
return "", false, err
78+
}
79+
changed = changed || stmtChanged
80+
}
81+
if !changed {
82+
return sql, false, nil
83+
}
84+
85+
out, err := restoreStatements(stmts)
86+
if err != nil {
87+
return "", false, err
88+
}
89+
return out, true, nil
90+
}
91+
92+
type rewriteVisitor struct {
93+
rules []Rule
94+
changed bool
95+
err error
96+
}
97+
98+
func (v *rewriteVisitor) Enter(node ast.Node) (ast.Node, bool) {
99+
if v.err != nil {
100+
return node, true
101+
}
102+
return node, false
103+
}
104+
105+
func (v *rewriteVisitor) Leave(node ast.Node) (ast.Node, bool) {
106+
if v.err != nil {
107+
return node, false
108+
}
109+
for _, rule := range v.rules {
110+
changed, err := rule.Apply(node)
111+
if err != nil {
112+
v.err = err
113+
return node, false
114+
}
115+
v.changed = v.changed || changed
116+
}
117+
return node, true
118+
}
119+
120+
func restoreStatements(stmts []ast.StmtNode) (string, error) {
121+
var out strings.Builder
122+
for i, stmt := range stmts {
123+
if i > 0 {
124+
out.WriteString(";\n")
125+
}
126+
err := stmt.Restore(&format.RestoreCtx{
127+
Flags: format.DefaultRestoreFlags | format.RestoreTiDBSpecialComment | format.RestoreStringWithoutDefaultCharset,
128+
In: &out,
129+
})
130+
if err != nil {
131+
return "", err
132+
}
133+
}
134+
return out.String(), nil
135+
}
Lines changed: 160 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,160 @@
1+
// Copyright 2026 PingCAP, Inc.
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// See the License for the specific language governing permissions and
12+
// limitations under the License.
13+
14+
package rewriter
15+
16+
import (
17+
"strings"
18+
"testing"
19+
20+
"github.com/pingcap/tidb/pkg/parser"
21+
"github.com/pingcap/tidb/pkg/parser/ast"
22+
"github.com/stretchr/testify/require"
23+
)
24+
25+
func TestRewriteSQLRemovesFunctionDefaultOnVarchar(t *testing.T) {
26+
rewriter := NewDefaultRewriter()
27+
28+
out, changed, err := rewriter.RewriteSQL("CREATE TABLE t(t VARCHAR(100) DEFAULT current_timestamp());")
29+
require.NoError(t, err)
30+
require.True(t, changed)
31+
require.NotContains(t, strings.ToLower(out), "default")
32+
33+
stmt := parseCreateTable(t, out)
34+
col := findColumn(stmt, "t")
35+
require.NotNil(t, col)
36+
require.False(t, hasColumnOption(col, ast.ColumnOptionDefaultValue))
37+
}
38+
39+
func TestRewriteSQLKeepsTimeFunctionDefaultOnTimeColumn(t *testing.T) {
40+
rewriter := NewDefaultRewriter()
41+
42+
out, changed, err := rewriter.RewriteSQL("CREATE TABLE t(ts TIMESTAMP DEFAULT current_timestamp());")
43+
require.NoError(t, err)
44+
require.False(t, changed)
45+
require.Contains(t, strings.ToLower(out), "default current_timestamp")
46+
}
47+
48+
func TestRewriteSQLDefaultRules(t *testing.T) {
49+
rewriter := NewDefaultRewriter()
50+
input := `CREATE TABLE t (
51+
id INT(11),
52+
txt TEXT DEFAULT 'x',
53+
v VARCHAR(800),
54+
j JSON,
55+
g JSON GENERATED ALWAYS AS (JSON_EXTRACT(j, '$.a')) VIRTUAL,
56+
zero_ts TIMESTAMP DEFAULT '0000-00-00 00:00:00',
57+
CHECK (json_valid(j)),
58+
KEY idx_txt (txt),
59+
KEY idx_v (v)
60+
) DEFAULT CHARSET=latin1 COLLATE=latin1_swedish_ci;`
61+
62+
out, changed, err := rewriter.RewriteSQL(input)
63+
require.NoError(t, err)
64+
require.True(t, changed)
65+
66+
stmt := parseCreateTable(t, out)
67+
require.Equal(t, "utf8mb4", findTableOption(stmt, ast.TableOptionCharset))
68+
require.Equal(t, "utf8mb4_0900_ai_ci", findTableOption(stmt, ast.TableOptionCollate))
69+
require.Equal(t, -1, findColumn(stmt, "id").Tp.GetFlen())
70+
require.Equal(t, 768, findColumn(stmt, "v").Tp.GetFlen())
71+
require.False(t, hasColumnOption(findColumn(stmt, "txt"), ast.ColumnOptionDefaultValue))
72+
require.False(t, hasColumnOption(findColumn(stmt, "g"), ast.ColumnOptionGenerated))
73+
require.False(t, hasColumnOption(findColumn(stmt, "zero_ts"), ast.ColumnOptionDefaultValue))
74+
require.False(t, hasJSONValidCheck(stmt))
75+
76+
idxTxt := findConstraint(stmt, "idx_txt")
77+
require.NotNil(t, idxTxt)
78+
require.Equal(t, 255, idxTxt.Keys[0].Length)
79+
idxV := findConstraint(stmt, "idx_v")
80+
require.NotNil(t, idxV)
81+
require.Equal(t, 768, idxV.Keys[0].Length)
82+
}
83+
84+
func TestRewriteSQLSkipsExpressionIndexPrefix(t *testing.T) {
85+
rewriter := NewDefaultRewriter()
86+
87+
_, _, err := rewriter.RewriteSQL("CREATE TABLE t(name VARCHAR(32), KEY idx_expr ((LOWER(name))));")
88+
require.NoError(t, err)
89+
}
90+
91+
func TestRewriteSQLRemovesParenthesizedJSONGeneratedColumn(t *testing.T) {
92+
rewriter := NewDefaultRewriter()
93+
94+
out, changed, err := rewriter.RewriteSQL(
95+
"CREATE TABLE t(j JSON, g JSON GENERATED ALWAYS AS ((JSON_EXTRACT(j, '$.a'))) VIRTUAL);",
96+
)
97+
require.NoError(t, err)
98+
require.True(t, changed)
99+
require.False(t, hasColumnOption(findColumn(parseCreateTable(t, out), "g"), ast.ColumnOptionGenerated))
100+
}
101+
102+
func TestNewRewriterForFlavor(t *testing.T) {
103+
require.NotNil(t, NewRewriterForFlavor("mariadb"))
104+
require.NotNil(t, NewRewriterForFlavor("MariaDB"))
105+
require.Nil(t, NewRewriterForFlavor("mysql"))
106+
}
107+
108+
func parseCreateTable(t *testing.T, sql string) *ast.CreateTableStmt {
109+
t.Helper()
110+
stmt, err := parser.New().ParseOneStmt(sql, "", "")
111+
require.NoError(t, err)
112+
create, ok := stmt.(*ast.CreateTableStmt)
113+
require.True(t, ok)
114+
return create
115+
}
116+
117+
func findColumn(stmt *ast.CreateTableStmt, name string) *ast.ColumnDef {
118+
for _, col := range stmt.Cols {
119+
if strings.EqualFold(col.Name.Name.O, name) {
120+
return col
121+
}
122+
}
123+
return nil
124+
}
125+
126+
func hasColumnOption(col *ast.ColumnDef, optionType ast.ColumnOptionType) bool {
127+
for _, opt := range col.Options {
128+
if opt.Tp == optionType {
129+
return true
130+
}
131+
}
132+
return false
133+
}
134+
135+
func hasJSONValidCheck(stmt *ast.CreateTableStmt) bool {
136+
for _, cons := range stmt.Constraints {
137+
if cons.Tp == ast.ConstraintCheck && isJSONValidExpr(cons.Expr) {
138+
return true
139+
}
140+
}
141+
return false
142+
}
143+
144+
func findConstraint(stmt *ast.CreateTableStmt, name string) *ast.Constraint {
145+
for _, cons := range stmt.Constraints {
146+
if strings.EqualFold(cons.Name, name) {
147+
return cons
148+
}
149+
}
150+
return nil
151+
}
152+
153+
func findTableOption(stmt *ast.CreateTableStmt, optionType ast.TableOptionType) string {
154+
for _, opt := range stmt.Options {
155+
if opt.Tp == optionType {
156+
return strings.ToLower(opt.StrValue)
157+
}
158+
}
159+
return ""
160+
}

0 commit comments

Comments
 (0)