Files
2026-07-26 14:01:20 -05:00

97 lines
2.3 KiB
Go

package main
import (
"context"
"database/sql"
"fmt"
"log"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"time"
_ "modernc.org/sqlite"
)
func main() {
dsn := os.Getenv("TAPM_DATABASE_DSN")
if dsn == "" {
dsn = "file:/data/tapm.db?_pragma=busy_timeout(5000)&_pragma=foreign_keys(1)&_pragma=journal_mode(WAL)"
}
directory := os.Getenv("TAPM_MIGRATIONS_DIR")
if directory == "" {
directory = "/app/migrations"
}
db, err := sql.Open("sqlite", dsn)
if err != nil {
log.Fatal(err)
}
defer db.Close()
db.SetMaxOpenConns(1)
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
if err := db.PingContext(ctx); err != nil {
log.Fatalf("database: %v", err)
}
if _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER NOT NULL PRIMARY KEY,
applied_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`); err != nil {
log.Fatal(err)
}
entries, err := os.ReadDir(directory)
if err != nil {
log.Fatal(err)
}
sort.Slice(entries, func(i, j int) bool { return entries[i].Name() < entries[j].Name() })
for _, entry := range entries {
if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".sql") {
continue
}
versionText, _, ok := strings.Cut(entry.Name(), "_")
if !ok {
log.Fatalf("migration %q must begin with a numeric version and underscore", entry.Name())
}
version, err := strconv.ParseUint(versionText, 10, 64)
if err != nil {
log.Fatalf("migration %q: %v", entry.Name(), err)
}
var applied int
if err := db.QueryRowContext(ctx,
`SELECT COUNT(*) FROM schema_migrations WHERE version = ?`, version,
).Scan(&applied); err != nil {
log.Fatal(err)
}
if applied != 0 {
continue
}
body, err := os.ReadFile(filepath.Join(directory, entry.Name()))
if err != nil {
log.Fatal(err)
}
tx, err := db.BeginTx(ctx, nil)
if err != nil {
log.Fatal(err)
}
if _, err := tx.ExecContext(ctx, string(body)); err != nil {
_ = tx.Rollback()
log.Fatalf("apply %s: %v", entry.Name(), err)
}
if _, err := tx.ExecContext(ctx,
`INSERT INTO schema_migrations (version) VALUES (?)`, version,
); err != nil {
_ = tx.Rollback()
log.Fatalf("record %s: %v", entry.Name(), err)
}
if err := tx.Commit(); err != nil {
log.Fatalf("commit %s: %v", entry.Name(), err)
}
fmt.Printf("applied migration %s\n", entry.Name())
}
}