// Command dbtool 是后台管理用的数据库一键迁移脚本。 // // 用途:在不同开发电脑 / 环境之间搬运 PostgreSQL 库(fashion)的「结构 + 数据」, // 省去手动 pg_dump / 重新 seed 的麻烦。完全复用后端既有的 pgx 驱动与 // configs/config.yml,不依赖任何外部二进制(pg_dump 等)。 // // 子命令: // // dbtool dump -out db_dump.json 导出当前库全部表(schema+data)为单个 JSON 文件 // dbtool import -in db_dump.json 读取 JSON 文件,DROP+CREATE+INSERT 回灌到目标库 // (可用 -data-only 只导数据,表结构须已由 AutoMigrate 创建) // // 连接信息来自 configs/config.yml 的 database 段(同后端服务),可用 -config 指定其它配置。 package main import ( "database/sql" "encoding/json" "errors" "flag" "fmt" "log" "os" "strings" "time" "github.com/jackc/pgx/v5/pgconn" _ "github.com/jackc/pgx/v5/stdlib" "fashionapi/internal/config" ) // columnMeta 是单列的元信息,用于重建 DDL 与导入时类型转换。 type columnMeta struct { Name string `json:"name"` Type string `json:"type"` // udt_name,如 int8 / varchar / timestamptz / bool / jsonb IsNullable bool `json:"is_nullable"` IsIdentity bool `json:"is_identity"` Default string `json:"default,omitempty"` } // tableDump 是单张表的导出结构。 type tableDump struct { Name string `json:"name"` Columns []columnMeta `json:"columns"` Pk []string `json:"pk"` 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"` } // verbose 控制是否打印每条执行的 SQL(用于排查)。由 dump/import 子命令的 -v 开关设置。 var verbose bool // vlog 在开启 verbose 时打印调试信息。 func vlog(format string, args ...any) { if verbose { log.Printf("[sql] "+format, args...) } } // pgDetail 提取 PostgreSQL 报错的位置/消息,便于定位语法错误。 func pgDetail(err error) string { var pgErr *pgconn.PgError if errors.As(err, &pgErr) { return fmt.Sprintf(" [position=%d message=%q where=%q]", pgErr.Position, pgErr.Message, pgErr.Where) } return "" } // quoteLit 安全包裹 SQL 字符串字面量(转义单引号),用于内联可信的内部标识符/表名。 func quoteLit(s string) string { return "'" + strings.ReplaceAll(s, "'", "''") + "'" } 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 - 后台数据库一键迁移脚本(PostgreSQL) 用法: dbtool dump [-config ] [-out ] 导出全库 schema+data 到 JSON dbtool import [-config ] [-in ] [-data-only] 从 JSON 回灌(默认 DROP+CREATE+INSERT) 说明: 连接信息读取 configs/config.yml 的 database 段(可用 -config 覆盖)。 import 默认连结构带数据全部重建;-data-only 仅导数据(表结构须已由后端 AutoMigrate 创建)。`) } // 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("pgx", 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)") v := fs.Bool("v", false, "打印每条执行的 SQL,便于排查") _ = fs.Parse(args) verbose = *v 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, "仅导数据(表结构须已由 AutoMigrate 创建)") _ = 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() // 关闭外键 / 触发器,避免插入顺序受约束(整库重建无需保序)。 if _, err := db.Exec("SET session_replication_role = 'replica'"); err != nil { log.Fatalf("✗ 关闭约束检查失败: %v", err) } defer db.Exec("SET session_replication_role = 'origin'") for _, t := range df.Tables { if !*dataOnly { log.Printf("→ 重建表 %s ...", t.Name) if _, err := db.Exec(fmt.Sprintf(`DROP TABLE IF EXISTS "%s" CASCADE`, t.Name)); err != nil { log.Fatalf("✗ 删表 %s 失败: %v", t.Name, err) } ddl, err := buildCreate(t) if err != nil { log.Fatalf("✗ 生成建表语句失败(%s): %v", t.Name, err) } if _, err := db.Exec(ddl); err != nil { log.Fatalf("✗ 建表 %s 失败: %v", t.Name, err) } } else { // 仅导数据:先清空目标表(约束已由 replica 角色关闭),再插入。 if _, err := db.Exec(fmt.Sprintf(`DELETE FROM "%s"`, t.Name)); 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 返回 public 模式下所有基表名。 func listTables(db *sql.DB) ([]string, error) { rows, err := db.Query(` SELECT table_name FROM information_schema.tables WHERE table_schema = 'public' AND table_type = 'BASE TABLE' ORDER BY table_name`) 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 导出单张表的列元信息 + 数据。 func dumpTable(db *sql.DB, name string) (tableDump, error) { cols, pk, err := dumpColumns(db, name) if err != nil { return tableDump{}, err } dataQuery := fmt.Sprintf(`SELECT * FROM "%s"`, name) vlog("dump data: %s", dataQuery) rows, err := db.Query(dataQuery) if err != nil { return tableDump{}, fmt.Errorf("读取数据失败(表 %s): %w\nSQL: %s", name, err, dataQuery) } defer rows.Close() colNames, err := rows.Columns() if err != nil { return tableDump{}, err } n := len(colNames) td := tableDump{Name: name, Columns: cols, Pk: pk, 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() } // dumpColumns 读取列元信息与主键。 func dumpColumns(db *sql.DB, table string) ([]columnMeta, []string, error) { colQuery := strings.Replace(` SELECT c.column_name, c.udt_name, (c.is_nullable = 'YES'), COALESCE(c.column_default, ''), COALESCE(c.column_default LIKE 'nextval(%', false) FROM information_schema.columns c WHERE c.table_schema = 'public' AND c.table_name = $1 ORDER BY c.ordinal_position`, "$1", quoteLit(table), 1) vlog("dump columns: table=%q query=%s", table, colQuery) rows, err := db.Query(colQuery) if err != nil { return nil, nil, fmt.Errorf("读取列元信息失败(表 %s): %w%s\nSQL: %s", table, err, pgDetail(err), strings.TrimSpace(colQuery)) } defer rows.Close() var cols []columnMeta for rows.Next() { var name, udt, def string var nullable, isIdent bool if err := rows.Scan(&name, &udt, &nullable, &def, &isIdent); err != nil { return nil, nil, err } cols = append(cols, columnMeta{ Name: name, Type: udt, IsNullable: nullable, IsIdentity: isIdent, Default: def, }) } if err := rows.Err(); err != nil { return nil, nil, err } pkQuery := strings.Replace(` SELECT kcu.column_name FROM information_schema.table_constraints tc JOIN information_schema.key_column_usage kcu ON kcu.constraint_name = tc.constraint_name AND kcu.table_schema = tc.table_schema WHERE tc.table_schema = 'public' AND tc.table_name = $1 AND tc.constraint_type = 'PRIMARY KEY' ORDER BY kcu.ordinal_position`, "$1", quoteLit(table), 1) vlog("dump pk: %s", strings.TrimSpace(pkQuery)) pkRows, err := db.Query(pkQuery) if err != nil { return nil, nil, fmt.Errorf("读取主键失败(表 %s): %w\nSQL: %s", table, err, strings.TrimSpace(pkQuery)) } defer pkRows.Close() var pk []string for pkRows.Next() { var c string if err := pkRows.Scan(&c); err != nil { return nil, nil, err } pk = append(pk, c) } return cols, pk, pkRows.Err() } // pgType 将 udt_name 映射为建表用的列类型。 func pgType(udt string) string { switch udt { case "int2": return "smallint" case "int4": return "integer" case "int8": return "bigint" case "numeric": return "numeric" case "float4": return "real" case "float8": return "double precision" case "bool": return "boolean" case "timestamp": return "timestamp" case "timestamptz": return "timestamptz" case "date": return "date" case "time": return "time" case "text": return "text" case "json": return "json" case "jsonb": return "jsonb" case "uuid": return "uuid" case "bytea": return "bytea" default: return udt // varchar / char / 未知类型原样返回 } } // buildCreate 由列元信息重建 CREATE TABLE 语句。 func buildCreate(t tableDump) (string, error) { var b strings.Builder b.WriteString(fmt.Sprintf(`CREATE TABLE IF NOT EXISTS "%s" (`, t.Name)) first := true for _, c := range t.Columns { if !first { b.WriteString(",") } first = false b.WriteString(fmt.Sprintf("\n \"%s\" %s", c.Name, pgType(c.Type))) if c.IsIdentity { b.WriteString(" GENERATED BY DEFAULT AS IDENTITY") } else if c.Default != "" { b.WriteString(" DEFAULT " + c.Default) } if !c.IsNullable && !c.IsIdentity { b.WriteString(" NOT NULL") } } if len(t.Pk) > 0 { b.WriteString(",\n PRIMARY KEY (" + strings.Join(quoteAll(t.Pk), ",") + ")") } b.WriteString("\n);") return b.String(), nil } func quoteAll(cols []string) []string { out := make([]string, len(cols)) for i, c := range cols { out[i] = `"` + c + `"` } return out } // castFor 返回列类型对应的 pg 类型转换后缀(用于导入时把字符串值转为正确类型)。 func castFor(udt string) string { switch udt { case "int2", "int4", "int8": return "bigint" case "numeric": return "numeric" case "float4", "float8": return "double precision" case "bool": return "boolean" case "timestamp": return "timestamp" case "timestamptz": return "timestamptz" case "date": return "date" case "time": return "time" case "json", "jsonb": return "jsonb" case "uuid": return "uuid" default: return "" // text / varchar / char 等字符串类型无需转换 } } // importRows 把一张表的数据批量 INSERT 进目标库(单表一个事务,每批多行)。 func importRows(db *sql.DB, t tableDump) error { if len(t.Rows) == 0 { return nil } hasIdentity := false for _, c := range t.Columns { if c.IsIdentity { hasIdentity = true break } } tx, err := db.Begin() if err != nil { return err } defer func() { _ = tx.Rollback() }() colList := make([]string, len(t.Columns)) casts := make([]string, len(t.Columns)) for i, c := range t.Columns { colList[i] = `"` + c.Name + `"` casts[i] = castFor(c.Type) } 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 ov := "" if hasIdentity { ov = " OVERRIDING SYSTEM VALUE" } sb.WriteString(fmt.Sprintf(`INSERT INTO "%s" (%s)%s VALUES `, t.Name, strings.Join(colList, ","), ov)) args := make([]any, 0, len(chunk)*len(t.Columns)) param := 1 for ri, row := range chunk { if ri > 0 { sb.WriteString(",") } sb.WriteString("(") for ci, val := range row { if ci > 0 { sb.WriteString(",") } if val == nil { sb.WriteString("NULL") } else { if casts[ci] != "" { sb.WriteString(fmt.Sprintf("$%d::%s", param, casts[ci])) } else { sb.WriteString(fmt.Sprintf("$%d", param)) } args = append(args, val) param++ } } sb.WriteString(")") } if _, err := tx.Exec(sb.String(), args...); err != nil { return err } } // 重置 identity 序列,避免后续自增插入与已导入的最大 ID 冲突。 for _, c := range t.Columns { if c.IsIdentity { seq := fmt.Sprintf("pg_get_serial_sequence('%s','%s')", t.Name, c.Name) if _, err := tx.Exec(fmt.Sprintf( `SELECT setval(%s, COALESCE((SELECT MAX("%s") FROM "%s"), 1))`, seq, c.Name, t.Name)); err != nil { log.Printf("! 重置序列 %s.%s 失败: %v", t.Name, c.Name, err) } } } return tx.Commit() }