clearlove2.1
1package migrate
2
3import (
4 "errors"
5 "strings"
6)
7
8var errEmptyFile = errors.New("文件内容为空")
9
10// dumpSource MySQL 导出文件(.sql)数据源。
11// tables 保存的是**有序**列名列表:部分 INSERT 不带列名列表,
12// 必须按建表语句的列顺序还原,否则字段会错位。
13type dumpSource struct {
14 tables map[string][]string
15 rows map[string][]Row
16}
17
18func (s *dumpSource) HasTable(table string) bool { _, ok := s.tables[table]; return ok }
19
20func (s *dumpSource) Kind() string { return "mysqldump" }
21
22func (s *dumpSource) Close() error { return nil }
23
24func (s *dumpSource) Rows(table string, _ ...string) ([]Row, error) {
25 return s.rows[table], nil
26}
27
28// parseDump 解析 MySQL 导出文件
29func parseDump(data []byte) (Source, error) {
30 s := &dumpSource{tables: map[string][]string{}, rows: map[string][]Row{}}
31 stmts := splitSQL(string(data))
32 for _, stmt := range stmts {
33 up := strings.ToUpper(strings.TrimSpace(stmt))
34 switch {
35 case strings.HasPrefix(up, "CREATE TABLE"):
36 name, cols := parseCreateTable(stmt)
37 if name == "" {
38 continue
39 }
40 s.tables[name] = cols
41 case strings.HasPrefix(up, "INSERT INTO"):
42 name, cols, tuples := parseInsert(stmt)
43 if name == "" || len(tuples) == 0 {
44 continue
45 }
46 if len(cols) == 0 {
47 // 没有列名列表时按建表语句的列顺序还原
48 cols = s.tables[name]
49 if len(cols) == 0 {
50 continue
51 }
52 }
53 if _, ok := s.tables[name]; !ok {
54 s.tables[name] = cols
55 }
56 for _, t := range tuples {
57 row := Row{}
58 for i, c := range cols {
59 if i < len(t) {
60 row[c] = t[i]
61 }
62 }
63 s.rows[name] = append(s.rows[name], row)
64 }
65 }
66 }
67 if len(s.tables) == 0 {
68 return nil, errors.New("未能从文件中解析出任何数据表,请确认这是 MySQL 导出的 .sql 文件")
69 }
70 return s, nil
71}
72
73// ---------- SQL 语句切分 ----------
74
75// splitSQL 按分号切分 SQL 语句,自动跳过字符串、标识符与注释中的分号
76func splitSQL(s string) []string {
77 var out []string
78 var cur strings.Builder
79 n := len(s)
80 for i := 0; i < n; i++ {
81 c := s[i]
82 switch c {
83 case '\'', '"', '`':
84 // 引用内容整体拷贝(含转义)
85 q := c
86 cur.WriteByte(c)
87 i++
88 for i < n {
89 if s[i] == '\\' && i+1 < n && q != '`' {
90 cur.WriteByte(s[i])
91 cur.WriteByte(s[i+1])
92 i += 2
93 continue
94 }
95 cur.WriteByte(s[i])
96 if s[i] == q {
97 i++
98 break
99 }
100 i++
101 }
102 i--
103 case '-':
104 // -- 行注释
105 if i+1 < n && s[i+1] == '-' {
106 for i < n && s[i] != '\n' {
107 i++
108 }
109 } else {
110 cur.WriteByte(c)
111 }
112 case '#':
113 for i < n && s[i] != '\n' {
114 i++
115 }
116 case '/':
117 // 块注释(含 /*! 条件注释)
118 if i+1 < n && s[i+1] == '*' {
119 i += 2
120 for i+1 < n && !(s[i] == '*' && s[i+1] == '/') {
121 i++
122 }
123 i++
124 } else {
125 cur.WriteByte(c)
126 }
127 case ';':
128 if stmt := strings.TrimSpace(cur.String()); stmt != "" {
129 out = append(out, stmt)
130 }
131 cur.Reset()
132 default:
133 cur.WriteByte(c)
134 }
135 }
136 if stmt := strings.TrimSpace(cur.String()); stmt != "" {
137 out = append(out, stmt)
138 }
139 return out
140}
141
142// ---------- CREATE TABLE ----------
143
144// parseCreateTable 解析建表语句,返回表名与列名列表
145func parseCreateTable(stmt string) (string, []string) {
146 up := strings.ToUpper(stmt)
147 i := strings.Index(up, "CREATE TABLE")
148 if i < 0 {
149 return "", nil
150 }
151 i += len("CREATE TABLE")
152 // 跳过 IF NOT EXISTS
153 rest := strings.ToUpper(stmt[i:])
154 if j := strings.Index(rest, "IF NOT EXISTS"); j >= 0 && j < 24 {
155 i += j + len("IF NOT EXISTS")
156 }
157 i = skipSpace(stmt, i)
158 name, i := readIdent(stmt, i)
159 i = skipSpace(stmt, i)
160 if name == "" || i >= len(stmt) || stmt[i] != '(' {
161 return name, nil
162 }
163 end := matchParen(stmt, i)
164 if end < 0 {
165 return name, nil
166 }
167 var cols []string
168 for _, part := range splitTopLevel(stmt[i+1 : end]) {
169 part = strings.TrimSpace(part)
170 // 只取以反引号开头的列定义,跳过 PRIMARY KEY / KEY / UNIQUE KEY / CONSTRAINT 等
171 if !strings.HasPrefix(part, "`") {
172 continue
173 }
174 if c, _ := readIdent(part, 0); c != "" {
175 cols = append(cols, c)
176 }
177 }
178 return name, cols
179}
180
181// ---------- INSERT INTO ----------
182
183// parseInsert 解析插入语句,返回表名、列名列表与值元组
184func parseInsert(stmt string) (string, []string, [][]string) {
185 up := strings.ToUpper(stmt)
186 i := strings.Index(up, "INSERT INTO")
187 if i < 0 {
188 return "", nil, nil
189 }
190 i += len("INSERT INTO")
191 i = skipSpace(stmt, i)
192 // 可选修饰符
193 for _, kw := range []string{"LOW_PRIORITY", "DELAYED", "HIGH_PRIORITY", "IGNORE"} {
194 if strings.HasPrefix(strings.ToUpper(stmt[i:]), kw) {
195 i += len(kw)
196 i = skipSpace(stmt, i)
197 }
198 }
199 table, i := readIdent(stmt, i)
200 if table == "" {
201 return "", nil, nil
202 }
203 i = skipSpace(stmt, i)
204 var cols []string
205 if i < len(stmt) && stmt[i] == '(' {
206 end := matchParen(stmt, i)
207 if end < 0 {
208 return table, nil, nil
209 }
210 for _, c := range strings.Split(stmt[i+1:end], ",") {
211 c = strings.TrimSpace(c)
212 c = strings.Trim(c, "`")
213 if c != "" {
214 cols = append(cols, c)
215 }
216 }
217 i = end + 1
218 }
219 rest := stmt[i:]
220 vi := strings.Index(strings.ToUpper(rest), "VALUES")
221 if vi < 0 {
222 return table, cols, nil
223 }
224 return table, cols, parseTuples(rest[vi+len("VALUES"):])
225}
226
227// ---------- 词法辅助 ----------
228
229func skipSpace(s string, i int) int {
230 for i < len(s) && (s[i] == ' ' || s[i] == '\t' || s[i] == '\n' || s[i] == '\r') {
231 i++
232 }
233 return i
234}
235
236// readIdent 读取标识符(支持反引号包裹)
237func readIdent(s string, i int) (string, int) {
238 i = skipSpace(s, i)
239 if i >= len(s) {
240 return "", i
241 }
242 if s[i] == '`' {
243 i++
244 var b strings.Builder
245 for i < len(s) {
246 if s[i] == '`' {
247 if i+1 < len(s) && s[i+1] == '`' {
248 b.WriteByte('`')
249 i += 2
250 continue
251 }
252 i++
253 break
254 }
255 b.WriteByte(s[i])
256 i++
257 }
258 return b.String(), i
259 }
260 start := i
261 for i < len(s) && !strings.ContainsRune(" \t\n\r(,;", rune(s[i])) {
262 i++
263 }
264 return s[start:i], i
265}
266
267// matchParen 返回与 s[i]('(')匹配的 ')' 下标,找不到返回 -1
268func matchParen(s string, i int) int {
269 depth := 0
270 for ; i < len(s); i++ {
271 switch s[i] {
272 case '(':
273 depth++
274 case ')':
275 depth--
276 if depth == 0 {
277 return i
278 }
279 case '\'', '"', '`':
280 q := s[i]
281 i++
282 for i < len(s) {
283 if s[i] == '\\' && i+1 < len(s) && q != '`' {
284 i += 2
285 continue
286 }
287 if s[i] == q {
288 break
289 }
290 i++
291 }
292 }
293 }
294 return -1
295}
296
297// splitTopLevel 按顶层逗号切分(忽略括号与引号内的逗号)
298func splitTopLevel(s string) []string {
299 var out []string
300 depth, start := 0, 0
301 for i := 0; i < len(s); i++ {
302 switch s[i] {
303 case '(':
304 depth++
305 case ')':
306 depth--
307 case '\'', '"', '`':
308 q := s[i]
309 i++
310 for i < len(s) {
311 if s[i] == '\\' && i+1 < len(s) && q != '`' {
312 i += 2
313 continue
314 }
315 if s[i] == q {
316 break
317 }
318 i++
319 }
320 case ',':
321 if depth == 0 {
322 out = append(out, s[start:i])
323 start = i + 1
324 }
325 }
326 }
327 out = append(out, s[start:])
328 return out
329}
330
331// parseTuples 解析 VALUES 后的 "(v1,v2),(v3,v4)" 列表
332func parseTuples(s string) [][]string {
333 var out [][]string
334 i, n := 0, len(s)
335 for i < n {
336 for i < n && s[i] != '(' {
337 i++
338 }
339 if i >= n {
340 break
341 }
342 i++ // 跳过 '('
343 var vals []string
344 for {
345 v, ni := parseValue(s, i)
346 vals = append(vals, v)
347 i = skipSpace(s, ni)
348 if i < n && s[i] == ',' {
349 i++
350 continue
351 }
352 if i < n && s[i] == ')' {
353 i++
354 }
355 break
356 }
357 out = append(out, vals)
358 }
359 return out
360}
361
362// parseValue 解析一个值,返回字符串值与新的下标(NULL 与空串统一为空串)
363func parseValue(s string, i int) (string, int) {
364 n := len(s)
365 i = skipSpace(s, i)
366 if i >= n {
367 return "", i
368 }
369 if s[i] == '\'' || s[i] == '"' {
370 q := s[i]
371 i++
372 var b strings.Builder
373 for i < n {
374 c := s[i]
375 if c == '\\' && i+1 < n {
376 b.WriteByte(decodeEscape(s[i+1]))
377 i += 2
378 continue
379 }
380 if c == q {
381 // '' 表示一个引号字符
382 if i+1 < n && s[i+1] == q {
383 b.WriteByte(q)
384 i += 2
385 continue
386 }
387 i++
388 break
389 }
390 b.WriteByte(c)
391 i++
392 }
393 return b.String(), i
394 }
395 start := i
396 for i < n && s[i] != ',' && s[i] != ')' {
397 i++
398 }
399 raw := strings.TrimSpace(s[start:i])
400 if strings.EqualFold(raw, "NULL") {
401 return "", i
402 }
403 return raw, i
404}
405
406// decodeEscape 解析 MySQL 反斜杠转义
407func decodeEscape(c byte) byte {
408 switch c {
409 case 'n':
410 return '\n'
411 case 'r':
412 return '\r'
413 case 't':
414 return '\t'
415 case '0':
416 return 0
417 case 'b':
418 return '\b'
419 case 'Z':
420 return 26
421 default:
422 return c // \\ \' \" \% \_ 等保持原字符
423 }
424}