Files
backend_v2/cmd/dbtool/main.go
toom1996 299fd974df updae
2026-09-21 19:19:03 +08:00

485 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)的「结构 + 索引 + 约束 + 数据」导出为
// **单个纯 SQL 文件**,拷到另一台机器后直接用 psql 灌入即可 —— 无需 pg_dump。
//
// dbtool dump [-config <yml>] [-out <file.sql>] [-with-extension]
//
// 导入方式(目标机器):
//
// psql -U fashion -d fashion -v ON_ERROR_STOP=1 -f db_dump.sql
//
// 结构从哪来:**唯一入口就是本工具的 dump 输出**。项目已不再在服务启动时自动迁移表结构
// (原 database.AutoMigrate / EnsureDedupSchema 已移除),重建一个库就是「dump 出来再灌进去」。
//
// 为什么没有 import 子命令:数据段用的是 SQL 标准的 COPY ... FROM stdin,它依赖 PostgreSQL
// 前端的 copy 子协议,Go 的 database/sql 无法执行这类脚本。导入统一交给 psql。
//
// 与 pg_dump 的关系:本工具不依赖任何外部二进制,等于把「读系统目录 → 拼 DDL → 导数据」
// 自己实现一遍,因此**只覆盖本项目实际用到的对象**:
//
// 表、列(类型含修饰符 / 默认值 / identity / NOT NULL)、主键、索引(含 pgvector 的 HNSW)、
// 数据(COPY)、序列当前值。
//
// 视图 / 触发器 / 外键 / 注释 / 权限不导出(当前库中也不存在)。
//
// 前置条件:目标库必须已启用 pgvector 扩展,否则 vector(64) 列建不出来。
// 默认**不**导出扩展语句(可用 -with-extension 带上)。若目标容器由 scripts/pgvector 的
// compose 启动,initdb/01-extensions.sql 会在数据卷首次初始化时自动创建扩展。
//
// 连接信息来自 configs/config.yml 的 database 段(同后端服务),可用 -config 指定其它配置。
package main
import (
"bufio"
"database/sql"
"flag"
"fmt"
"log"
"os"
"strings"
"time"
_ "github.com/jackc/pgx/v5/stdlib"
"fashionapi/internal/config"
)
// column 是单列的结构信息,用于生成 CREATE TABLE、COPY 列清单与序列重置语句。
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
}
func main() {
if len(os.Args) < 2 {
usage()
os.Exit(2)
}
switch os.Args[1] {
case "dump":
runDump(os.Args[2:])
default:
usage()
os.Exit(2)
}
}
func usage() {
fmt.Println(`dbtool - 后台数据库导出脚本(PostgreSQL)
用法:
dbtool dump [-config <yml>] [-out <file>] [-with-extension]
导出「结构 + 主键 / 索引 + 数据」为单个纯 SQL 文件
导入(目标机器):
psql -U fashion -d fashion -v ON_ERROR_STOP=1 -f <file>
说明:
连接信息读取 configs/config.yml 的 database 段(可用 -config 覆盖)。
-with-extension 会在文件开头加 CREATE EXTENSION IF NOT EXISTS vector;
默认不加,此时目标库须已启用 pgvector,否则 vector 列建不出来。
目标库必须是**空库**(脚本按 CREATE TABLE IF NOT EXISTS + COPY 写入,不会清表)。`)
}
// 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)
if *withExt {
fmt.Fprintln(w, "CREATE EXTENSION IF NOT EXISTS vector;")
fmt.Fprintln(w)
}
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 := writeCopyData(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)
}
}
// ④ 序列当前值:COPY 写了显式 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)
}
// writeHeader 写文件头与几个影响字面量解析的会话设置。
func writeHeader(w *bufio.Writer, dcfg *config.DatabaseConfig) {
fmt.Fprintf(w, "-- dbtool 导出:%s\n", dcfg.Addr())
fmt.Fprintf(w, "-- 生成时间:%s\n", time.Now().Format(time.RFC3339))
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 standard_conforming_strings = on;")
fmt.Fprintln(w, "SET check_function_bodies = false;")
fmt.Fprintln(w, "SET client_min_messages = warning;")
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)
}
// writeCopyData 用 COPY ... FROM stdin 导出单表数据,返回行数。
//
// 全部值按服务端文本表示搬运;NULL 写成 \N,其余转义反斜杠 / 制表符 / 换行 / 回车。
func writeCopyData(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, ", ")
fmt.Fprintf(w, "COPY %s (%s) FROM stdin;\n", qname(table), colList)
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
for rows.Next() {
if err := rows.Scan(ptrs...); err != nil {
return count, err
}
for i := range raw {
if i > 0 {
w.WriteByte('\t')
}
if raw[i] == nil {
w.WriteString(`\N`)
continue
}
w.WriteString(escapeCopy(raw[i]))
}
w.WriteByte('\n')
count++
}
if err := rows.Err(); err != nil {
return count, err
}
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。
//
// COPY 写入的是显式 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
}
}
// escapeCopy 转义 COPY 文本格式里的特殊字符。
// 先处理反斜杠,保证字面量 `\N` 被写成 `\\N` 而不会变成 NULL。
func escapeCopy(b []byte) string {
var sb strings.Builder
sb.Grow(len(b) + 8)
for _, c := range b {
switch c {
case '\\':
sb.WriteString(`\\`)
case '\t':
sb.WriteString(`\t`)
case '\n':
sb.WriteString(`\n`)
case '\r':
sb.WriteString(`\r`)
default:
sb.WriteByte(c)
}
}
return sb.String()
}
// 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, "'", "''") + "'" }