From 6d3d23c3f9df182694569cf44832518861a4306c Mon Sep 17 00:00:00 2001 From: XingfenD Date: Mon, 14 Sep 2026 19:36:21 +0800 Subject: [PATCH] feat(db): ordered migration system with advisory lock and baseline detection --- backend/internal/db/db.go | 130 +++++++++++++++++- .../0001_baseline.sql} | 0 2 files changed, 125 insertions(+), 5 deletions(-) rename backend/internal/db/{schema.sql => migrations/0001_baseline.sql} (100%) diff --git a/backend/internal/db/db.go b/backend/internal/db/db.go index ce3228d..d36cf80 100644 --- a/backend/internal/db/db.go +++ b/backend/internal/db/db.go @@ -2,14 +2,26 @@ package db import ( "context" - _ "embed" + "embed" "fmt" + "io/fs" + "log" + "regexp" + "sort" + "strings" "github.com/jackc/pgx/v5/pgxpool" ) -//go:embed schema.sql -var schema string +//go:embed migrations +var migrationsFS embed.FS + +// advisoryLockKey is a fixed int64 used with pg_advisory_lock to serialize +// migrations across --scale api=N replicas. Value is arbitrary but must be +// unique within the database (pick a project-specific constant). +const advisoryLockKey int64 = 0x424C4D49 // "BLMI" + +var migrationNameRe = regexp.MustCompile(`^\d{4}_[a-z0-9_]+\.sql$`) func Connect(ctx context.Context, url string) (*pgxpool.Pool, error) { cfg, err := pgxpool.ParseConfig(url) @@ -21,8 +33,116 @@ func Connect(ctx context.Context, url string) (*pgxpool.Pool, error) { } func Migrate(ctx context.Context, p *pgxpool.Pool) error { - if _, err := p.Exec(ctx, schema); err != nil { - return fmt.Errorf("migrate: %w", err) + // 1. Acquire advisory lock — serializes concurrent replicas. + if _, err := p.Exec(ctx, "SELECT pg_advisory_lock($1)", advisoryLockKey); err != nil { + return fmt.Errorf("advisory lock: %w", err) + } + defer func() { + if _, err := p.Exec(ctx, "SELECT pg_advisory_unlock($1)", advisoryLockKey); err != nil { + log.Printf("advisory unlock: %v", err) + } + }() + + // 2. Create tracking table. + if _, err := p.Exec(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations ( + version BIGINT PRIMARY KEY, + name TEXT NOT NULL, + applied_at TIMESTAMPTZ NOT NULL DEFAULT now() + )`); err != nil { + return fmt.Errorf("create schema_migrations: %w", err) + } + + // 3. Read embedded migration files, validate names. + entries, err := fs.ReadDir(migrationsFS, "migrations") + if err != nil { + return fmt.Errorf("read migrations dir: %w", err) + } + var files []string + for _, e := range entries { + name := e.Name() + if !migrationNameRe.MatchString(name) { + panic(fmt.Sprintf("invalid migration filename: %q (must match %s)", name, migrationNameRe)) + } + files = append(files, name) + } + sort.Strings(files) + + // 4. Baseline detection: if schema_migrations is empty but 'books' table exists, + // this is an existing database — mark 0001 as applied without re-running DDL. + var count int + if err := p.QueryRow(ctx, "SELECT count(*) FROM schema_migrations").Scan(&count); err != nil { + return fmt.Errorf("count migrations: %w", err) + } + if count == 0 { + var hasBooks bool + err := p.QueryRow(ctx, "SELECT to_regclass('books') IS NOT NULL").Scan(&hasBooks) + if err != nil { + return fmt.Errorf("check books table: %w", err) + } + if hasBooks && len(files) > 0 && strings.HasPrefix(files[0], "0001_") { + if _, err := p.Exec(ctx, + "INSERT INTO schema_migrations (version, name) VALUES ($1, $2)", + 1, files[0]); err != nil { + return fmt.Errorf("baseline insert: %w", err) + } + log.Printf("migration baseline: marked %s as applied (existing database)", files[0]) + files = files[1:] + } + } + + // 5. Build set of already-applied versions. + applied := map[int64]bool{} + rows, err := p.Query(ctx, "SELECT version FROM schema_migrations") + if err != nil { + return fmt.Errorf("list applied: %w", err) + } + defer rows.Close() + for rows.Next() { + var v int64 + if err := rows.Scan(&v); err != nil { + return fmt.Errorf("scan applied: %w", err) + } + applied[v] = true + } + if err := rows.Err(); err != nil { + return fmt.Errorf("rows applied: %w", err) + } + + // 6. Apply pending migrations in order, each in its own transaction. + for _, name := range files { + version := parseVersion(name) + if applied[version] { + continue + } + sql, err := fs.ReadFile(migrationsFS, "migrations/"+name) + if err != nil { + return fmt.Errorf("read %s: %w", name, err) + } + tx, err := p.Begin(ctx) + if err != nil { + return fmt.Errorf("begin %s: %w", name, err) + } + if _, err := tx.Exec(ctx, string(sql)); err != nil { + tx.Rollback(ctx) + return fmt.Errorf("exec %s: %w", name, err) + } + if _, err := tx.Exec(ctx, + "INSERT INTO schema_migrations (version, name) VALUES ($1, $2)", + version, name); err != nil { + tx.Rollback(ctx) + return fmt.Errorf("record %s: %w", name, err) + } + if err := tx.Commit(ctx); err != nil { + return fmt.Errorf("commit %s: %w", name, err) + } + log.Printf("migration applied: %s", name) } return nil } + +func parseVersion(name string) int64 { + parts := strings.SplitN(name, "_", 2) + var v int64 + fmt.Sscanf(parts[0], "%d", &v) + return v +} diff --git a/backend/internal/db/schema.sql b/backend/internal/db/migrations/0001_baseline.sql similarity index 100% rename from backend/internal/db/schema.sql rename to backend/internal/db/migrations/0001_baseline.sql