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