536 lines
15 KiB
Go
536 lines
15 KiB
Go
// 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()
|
||
}
|