Files
backend_v2/cmd/dbtool/main.go
toom1996 6a5a4378ab update
2026-09-07 00:02:19 +08:00

280 lines
7.7 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// Command dbtool 是后台管理用的数据库一键迁移脚本。
//
// 用途:在不同开发电脑之间搬运 MySQL 库(db_dev)的「结构 + 数据」,
// 省去手动 mysqldump / 重新 seed 的麻烦。完全复用后端既有的 mysql 驱动与
// configs/config.yml,不依赖任何外部二进制(mysqldump 等)。
//
// 子命令:
//
// dbtool dump -out db_dump.json 导出当前库全部表(schema+data)为单个 JSON 文件
// dbtool import -in db_dump.json 读取 JSON 文件,DROP+CREATE+INSERT 回灌到目标库
// (可用 -data-only 只导数据,表结构需已存在)
//
// 连接信息来自 configs/config.yml 的 database 段(同后端服务),可用 -config 指定其它配置。
package main
import (
"database/sql"
"encoding/json"
"flag"
"fmt"
"log"
"os"
"strings"
"time"
_ "github.com/go-sql-driver/mysql"
"fashionapi/internal/config"
)
// tableDump 是单张表的导出结构。
type tableDump struct {
Name string `json:"name"`
Create string `json:"create"` // SHOW CREATE TABLE 得到的完整 DDL
Columns []string `json:"columns"` // SELECT * 得到的列顺序,导入时按此生成 INSERT
Rows [][]any `json:"rows"` // 每行为一组值;nil 表示 NULL,string 表示值
}
// dumpFile 是 dump 输出的顶层结构。
type dumpFile struct {
Version int `json:"version"`
GeneratedAt string `json:"generated_at"`
Database string `json:"database"`
Tables []tableDump `json:"tables"`
}
func main() {
if len(os.Args) < 2 {
usage()
os.Exit(2)
}
switch os.Args[1] {
case "dump":
runDump(os.Args[2:])
case "import":
runImport(os.Args[2:])
default:
usage()
os.Exit(2)
}
}
func usage() {
fmt.Println(`dbtool - 后台数据库一键迁移脚本
用法:
dbtool dump [-config <yml>] [-out <file>] 导出全库 schema+data 到 JSON
dbtool import [-config <yml>] [-in <file>] [-data-only] 从 JSON 回灌(DROP+CREATE+INSERT)
说明:
连接信息读取 configs/config.yml 的 database 段(可用 -config 覆盖)。
import 默认连结构带数据全部重建;-data-only 仅导数据(目标库表结构须已存在)。`)
}
// openDB 按 config 加载 DSN 并探活。
func openDB(cfgPath string) (*sql.DB, *config.DatabaseConfig, error) {
cfg, err := config.Load(cfgPath)
if err != nil {
return nil, nil, fmt.Errorf("加载配置失败: %w", err)
}
db, err := sql.Open("mysql", cfg.Database.DSN())
if err != nil {
return nil, nil, fmt.Errorf("打开数据库连接失败: %w", err)
}
if err := db.Ping(); err != nil {
_ = db.Close()
return nil, nil, fmt.Errorf("数据库连接失败(%s): %w", cfg.Database.Addr(), err)
}
return db, &cfg.Database, nil
}
// runDump 导出全库。
func runDump(args []string) {
fs := flag.NewFlagSet("dump", flag.ExitOnError)
out := fs.String("out", "db_dump.json", "导出文件路径")
cfgPath := fs.String("config", "", "配置文件路径(默认 configs/config.yml)")
_ = fs.Parse(args)
db, dcfg, err := openDB(*cfgPath)
if err != nil {
log.Fatalf("✗ %v", err)
}
defer db.Close()
tables, err := listTables(db)
if err != nil {
log.Fatalf("✗ 列举表失败: %v", err)
}
df := dumpFile{
Version: 1,
GeneratedAt: time.Now().UTC().Format(time.RFC3339),
Database: dcfg.Name,
Tables: make([]tableDump, 0, len(tables)),
}
for _, t := range tables {
log.Printf("→ 导出表 %s ...", t)
td, err := dumpTable(db, t)
if err != nil {
log.Fatalf("✗ 导出表 %s 失败: %v", t, err)
}
df.Tables = append(df.Tables, td)
}
data, err := json.MarshalIndent(df, "", " ")
if err != nil {
log.Fatalf("✗ 序列化失败: %v", err)
}
if err := os.WriteFile(*out, data, 0o644); err != nil {
log.Fatalf("✗ 写文件 %s 失败: %v", *out, err)
}
log.Printf("✓ 已导出 %d 张表 -> %s", len(df.Tables), *out)
}
// runImport 回灌。
func runImport(args []string) {
fs := flag.NewFlagSet("import", flag.ExitOnError)
in := fs.String("in", "db_dump.json", "导入文件路径")
cfgPath := fs.String("config", "", "配置文件路径(默认 configs/config.yml)")
dataOnly := fs.Bool("data-only", false, "仅导数据(表结构须已存在)")
_ = fs.Parse(args)
raw, err := os.ReadFile(*in)
if err != nil {
log.Fatalf("✗ 读文件 %s 失败: %v", *in, err)
}
var df dumpFile
if err := json.Unmarshal(raw, &df); err != nil {
log.Fatalf("✗ 解析 %s 失败: %v", *in, err)
}
db, _, err := openDB(*cfgPath)
if err != nil {
log.Fatalf("✗ %v", err)
}
defer db.Close()
// 关掉外键检查,避免建表/插数据顺序受 FK 约束(全量重建无需保序)。
if _, err := db.Exec("SET FOREIGN_KEY_CHECKS=0"); err != nil {
log.Fatalf("✗ 关闭外键检查失败: %v", err)
}
defer db.Exec("SET FOREIGN_KEY_CHECKS=1")
for _, t := range df.Tables {
if !*dataOnly {
log.Printf("→ 重建表 %s ...", t.Name)
if _, err := db.Exec(fmt.Sprintf("DROP TABLE IF EXISTS `%s`", t.Name)); err != nil {
log.Fatalf("✗ 删表 %s 失败: %v", t.Name, err)
}
if _, err := db.Exec(t.Create); err != nil {
log.Fatalf("✗ 建表 %s 失败: %v", t.Name, err)
}
}
if err := importRows(db, t); err != nil {
log.Fatalf("✗ 导数据到 %s 失败: %v", t.Name, err)
}
log.Printf(" ✓ %s: %d 行", t.Name, len(t.Rows))
}
log.Printf("✓ 导入完成(%d 张表)", len(df.Tables))
}
// listTables 返回当前库所有表名。
func listTables(db *sql.DB) ([]string, error) {
rows, err := db.Query("SHOW TABLES")
if err != nil {
return nil, err
}
defer rows.Close()
var out []string
for rows.Next() {
var name string
if err := rows.Scan(&name); err != nil {
return nil, err
}
out = append(out, name)
}
return out, rows.Err()
}
// dumpTable 导出单张表的 DDL + 数据。
func dumpTable(db *sql.DB, name string) (tableDump, error) {
// 1) DDL
var dummy, ddl string
if err := db.QueryRow(fmt.Sprintf("SHOW CREATE TABLE `%s`", name)).Scan(&dummy, &ddl); err != nil {
return tableDump{}, err
}
// 2) 数据
rows, err := db.Query(fmt.Sprintf("SELECT * FROM `%s`", name))
if err != nil {
return tableDump{}, err
}
defer rows.Close()
cols, err := rows.Columns()
if err != nil {
return tableDump{}, err
}
n := len(cols)
td := tableDump{Name: name, Create: ddl, Columns: cols, Rows: make([][]any, 0)}
scanPtrs := make([]any, n)
raw := make([]sql.RawBytes, n)
for i := range raw {
scanPtrs[i] = &raw[i]
}
for rows.Next() {
if err := rows.Scan(scanPtrs...); err != nil {
return tableDump{}, err
}
row := make([]any, n)
for i := range raw {
if raw[i] == nil {
row[i] = nil // NULL
} else {
row[i] = string(raw[i]) // 全部按字符串搬运,导入时由驱动按列类型适配
}
}
td.Rows = append(td.Rows, row)
}
return td, rows.Err()
}
// importRows 把一张表的数据批量 INSERT 进目标库(单表一个事务,每批多行,减少网络往返)。
func importRows(db *sql.DB, t tableDump) error {
if len(t.Rows) == 0 {
return nil
}
tx, err := db.Begin()
if err != nil {
return err
}
colList := "`" + strings.Join(t.Columns, "`,`") + "`"
ph := "(" + strings.TrimSuffix(strings.Repeat("?,", len(t.Columns)), ",") + ")"
const batchSize = 200
for start := 0; start < len(t.Rows); start += batchSize {
end := start + batchSize
if end > len(t.Rows) {
end = len(t.Rows)
}
chunk := t.Rows[start:end]
var sb strings.Builder
sb.WriteString(fmt.Sprintf("INSERT INTO `%s` (%s) VALUES ", t.Name, colList))
args := make([]any, 0, len(chunk)*len(t.Columns))
for i, r := range chunk {
if i > 0 {
sb.WriteString(",")
}
sb.WriteString(ph)
args = append(args, r...)
}
if _, err := tx.Exec(sb.String(), args...); err != nil {
_ = tx.Rollback()
return err
}
}
return tx.Commit()
}