Files
backend_v2/cmd/dbtool/main.go
toom1996 dffcf9f9e5 update
2026-09-21 18:42:50 +08:00

506 lines
16 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 导出当前库全部表的数据 + 列元信息为单个 JSON 文件
// dbtool import -in db_dump.json 把 JSON 数据灌入目标库(结构先由 AutoMigrate 收敛)
//
// 结构在哪里定义:**只在代码里** —— model 的 GORM tag + database.EnsureDedupSchema。
// import 会先调 database.AutoMigrate + EnsureDedupSchema,把目标库收敛到与当前代码
// 一致(建表 / 加列 / 加索引 / 建 pgvector 扩展 / 建 HNSW 索引),然后再灌数据。
//
// 历史实现曾让 import 自己读 information_schema 拼 DDL 来重建表,那条路必然丢信息
// (类型修饰符 vector(64)、identity 自增、索引、扩展),已移除。
//
// 连接信息来自 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"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"fashionapi/internal/config"
"fashionapi/internal/database"
)
// columnMeta 是单列的元信息,用于导入时的类型转换与 identity 判定。
type columnMeta struct {
Name string `json:"name"`
// Type 为 udt_name,如 int8 / varchar / timestamptz / vector。
// 注意:它**不含**类型修饰符(vector(64) 会退化成 vector),
// 因此不可用于重建 DDL —— 结构一律交给 database.AutoMigrate。
Type string `json:"type"`
IsNullable bool `json:"is_nullable"`
// IsIdentity 取 information_schema.columns.is_identity。
// 必须读这一列:identity 列的 column_default 是 NULL,
// 用 column_default LIKE 'nextval(%' 去猜会把它全部漏判为「非自增」,
// 于是重建出的表 id 没有默认值、插入即违反 NOT NULL。
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>] 导出全库数据到 JSON
dbtool import [-config <yml>] [-in <file>] 把 JSON 数据灌入目标库
说明:
连接信息读取 configs/config.yml 的 database 段(可用 -config 覆盖)。
import 会先把目标库结构收敛到与当前代码一致(AutoMigrate + EnsureDedupSchema),
再清空同名表并按 JSON 重灌数据。结构定义只在代码里,dump 文件不含 DDL。`)
}
// 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 把 JSON 中的数据灌入目标库。
//
// 结构不在这里重建:先调 ensureSchema 用后端的 AutoMigrate + EnsureDedupSchema
// 把目标库收敛到与当前代码一致(建表 / 加列 / 加索引 / 建 pgvector 扩展 / 建 HNSW 索引),
// 本函数只负责搬数据。这样结构只有一个来源(model + EnsureDedupSchema),
// 不会再出现「自拼 DDL 丢类型修饰符 / 丢 identity 自增 / 丢索引 / 不建扩展」那一类问题。
//
// 注意:导入先清空目标同名表再写入,属「以 dump 为准的整体覆盖」。
func runImport(args []string) {
fs := flag.NewFlagSet("import", flag.ExitOnError)
in := fs.String("in", "db_dump.json", "导入文件路径")
cfgPath := fs.String("config", "", "配置文件路径(默认 configs/config.yml)")
_ = 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, dcfg, err := openDB(*cfgPath)
if err != nil {
log.Fatalf("✗ %v", err)
}
defer db.Close()
// ① 结构:与后端同一套定义(含 CREATE EXTENSION vector、HNSW 索引、老库兼容迁移)。
log.Printf("→ 收敛目标库结构(AutoMigrate + EnsureDedupSchema)...")
if err := ensureSchema(dcfg); err != nil {
log.Fatalf("✗ 结构初始化失败: %v", err)
}
// ② 数据:关闭外键 / 触发器,避免插入顺序受约束(整库覆盖无需保序)。
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 {
// 先清空目标表(约束已由 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))
}
// ensureSchema 用后端的 database.AutoMigrate + EnsureDedupSchema 把目标库结构
// 收敛到与当前代码一致。这是「结构定义只有一个来源」的落点:
// 表 / 列来自 model 的 GORM tag,扩展与 HNSW 索引来自 EnsureDedupSchema。
func ensureSchema(dcfg *config.DatabaseConfig) error {
gdb, err := gorm.Open(postgres.Open(dcfg.DSN()), &gorm.Config{
Logger: logger.Default.LogMode(logger.Warn),
SkipDefaultTransaction: true,
})
if err != nil {
return fmt.Errorf("连接数据库失败: %w", err)
}
sqlDB, err := gdb.DB()
if err != nil {
return fmt.Errorf("获取底层连接池失败: %w", err)
}
defer sqlDB.Close()
if err := database.AutoMigrate(gdb); err != nil {
return fmt.Errorf("AutoMigrate: %w", err)
}
if err := database.EnsureDedupSchema(gdb); err != nil {
return fmt.Errorf("EnsureDedupSchema: %w", err)
}
return nil
}
// 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) {
// is_identity 必须直接读 information_schema.columns.is_identity:
// identity 列的 column_default 是 NULL,用 column_default LIKE 'nextval(%' 判断
// 会把 GENERATED BY DEFAULT AS IDENTITY 的列全部漏判成「非自增」。
colQuery := strings.Replace(`
SELECT c.column_name,
c.udt_name,
(c.is_nullable = 'YES'),
COALESCE(c.column_default, ''),
(c.is_identity = 'YES')
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, isIdentity bool
if err := rows.Scan(&name, &udt, &nullable, &def, &isIdentity); err != nil {
return nil, nil, err
}
cols = append(cols, columnMeta{
Name: name, Type: udt, IsNullable: nullable, IsIdentity: isIdentity, 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()
}
// 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()
}