Files
backend_v2/cmd/dbtool/main.go
toom1996 f7ae917603 update
2026-09-17 22:51:00 +08:00

536 lines
15 KiB
Go
Raw 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 是后台管理用的数据库一键迁移脚本。
//
// 用途:在不同开发电脑 / 环境之间搬运 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 <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 仅导数据(表结构须已由后端 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()
}