package store import ( "context" "embed" "fmt" "io/fs" "log/slog" "sort" "strconv" "strings" ) //go:embed migrations/*.sql var migrationsFS embed.FS // migration is one numbered SQL file: version parsed from the filename prefix. type migration struct { version int name string sql string } // Migrate applies every embedded migration whose version has not yet been // recorded, in ascending order, each in its own transaction. It is idempotent: // a second run with no new files is a no-op. Migration files are named // NNNN_description.sql (e.g. 0001_init.sql); NNNN is the version. func (d *DB) Migrate(ctx context.Context) error { if _, err := d.sql.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (version INTEGER PRIMARY KEY)`, ); err != nil { return fmt.Errorf("store: create schema_migrations: %w", err) } applied, err := d.appliedVersions(ctx) if err != nil { return err } migrations, err := loadMigrations() if err != nil { return err } for _, m := range migrations { if applied[m.version] { continue } if err := d.applyMigration(ctx, m); err != nil { return err } slog.Info("store: applied migration", "version", m.version, "name", m.name) } return nil } func (d *DB) appliedVersions(ctx context.Context) (map[int]bool, error) { rows, err := d.sql.QueryContext(ctx, `SELECT version FROM schema_migrations`) if err != nil { return nil, fmt.Errorf("store: read schema_migrations: %w", err) } defer rows.Close() applied := map[int]bool{} for rows.Next() { var v int if err := rows.Scan(&v); err != nil { return nil, fmt.Errorf("store: scan schema_migrations: %w", err) } applied[v] = true } return applied, rows.Err() } func (d *DB) applyMigration(ctx context.Context, m migration) error { tx, err := d.sql.BeginTx(ctx, nil) if err != nil { return fmt.Errorf("store: begin migration %d: %w", m.version, err) } defer tx.Rollback() //nolint:errcheck // no-op after a successful commit if _, err := tx.ExecContext(ctx, m.sql); err != nil { return fmt.Errorf("store: apply migration %d (%s): %w", m.version, m.name, err) } if _, err := tx.ExecContext(ctx, `INSERT INTO schema_migrations (version) VALUES (?)`, m.version, ); err != nil { return fmt.Errorf("store: record migration %d: %w", m.version, err) } return tx.Commit() } // loadMigrations reads and parses every embedded migration file, sorted by // version ascending. func loadMigrations() ([]migration, error) { entries, err := fs.ReadDir(migrationsFS, "migrations") if err != nil { return nil, fmt.Errorf("store: read migrations dir: %w", err) } var migrations []migration for _, e := range entries { if e.IsDir() || !strings.HasSuffix(e.Name(), ".sql") { continue } version, err := parseVersion(e.Name()) if err != nil { return nil, err } body, err := fs.ReadFile(migrationsFS, "migrations/"+e.Name()) if err != nil { return nil, fmt.Errorf("store: read migration %s: %w", e.Name(), err) } migrations = append(migrations, migration{version: version, name: e.Name(), sql: string(body)}) } sort.Slice(migrations, func(i, j int) bool { return migrations[i].version < migrations[j].version }) return migrations, nil } // parseVersion extracts the leading integer from a migration filename such as // "0001_init.sql". func parseVersion(name string) (int, error) { prefix, _, ok := strings.Cut(name, "_") if !ok { return 0, fmt.Errorf("store: migration %q missing NNNN_ prefix", name) } v, err := strconv.Atoi(prefix) if err != nil { return 0, fmt.Errorf("store: migration %q has non-numeric version: %w", name, err) } return v, nil }