485 lines
15 KiB
Go
485 lines
15 KiB
Go
// 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, "'", "''") + "'" }
|