780 lines
27 KiB
Go
780 lines
27 KiB
Go
// Command dbtool 是后台管理用的数据库导出 / 导入脚本。
|
||
//
|
||
// 用途:把一个 PostgreSQL 库(fashion)的「结构 + 索引 + 约束 + 视图 + 数据」导出为
|
||
// **单个纯 SQL 文件**,拷到另一台机器后用 dbtool 或 psql 灌入即可 —— 无需 pg_dump。
|
||
//
|
||
// dbtool dump [-config <yml>] [-out <file.sql>] [-with-extension] [-clean] 导出
|
||
// dbtool import [-config <yml>] [-in <file.sql>] [-clean] 导入
|
||
//
|
||
// 导入(目标机器,二选一):
|
||
//
|
||
// 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)、
|
||
// 视图(按依赖顺序建;-clean 时按逆依赖顺序先删)、数据、序列当前值。
|
||
//
|
||
// 触发器 / 外键 / 注释 / 权限不导出。
|
||
//
|
||
// 视图为什么必须导出:公开读走 public_* 只读视图(见 db/migrations/2026-09-22-01),
|
||
// 而按下文口径「重建一个库就是 dump 出来再灌进去」——视图不进 dump,重建出的库就会缺视图,
|
||
// 公开查询直接报 relation does not exist。
|
||
//
|
||
// 目标库为空还是已有数据:
|
||
//
|
||
// 默认只 CREATE TABLE IF NOT EXISTS + INSERT,**不清表**,因此要求目标库为空
|
||
// (对已有数据的库会主键冲突)。要覆盖一个已有数据的库,用 -clean:
|
||
//
|
||
// dump -clean 在结构段前输出 DROP TABLE IF EXISTS ... CASCADE,只针对本次
|
||
// dump 里的表(与 pg_dump --clean 口径一致),清理语义写在文件里,
|
||
// 因此 psql -f 同样能灌进脏库。
|
||
// import -clean 导入前 DROP public 下全部表,不必重新导出文件;适合「手上只有
|
||
// 一份旧文件、目标库状态不明」的情况。
|
||
//
|
||
// 约束段另有一层幂等保护:先 DROP CONSTRAINT IF EXISTS 再 ADD。这样即使目标库里
|
||
// 已存在同名主键(例如早先被 AutoMigrate 建过、或上一次导入中断留下的),
|
||
// 也不会报 relation "<name>" already exists。
|
||
//
|
||
// 前置条件:目标库必须已启用 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"
|
||
"sort"
|
||
"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] [-clean]
|
||
导出「结构 + 主键 / 索引 / 视图 + 数据」为单个纯 SQL 文件
|
||
dbtool import [-config <yml>] [-in <file>] [-clean]
|
||
把 dump 出来的 SQL 文件灌入目标库
|
||
|
||
说明:
|
||
连接信息读取 configs/config.yml 的 database 段(可用 -config 覆盖)。
|
||
-with-extension(dump)在文件开头加 CREATE EXTENSION IF NOT EXISTS vector;
|
||
默认不加,此时目标库须已启用 pgvector,否则 vector 列建不出来。
|
||
-clean 有两个位置,都会丢弃目标库中对应表的既有数据:
|
||
dump -clean 在结构段前写 DROP TABLE IF EXISTS,使文件可灌进已有数据的库
|
||
import -clean 导入前先 DROP public 下全部表,适合「只有旧文件、目标库状态不明」
|
||
不加 -clean 时要求目标库为空(或至少同名表为空),否则会主键冲突。
|
||
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")
|
||
clean := fs.Bool("clean", false, "在结构段前加 DROP TABLE IF EXISTS,使脚本可灌进已有数据的库(这些表的既有数据会丢失)")
|
||
_ = 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, *clean)
|
||
|
||
tables, err := listTables(db)
|
||
if err != nil {
|
||
log.Fatalf("✗ 列举表失败: %v", err)
|
||
}
|
||
if len(tables) == 0 {
|
||
log.Fatalf("✗ 库中没有表(schema=public),请确认 -config 指向的库是否正确")
|
||
}
|
||
views, err := listViews(db)
|
||
if err != nil {
|
||
log.Fatalf("✗ 列举视图失败: %v", err)
|
||
}
|
||
|
||
cols := make(map[string][]column, len(tables))
|
||
|
||
// ① 结构(-clean 时先输出清理段)
|
||
fmt.Fprintln(w, "-- ==================== 结构 ====================")
|
||
if *clean {
|
||
writeClean(w, tables, views)
|
||
}
|
||
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)
|
||
}
|
||
// 视图放在全部基表之后建:视图可能引用基表,也可能引用另一个视图(listViews 已按依赖排序)。
|
||
for _, v := range views {
|
||
writeCreateView(w, v)
|
||
}
|
||
log.Printf("→ 结构:%d 张表、%d 个视图", len(tables), len(views))
|
||
|
||
// ② 数据
|
||
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)")
|
||
clean := fs.Bool("clean", false, "导入前先 DROP public 下全部表(覆盖式还原;目标库这些数据会丢失)")
|
||
_ = 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()
|
||
|
||
if *clean {
|
||
// 与脚本拼成一段再执行:两者落在同一个隐式事务里,要么都成功、要么都回滚。
|
||
script = append([]byte(dropAllTablesSQL+"\n"), script...)
|
||
log.Printf("→ 导入前先清空 public 下全部表(-clean)")
|
||
}
|
||
|
||
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)
|
||
}
|
||
|
||
// dropAllTablesSQL 删除 public 下全部基表(CASCADE 连带索引 / 约束 / 归属该表的序列)。
|
||
//
|
||
// 只删表,不动扩展:pgvector 的 vector 类型是 extension 对象,不属于任何表,因此
|
||
// DROP TABLE 之后扩展仍然可用,无需重新 CREATE EXTENSION。
|
||
const dropAllTablesSQL = `
|
||
DO $$
|
||
DECLARE
|
||
r record;
|
||
BEGIN
|
||
FOR r IN
|
||
SELECT c.relname FROM pg_class c
|
||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||
WHERE n.nspname = 'public' AND c.relkind = 'r'
|
||
LOOP
|
||
EXECUTE format('DROP TABLE IF EXISTS public.%I CASCADE', r.relname);
|
||
END LOOP;
|
||
END $$;`
|
||
|
||
// 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
|
||
})
|
||
}
|
||
|
||
// writeClean 输出清理段:DROP 掉本次要重建的视图与表,让脚本可以灌进「已有数据的库」。
|
||
//
|
||
// 只 DROP 本次 dump 里出现的对象(与 pg_dump --clean 口径一致)——目标库中多出来的对象保持不动。
|
||
// 视图必须**先于**它依赖的基表被删:views 是按依赖顺序(被依赖者在前)排的,故这里逆序删。
|
||
// 表用 CASCADE:索引 / 约束 / 归属该表的序列随表一起删除,不必逐个列举;
|
||
// 顺序也无所谓,因为 CASCADE 会处理依赖,且整段脚本在同一个隐式事务里执行。
|
||
func writeClean(w *bufio.Writer, tables []string, views []viewDef) {
|
||
fmt.Fprintln(w, "-- 清理(-clean):DROP 下列视图与表,目标库中这些表的既有数据将丢失")
|
||
for i := len(views) - 1; i >= 0; i-- {
|
||
fmt.Fprintf(w, "DROP VIEW IF EXISTS %s;\n", qname(views[i].Name))
|
||
}
|
||
for _, t := range tables {
|
||
fmt.Fprintf(w, "DROP TABLE IF EXISTS %s CASCADE;\n", qname(t))
|
||
}
|
||
fmt.Fprintln(w)
|
||
}
|
||
|
||
// viewDef 单个视图的名字与定义(定义取自 pg_get_viewdef,已去掉前导空白)。
|
||
type viewDef struct {
|
||
Name string
|
||
Def string
|
||
}
|
||
|
||
// listViews 返回 public 下全部视图,按**依赖顺序**排列(被依赖者在前)。
|
||
//
|
||
// 视图可以依赖基表,也可以依赖另一个视图;重建时必须先建被依赖者,否则 CREATE VIEW 会报
|
||
// relation does not exist。顺序由 pg_depend 推出的「视图→视图」依赖做拓扑排序得到。
|
||
func listViews(db *sql.DB) ([]viewDef, error) {
|
||
rows, err := db.Query(`
|
||
SELECT c.relname, pg_get_viewdef(c.oid, true)
|
||
FROM pg_class c
|
||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||
WHERE n.nspname = 'public' AND c.relkind = 'v'
|
||
ORDER BY c.relname`)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
|
||
defs := map[string]string{}
|
||
var names []string
|
||
for rows.Next() {
|
||
var name, def string
|
||
if err := rows.Scan(&name, &def); err != nil {
|
||
return nil, err
|
||
}
|
||
defs[name] = def
|
||
names = append(names, name)
|
||
}
|
||
if err := rows.Err(); err != nil {
|
||
return nil, err
|
||
}
|
||
if len(names) == 0 {
|
||
return nil, nil
|
||
}
|
||
|
||
deps, err := listViewDeps(db)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
ordered, err := topoSortViews(names, deps)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
out := make([]viewDef, 0, len(ordered))
|
||
for _, name := range ordered {
|
||
// pg_get_viewdef 末尾自带分号,去掉后由 writeCreateView 统一补,避免出现 ";;"。
|
||
body := strings.TrimSpace(strings.TrimRight(strings.TrimSpace(defs[name]), ";"))
|
||
out = append(out, viewDef{Name: name, Def: body})
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// listViewDeps 返回「视图 → 它直接引用的另一个视图」的依赖边(忽略对基表的引用)。
|
||
func listViewDeps(db *sql.DB) (map[string][]string, error) {
|
||
rows, err := db.Query(`
|
||
SELECT DISTINCT v.relname, d.relname
|
||
FROM pg_depend dep
|
||
JOIN pg_rewrite rw ON rw.oid = dep.objid
|
||
JOIN pg_class v ON v.oid = rw.ev_class
|
||
JOIN pg_class d ON d.oid = dep.refobjid
|
||
JOIN pg_namespace n ON n.oid = v.relnamespace
|
||
WHERE dep.classid = 'pg_rewrite'::regclass
|
||
AND dep.refclassid = 'pg_class'::regclass
|
||
AND dep.deptype = 'n'
|
||
AND v.relkind = 'v' AND d.relkind = 'v'
|
||
AND v.oid <> d.oid
|
||
AND n.nspname = 'public'`)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
|
||
deps := map[string][]string{}
|
||
for rows.Next() {
|
||
var v, d string
|
||
if err := rows.Scan(&v, &d); err != nil {
|
||
return nil, err
|
||
}
|
||
deps[v] = append(deps[v], d)
|
||
}
|
||
return deps, rows.Err()
|
||
}
|
||
|
||
// topoSortViews 对视图做拓扑排序:被依赖者在前。每步用名字排序兜底,保证导出可复现。
|
||
func topoSortViews(names []string, deps map[string][]string) ([]string, error) {
|
||
known := make(map[string]bool, len(names))
|
||
indeg := make(map[string]int, len(names))
|
||
for _, n := range names {
|
||
known[n] = true
|
||
indeg[n] = 0
|
||
}
|
||
|
||
adj := map[string][]string{}
|
||
for v, ds := range deps {
|
||
if !known[v] {
|
||
continue
|
||
}
|
||
for _, d := range ds {
|
||
if !known[d] || d == v {
|
||
continue
|
||
}
|
||
adj[d] = append(adj[d], v) // d 必须先建,建完 d 才轮到 v
|
||
indeg[v]++
|
||
}
|
||
}
|
||
|
||
ready := make([]string, 0, len(names))
|
||
for _, n := range names {
|
||
if indeg[n] == 0 {
|
||
ready = append(ready, n)
|
||
}
|
||
}
|
||
sort.Strings(ready)
|
||
|
||
out := make([]string, 0, len(names))
|
||
for len(ready) > 0 {
|
||
n := ready[0]
|
||
ready = ready[1:]
|
||
out = append(out, n)
|
||
|
||
next := append([]string(nil), adj[n]...)
|
||
sort.Strings(next)
|
||
for _, m := range next {
|
||
indeg[m]--
|
||
if indeg[m] == 0 {
|
||
ready = append(ready, m)
|
||
}
|
||
}
|
||
sort.Strings(ready)
|
||
}
|
||
if len(out) != len(names) {
|
||
return nil, fmt.Errorf("视图依赖存在环,无法排序(已完成 %d/%d)", len(out), len(names))
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// writeCreateView 输出一个视图的 CREATE VIEW。
|
||
// 定义来自 pg_get_viewdef(形如「 SELECT ...」),去掉前导空白后接到 AS 之后即可。
|
||
func writeCreateView(w *bufio.Writer, v viewDef) {
|
||
fmt.Fprintf(w, "CREATE VIEW %s AS\n%s;\n\n", qname(v.Name), v.Def)
|
||
}
|
||
|
||
// writeHeader 写文件头与几个会话设置。
|
||
func writeHeader(w *bufio.Writer, dcfg *config.DatabaseConfig, withExtension, clean 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, "--")
|
||
if clean {
|
||
fmt.Fprintln(w, "-- 注意:本文件由 dump -clean 生成,结构段前会 DROP 同名表 —— 目标库中这些表的既有数据将丢失。")
|
||
} else {
|
||
fmt.Fprintln(w, "-- 注意:本文件不含 DROP,目标库必须是空库;重复执行会主键冲突。")
|
||
}
|
||
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
|
||
}
|
||
// 先 DROP CONSTRAINT IF EXISTS 再 ADD,让这一段幂等:
|
||
// 目标库若已存在同名约束(例如早先被 AutoMigrate 建过、或导入中断留下的),
|
||
// 不先删就会报 relation "<name>" already exists —— 约束底层的索引名在 schema 内唯一。
|
||
fmt.Fprintf(w, "ALTER TABLE ONLY %s DROP CONSTRAINT IF EXISTS %s;\n", qname(table), qi(name))
|
||
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, "'", "''") + "'" }
|