package database import ( "context" "database/sql" "errors" "fmt" "io/fs" "log/slog" "path" "sort" "strconv" "strings" ) // bootstrapVersion is 000.sql: the migration that creates the ledger // the others are recorded in. const bootstrapVersion = 0 // errBadMigrationName is returned for a schema file whose name does not // start with a version number. It is a build-time mistake, not a // runtime condition, and it fails startup rather than being skipped — // a migration silently not applied is the failure mode this whole // mechanism exists to prevent. var errBadMigrationName = errors.New( "migration filename does not start with a version number", ) // ParseMigrationVersion extracts the leading integer from a migration // filename: "001_widgets.sql" is version 1. Exported so that a project // seeded from this template can validate its own schema directory in a // test. func ParseMigrationVersion(name string) (int, error) { base := name if i := strings.IndexAny(base, "_."); i > 0 { base = base[:i] } version, err := strconv.Atoi(base) if err != nil { return 0, fmt.Errorf("%w: %q", errBadMigrationName, name) } return version, nil } // migrationSet is one embedded directory of numbered .sql migrations // (000 bootstrap plus schema files). type migrationSet struct { fsys fs.FS dir string } // collect returns the set's migration filenames sorted // lexicographically, which is why they are zero-padded. func (m migrationSet) collect() ([]string, error) { entries, err := fs.ReadDir(m.fsys, m.dir) if err != nil { return nil, fmt.Errorf("failed to read schema directory: %w", err) } var migrations []string for _, entry := range entries { if !entry.IsDir() && strings.HasSuffix(entry.Name(), ".sql") { migrations = append(migrations, entry.Name()) } } sort.Strings(migrations) return migrations, nil } // bootstrap ensures the schema_migrations table exists by applying // 000.sql if the table is missing. func (m migrationSet) bootstrap( ctx context.Context, db *sql.DB, log *slog.Logger, ) error { var tableExists int err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='schema_migrations'", ).Scan(&tableExists) if err != nil { return fmt.Errorf("failed to check for migrations table: %w", err) } if tableExists > 0 { return nil } content, err := fs.ReadFile(m.fsys, path.Join(m.dir, "000.sql")) if err != nil { return fmt.Errorf("failed to read bootstrap migration 000.sql: %w", err) } if log != nil { log.Info("applying bootstrap migration", "version", bootstrapVersion) } _, err = db.ExecContext(ctx, string(content)) if err != nil { return fmt.Errorf("failed to apply bootstrap migration: %w", err) } return nil } // applied reports whether the numbered migration has been recorded. func (m migrationSet) applied( ctx context.Context, db *sql.DB, version int, ) (bool, error) { var count int err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM schema_migrations WHERE version = ?", version, ).Scan(&count) if err != nil { return false, fmt.Errorf("failed to check migration status: %w", err) } return count > 0, nil } // applyOne reads, executes, and records one migration file. func (m migrationSet) applyOne( ctx context.Context, db *sql.DB, migration string, version int, ) error { content, err := fs.ReadFile(m.fsys, path.Join(m.dir, migration)) if err != nil { return fmt.Errorf("failed to read migration %s: %w", migration, err) } _, execErr := db.ExecContext(ctx, string(content)) if execErr != nil { return fmt.Errorf("failed to apply migration %s: %w", migration, execErr) } _, recErr := db.ExecContext(ctx, "INSERT INTO schema_migrations (version) VALUES (?)", version, ) if recErr != nil { return fmt.Errorf("failed to record migration %s: %w", migration, recErr) } return nil } // apply runs all pending migrations of the set, in order. Idempotent: // a second run over the same database applies nothing. func (m migrationSet) apply(ctx context.Context, db *sql.DB, log *slog.Logger) error { err := m.bootstrap(ctx, db, log) if err != nil { return err } migrations, err := m.collect() if err != nil { return err } for _, migration := range migrations { version, parseErr := ParseMigrationVersion(migration) if parseErr != nil { return parseErr } done, checkErr := m.applied(ctx, db, version) if checkErr != nil { return checkErr } if done { if log != nil { log.Debug("migration already applied", "version", version) } continue } if log != nil { log.Info("applying migration", "version", version) } applyErr := m.applyOne(ctx, db, migration, version) if applyErr != nil { return applyErr } if log != nil { log.Info("migration applied successfully", "version", version) } } return nil }