feat(db): ordered migration system with advisory lock and baseline detection
This commit is contained in:
+125
-5
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user