Files

76 lines
2.0 KiB
Go

package postgres
import (
"context"
"embed"
"fmt"
"io/fs"
"sort"
"strconv"
"strings"
)
//go:embed migrations/*.up.sql
var migrationFiles embed.FS
const migrationLockID int64 = 674301
func Migrate(ctx context.Context, database Database) (err error) {
tx, err := database.Begin(ctx)
if err != nil {
return fmt.Errorf("begin migrations: %w", err)
}
defer func() {
if err != nil {
_ = tx.Rollback(ctx)
}
}()
if _, err = tx.Exec(ctx, "SELECT pg_advisory_xact_lock($1)", migrationLockID); err != nil {
return fmt.Errorf("lock migrations: %w", err)
}
if _, err = tx.Exec(ctx, `CREATE TABLE IF NOT EXISTS ledger_schema_migrations (
version bigint PRIMARY KEY,
applied_at timestamptz NOT NULL DEFAULT clock_timestamp()
)`); err != nil {
return fmt.Errorf("create migration ledger: %w", err)
}
files, err := fs.Glob(migrationFiles, "migrations/*.up.sql")
if err != nil {
return fmt.Errorf("list migrations: %w", err)
}
sort.Strings(files)
for _, name := range files {
versionText := strings.SplitN(strings.TrimPrefix(name, "migrations/"), "_", 2)[0]
version, parseErr := strconv.ParseInt(versionText, 10, 64)
if parseErr != nil {
return fmt.Errorf("parse migration %s: %w", name, parseErr)
}
var applied bool
if err = tx.QueryRow(ctx, "SELECT EXISTS (SELECT 1 FROM ledger_schema_migrations WHERE version = $1)", version).Scan(&applied); err != nil {
return fmt.Errorf("check migration %s: %w", name, err)
}
if applied {
continue
}
contents, readErr := migrationFiles.ReadFile(name)
if readErr != nil {
return fmt.Errorf("read migration %s: %w", name, readErr)
}
if _, err = tx.Exec(ctx, string(contents)); err != nil {
return fmt.Errorf("apply migration %s: %w", name, err)
}
if _, err = tx.Exec(ctx, "INSERT INTO ledger_schema_migrations (version) VALUES ($1) ON CONFLICT DO NOTHING", version); err != nil {
return fmt.Errorf("record migration %s: %w", name, err)
}
}
if err = tx.Commit(ctx); err != nil {
return fmt.Errorf("commit migrations: %w", err)
}
return nil
}