package db import ( "context" "embed" "fmt" "io/fs" "log" "regexp" "sort" "strings" "github.com/jackc/pgx/v5/pgxpool" ) //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) if err != nil { return nil, err } cfg.MaxConns = 10 return pgxpool.NewWithConfig(ctx, cfg) } func Migrate(ctx context.Context, p *pgxpool.Pool) error { // 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 }