Files
backend_v2/cmd/dbtool/main.go
toom1996 e5ce8aed3e update
2026-09-21 19:42:17 +08:00

549 lines
18 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)的「结构 + 索引 + 约束 + 数据」导出为
// **单个纯 SQL 文件**,拷到另一台机器后用 dbtool 或 psql 灌入即可 —— 无需 pg_dump。
//
// dbtool dump [-config <yml>] [-out <file.sql>] [-with-extension] 导出
// dbtool import [-config <yml>] [-in <file.sql>] 导入
//
// 导入(目标机器,二选一):
//
// dbtool import -in db_dump.sql
// psql -U fashion -d fashion -v ON_ERROR_STOP=1 -f db_dump.sql
//
// 结构从哪来:**唯一入口就是本工具的 dump 输出**。项目已不再在服务启动时自动迁移表结构
// (原 database.AutoMigrate / EnsureDedupSchema 已移除),重建一个库就是「dump 出来再灌进去」。
//
// 数据为什么用 INSERT 而不是 COPY:pg_dump 用的是 COPY ... FROM stdin,它依赖 PostgreSQL
// 前端的 copy 子协议,而 pgx 底层(pgconn)在简单查询协议的多语句执行中并不处理
// CopyInResponse(源码中无该分支),Go 侧无法执行含 COPY 的脚本。改用多行 INSERT 后,
// 同一份文件既能被 dbtool import 执行,也能被 psql 执行。
//
// 与 pg_dump 的关系:本工具不依赖任何外部二进制,等于把「读系统目录 → 拼 DDL → 导数据」
// 自己实现一遍,因此**只覆盖本项目实际用到的对象**:
//
// 表、列(类型含修饰符 / 默认值 / identity / NOT NULL)、主键、索引(含 pgvector 的 HNSW)、
// 数据、序列当前值。
//
// 视图 / 触发器 / 外键 / 注释 / 权限不导出(当前库中也不存在)。
//
// 前置条件:目标库必须为空,且已启用 pgvector 扩展,否则 vector(64) 列建不出来。
// 默认**不**导出扩展语句(可用 -with-extension 带上)。若目标容器由 scripts/pgvector 的
// compose 启动,initdb/01-extensions.sql 会在数据卷首次初始化时自动创建扩展。
//
// 连接信息来自 configs/config.yml 的 database 段(同后端服务),可用 -config 指定其它配置。
package main
import (
"bufio"
"context"
"database/sql"
"flag"
"fmt"
"log"
"os"
"strings"
"time"
"github.com/jackc/pgx/v5/stdlib"
"fashionapi/internal/config"
)
// column 是单列的结构信息,用于生成 CREATE TABLE、INSERT 列清单与序列重置语句。
type column struct {
Name string
// Type 为 format_type(atttypid, atttypmod) 的结果,**含**类型修饰符:
// vector(64) / character varying(255) / bigint。
// 不能用 information_schema.udt_name —— 它会丢掉修饰符(vector(64) 变成 vector),
// 建出来的列没有维度,随后 HNSW 索引会报 "column does not have dimensions"。
Type string
// NotNull 取自 pg_attribute.attnotnull。
NotNull bool
// Default 取自 pg_get_expr(adbin, adrelid),如 0 / ''::character varying / nextval('x'::regclass)。
Default string
// Identity 取自 pg_attribute.attidentity:'d'=BY DEFAULT、'a'=ALWAYS、''=非 identity。
Identity string
// AutoIncrement 表示该列由序列自动赋值(identity 列,或默认值是 nextval 的 serial 列),
// 导出后会追加 setval 把序列推到 MAX(col)+1。
//
// 判定必须读 attidentity:identity 列的 column_default 是 NULL,
// 用 column_default LIKE 'nextval(%' 去猜会把它们全部漏判,重建出的表 id 就没有自增。
AutoIncrement bool
}
// insertBatchRows 单条 INSERT 里最多写多少行。
// 批量写可以让文件更紧凑、导入更快;值都是字面量而非绑定参数,不受参数个数上限限制。
const insertBatchRows = 100
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>] [-with-extension]
导出「结构 + 主键 / 索引 + 数据」为单个纯 SQL 文件
dbtool import [-config <yml>] [-in <file>]
把 dump 出来的 SQL 文件灌入目标库
说明:
连接信息读取 configs/config.yml 的 database 段(可用 -config 覆盖)。
-with-extension 会在导出文件开头加 CREATE EXTENSION IF NOT EXISTS vector;
默认不加,此时目标库须已启用 pgvector,否则 vector 列建不出来。
目标库必须是**空库**(脚本按 CREATE TABLE IF NOT EXISTS + INSERT 写入,不会清表)。
import 等价于 psql -v ON_ERROR_STOP=1 -f;整个脚本在一个隐式事务里执行,出错整体回滚。`)
}
// 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 导出「结构 → 数据 → 主键/索引 → 序列」四段到单个 SQL 文件。
//
// 段落顺序刻意与 pg_dump 一致:先建表、再灌数据、最后建约束与索引。
// 索引放在数据之后建,既快又不会在导入过程中被反复维护。
func runDump(args []string) {
fs := flag.NewFlagSet("dump", flag.ExitOnError)
out := fs.String("out", "db_dump.sql", "导出文件路径")
cfgPath := fs.String("config", "", "配置文件路径(默认 configs/config.yml)")
withExt := fs.Bool("with-extension", false, "在文件开头加 CREATE EXTENSION IF NOT EXISTS vector")
_ = fs.Parse(args)
db, dcfg, err := openDB(*cfgPath)
if err != nil {
log.Fatalf("✗ %v", err)
}
defer db.Close()
f, err := os.Create(*out)
if err != nil {
log.Fatalf("✗ 创建文件 %s 失败: %v", *out, err)
}
defer f.Close()
w := bufio.NewWriter(f)
writeHeader(w, dcfg, *withExt)
tables, err := listTables(db)
if err != nil {
log.Fatalf("✗ 列举表失败: %v", err)
}
if len(tables) == 0 {
log.Fatalf("✗ 库中没有表(schema=public),请确认 -config 指向的库是否正确")
}
cols := make(map[string][]column, len(tables))
// ① 结构
fmt.Fprintln(w, "-- ==================== 结构 ====================")
for _, t := range tables {
cs, err := listColumns(db, t)
if err != nil {
log.Fatalf("✗ 读取 %s 的列失败: %v", t, err)
}
cols[t] = cs
writeCreateTable(w, t, cs)
}
log.Printf("→ 结构:%d 张表", len(tables))
// ② 数据
fmt.Fprintln(w, "-- ==================== 数据 ====================")
var total int64
for _, t := range tables {
n, err := writeInsertData(w, db, t, cols[t])
if err != nil {
log.Fatalf("✗ 导出 %s 数据失败: %v", t, err)
}
total += n
log.Printf(" ✓ %s: %d 行", t, n)
}
// ③ 主键 / 索引
fmt.Fprintln(w, "-- ==================== 主键 / 索引 ====================")
for _, t := range tables {
if err := writeConstraints(w, db, t); err != nil {
log.Fatalf("✗ 导出 %s 约束失败: %v", t, err)
}
if err := writeIndexes(w, db, t); err != nil {
log.Fatalf("✗ 导出 %s 索引失败: %v", t, err)
}
}
// ④ 序列当前值:INSERT 写了显式 id,不会推进序列,必须手工推到 MAX+1。
fmt.Fprintln(w, "-- ==================== 序列当前值 ====================")
for _, t := range tables {
writeSequenceResets(w, t, cols[t])
}
if err := w.Flush(); err != nil {
log.Fatalf("✗ 写文件 %s 失败: %v", *out, err)
}
log.Printf("✓ 已导出 %d 张表、%d 行 -> %s", len(tables), total, *out)
}
// runImport 把 dump 出来的 SQL 文件整体灌入目标库。
//
// 直接用 pgx 的「简单查询协议」把整个脚本一次性发给服务端:
// - 服务端自己解析多语句,客户端无需写 SQL 解析器;
// - 不含 COPY,因此不涉及 copy 子协议;
// - 多语句在同一个隐式事务里执行,中途出错整体回滚,不会留下半截的库。
func runImport(args []string) {
fs := flag.NewFlagSet("import", flag.ExitOnError)
in := fs.String("in", "db_dump.sql", "导入文件路径")
cfgPath := fs.String("config", "", "配置文件路径(默认 configs/config.yml)")
_ = fs.Parse(args)
script, err := os.ReadFile(*in)
if err != nil {
log.Fatalf("✗ 读文件 %s 失败: %v", *in, err)
}
db, dcfg, err := openDB(*cfgPath)
if err != nil {
log.Fatalf("✗ %v", err)
}
defer db.Close()
log.Printf("→ 导入 %s(%d KB)-> %s ...", *in, len(script)/1024, dcfg.Addr())
if err := execScript(context.Background(), db, string(script)); err != nil {
log.Fatalf("✗ 导入失败(已整体回滚): %v", err)
}
log.Printf("✓ 导入完成(%d KB)", len(script)/1024)
}
// execScript 用 pgx 简单查询协议执行整段脚本。
//
// 必须走底层 *pgx.Conn:database/sql 的 Exec 用扩展协议,不支持一次发多条语句。
func execScript(ctx context.Context, db *sql.DB, script string) error {
conn, err := db.Conn(ctx)
if err != nil {
return fmt.Errorf("获取连接失败: %w", err)
}
defer conn.Close()
return conn.Raw(func(dc any) error {
std, ok := dc.(*stdlib.Conn)
if !ok {
return fmt.Errorf("驱动连接类型异常: %T", dc)
}
_, err := std.Conn().PgConn().Exec(ctx, script).ReadAll()
return err
})
}
// writeHeader 写文件头与几个会话设置。
func writeHeader(w *bufio.Writer, dcfg *config.DatabaseConfig, withExtension bool) {
fmt.Fprintf(w, "-- dbtool 导出:%s\n", dcfg.Addr())
fmt.Fprintf(w, "-- 生成时间:%s\n", time.Now().Format(time.RFC3339))
fmt.Fprintln(w, "--")
fmt.Fprintln(w, "-- 导入(二选一):")
fmt.Fprintf(w, "-- dbtool import -in <本文件>\n")
fmt.Fprintf(w, "-- psql -U %s -d %s -v ON_ERROR_STOP=1 -f <本文件>\n", dcfg.User, dcfg.Name)
fmt.Fprintln(w, "--")
fmt.Fprintln(w, "-- 注意:目标库必须是空库;本脚本不清表,重复执行会主键冲突。")
fmt.Fprintln(w, "SET client_encoding = 'UTF8';")
fmt.Fprintln(w, "SET client_min_messages = warning;")
fmt.Fprintln(w)
if withExtension {
fmt.Fprintln(w, "CREATE EXTENSION IF NOT EXISTS vector;")
fmt.Fprintln(w)
}
}
// listTables 返回 public 下全部基表名(按名字排序,保证导出可复现)。
func listTables(db *sql.DB) ([]string, error) {
rows, err := db.Query(`
SELECT c.relname
FROM pg_class c
JOIN pg_namespace n ON n.oid = c.relnamespace
WHERE n.nspname = 'public' AND c.relkind = 'r'
ORDER BY c.relname`)
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()
}
// listColumns 读取单表全部列的完整结构信息。
func listColumns(db *sql.DB, table string) ([]column, error) {
q := fmt.Sprintf(`
SELECT a.attname,
format_type(a.atttypid, a.atttypmod),
a.attnotnull,
COALESCE(a.attidentity, ''),
COALESCE(pg_get_expr(d.adbin, d.adrelid), '')
FROM pg_attribute a
LEFT JOIN pg_attrdef d ON d.adrelid = a.attrelid AND d.adnum = a.attnum
WHERE a.attrelid = %s::regclass AND a.attnum > 0 AND NOT a.attisdropped
ORDER BY a.attnum`, ql(qname(table)))
rows, err := db.Query(q)
if err != nil {
return nil, err
}
defer rows.Close()
var out []column
for rows.Next() {
var c column
if err := rows.Scan(&c.Name, &c.Type, &c.NotNull, &c.Identity, &c.Default); err != nil {
return nil, err
}
c.AutoIncrement = c.Identity != "" || isSerial(c)
out = append(out, c)
}
return out, rows.Err()
}
// writeCreateTable 由列信息生成 CREATE TABLE。
//
// 两类自增列统一归一为 GENERATED BY DEFAULT AS IDENTITY:
// - 原生 identity 列(attidentity 非空);
// - serial 列(默认值是 nextval,如 image_embeddings.id)。
//
// 归一后不必再单独导出 CREATE SEQUENCE —— identity 子句会自动建序列,
// 也避免了「默认值引用一个还不存在的序列」导致建表失败。
func writeCreateTable(w *bufio.Writer, table string, cols []column) {
fmt.Fprintf(w, "CREATE TABLE IF NOT EXISTS %s (\n", qname(table))
for i, c := range cols {
fmt.Fprintf(w, " %s %s", qi(c.Name), c.Type)
switch {
case c.AutoIncrement:
fmt.Fprint(w, " GENERATED BY DEFAULT AS IDENTITY")
case c.Default != "":
fmt.Fprintf(w, " DEFAULT %s", c.Default)
}
// identity 列本身即 NOT NULL,不重复声明。
if c.NotNull && !c.AutoIncrement {
fmt.Fprint(w, " NOT NULL")
}
if i < len(cols)-1 {
fmt.Fprint(w, ",")
}
fmt.Fprintln(w)
}
fmt.Fprintln(w, ");")
fmt.Fprintln(w)
}
// writeInsertData 用多行 INSERT 导出单表数据,返回行数。
//
// 值一律按服务端文本表示写成字面量:NULL 直接写 NULL,其余走 escapeLit。
// 这样既绕开了 COPY 的 copy 子协议限制,也让文件对人类可读、可手工改。
func writeInsertData(w *bufio.Writer, db *sql.DB, table string, cols []column) (int64, error) {
names := make([]string, len(cols))
for i, c := range cols {
names[i] = qi(c.Name)
}
colList := strings.Join(names, ", ")
rows, err := db.Query(fmt.Sprintf("SELECT %s FROM %s", colList, qname(table)))
if err != nil {
return 0, err
}
defer rows.Close()
raw := make([]sql.RawBytes, len(cols))
ptrs := make([]any, len(cols))
for i := range raw {
ptrs[i] = &raw[i]
}
var count int64
var pending int // 当前这条 INSERT 已写入的行数
for rows.Next() {
if err := rows.Scan(ptrs...); err != nil {
return count, err
}
if pending == 0 {
fmt.Fprintf(w, "INSERT INTO %s (%s) VALUES\n", qname(table), colList)
} else {
fmt.Fprint(w, ",\n")
}
fmt.Fprint(w, " (")
for i := range raw {
if i > 0 {
fmt.Fprint(w, ", ")
}
if raw[i] == nil {
fmt.Fprint(w, "NULL")
} else {
fmt.Fprint(w, escapeLit(string(raw[i])))
}
}
fmt.Fprint(w, ")")
pending++
count++
if pending == insertBatchRows {
fmt.Fprintln(w, ";")
pending = 0
}
}
if err := rows.Err(); err != nil {
return count, err
}
if pending > 0 {
fmt.Fprintln(w, ";")
}
fmt.Fprintln(w)
return count, nil
}
// writeConstraints 导出主键等约束定义(pg_get_constraintdef 给出完整定义,无需自己拼)。
func writeConstraints(w *bufio.Writer, db *sql.DB, table string) error {
q := fmt.Sprintf(`
SELECT con.conname, pg_get_constraintdef(con.oid)
FROM pg_constraint con
WHERE con.conrelid = %s::regclass
ORDER BY CASE con.contype
WHEN 'p' THEN 1 WHEN 'u' THEN 2 WHEN 'c' THEN 3 WHEN 'x' THEN 4 WHEN 'f' THEN 5 ELSE 9
END,
con.conname`, ql(qname(table)))
rows, err := db.Query(q)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var name, def string
if err := rows.Scan(&name, &def); err != nil {
return err
}
fmt.Fprintf(w, "ALTER TABLE ONLY %s ADD CONSTRAINT %s %s;\n", qname(table), qi(name), def)
}
fmt.Fprintln(w)
return rows.Err()
}
// writeIndexes 导出不承载约束的索引(含 pgvector 的 HNSW)。
//
// 排除承载约束的索引:主键索引已由 writeConstraints 以 ADD CONSTRAINT 的形式重建,
// 再 CREATE INDEX 一次会重复。
func writeIndexes(w *bufio.Writer, db *sql.DB, table string) error {
q := fmt.Sprintf(`
SELECT i.relname, pg_get_indexdef(i.oid)
FROM pg_index x
JOIN pg_class i ON i.oid = x.indexrelid
WHERE x.indrelid = %s::regclass
AND NOT EXISTS (SELECT 1 FROM pg_constraint con WHERE con.conindid = x.indexrelid)
ORDER BY i.relname`, ql(qname(table)))
rows, err := db.Query(q)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var name, def string
if err := rows.Scan(&name, &def); err != nil {
return err
}
fmt.Fprintf(w, "%s;\n", injectIfNotExists(def))
}
fmt.Fprintln(w)
return rows.Err()
}
// writeSequenceResets 为每个自增列把序列推到 MAX(col)+1。
//
// INSERT 写入的是显式 id,不会推进序列;不重置的话,后续 INSERT 会从 1 开始并与存量主键冲突。
// setval 用 (值, false) 形式:false 表示「这个值还没被取走」,故下一次 nextval 正好是 max+1;
// 空表时为 1,即从 1 开始。
func writeSequenceResets(w *bufio.Writer, table string, cols []column) {
for _, c := range cols {
if !c.AutoIncrement {
continue
}
fmt.Fprintf(w,
"SELECT pg_catalog.setval(pg_get_serial_sequence(%s, %s), COALESCE((SELECT MAX(%s) FROM %s), 0) + 1, false);\n",
ql(qname(table)), ql(c.Name), qi(c.Name), qname(table))
}
}
// isSerial 判断是否为 serial 列(默认值 nextval,且类型为整型)。
func isSerial(c column) bool {
if !strings.HasPrefix(c.Default, "nextval(") {
return false
}
switch c.Type {
case "bigint", "integer", "smallint":
return true
default:
return false
}
}
// escapeLit 把值编码成 SQL 字符串字面量,用 E'...' 转义串语法。
//
// 选 E'...' 而不是普通 '...':普通字面量里反斜杠的含义取决于会话的
// standard_conforming_strings,而 E'...' 下反斜杠**总是**转义符,行为与设置无关。
// 因此只需处理反斜杠与单引号两个字符,任何文本(含换行、制表符、引号、反斜杠)都能安全往返。
func escapeLit(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, `'`, `''`)
return "E'" + s + "'"
}
// injectIfNotExists 给 pg_get_indexdef 的输出补上 IF NOT EXISTS。
// pg_get_indexdef 不会输出它(那是 pg_dump 的清理语义),补上可让索引段落可重复执行。
func injectIfNotExists(def string) string {
if rest, ok := strings.CutPrefix(def, "CREATE UNIQUE INDEX "); ok {
return "CREATE UNIQUE INDEX IF NOT EXISTS " + rest
}
if rest, ok := strings.CutPrefix(def, "CREATE INDEX "); ok {
return "CREATE INDEX IF NOT EXISTS " + rest
}
return def
}
// qname 返回 schema 限定且带引号的表名:public."brands"。
func qname(table string) string { return "public." + qi(table) }
// qi 安全包裹 SQL 标识符(表名 / 列名 / 索引名)。
func qi(s string) string { return `"` + strings.ReplaceAll(s, `"`, `""`) + `"` }
// ql 安全包裹 SQL 字符串字面量(普通形式,用于表名/列名等受控内容)。
func ql(s string) string { return "'" + strings.ReplaceAll(s, "'", "''") + "'" }