// Command dbtool 是后台管理用的数据库导出脚本。 // // 用途:把一个 PostgreSQL 库(fashion)的「结构 + 索引 + 约束 + 数据」导出为 // **单个纯 SQL 文件**,拷到另一台机器后直接用 psql 灌入即可 —— 无需 pg_dump。 // // dbtool dump [-config ] [-out ] [-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 ] [-out ] [-with-extension] 导出「结构 + 主键 / 索引 + 数据」为单个纯 SQL 文件 导入(目标机器): psql -U fashion -d fashion -v ON_ERROR_STOP=1 -f 说明: 连接信息读取 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, "'", "''") + "'" }