// Command dbtool 是后台管理用的数据库导出 / 导入脚本。 // // 用途:把一个 PostgreSQL 库(fashion)的「结构 + 索引 + 约束 + 数据」导出为 // **单个纯 SQL 文件**,拷到另一台机器后用 dbtool 或 psql 灌入即可 —— 无需 pg_dump。 // // dbtool dump [-config ] [-out ] [-with-extension] 导出 // dbtool import [-config ] [-in ] 导入 // // 导入(目标机器,二选一): // // 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 ] [-out ] [-with-extension] 导出「结构 + 主键 / 索引 + 数据」为单个纯 SQL 文件 dbtool import [-config ] [-in ] 把 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, "'", "''") + "'" }