Files
2026-08-22 02:59:16 +02:00

109 lines
2.8 KiB
Go

package main
import (
"context"
"database/sql"
"flag"
"fmt"
"log"
"os"
"path/filepath"
"sort"
"strings"
"github.com/example/sndit/backend/internal/config"
"github.com/example/sndit/backend/internal/db"
)
func main() {
directory := flag.String("dir", "migrations", "directory containing SQL migrations")
flag.Parse()
cfg := config.Load()
ctx := context.Background()
database, err := db.New(ctx, cfg.DBDriver, cfg.DatabaseDSN())
if err != nil {
log.Fatalf("database unavailable: %v", err)
}
defer database.Close()
if err := run(ctx, database, *directory, cfg.DBDriver); err != nil {
log.Fatal(err)
}
}
func run(ctx context.Context, database *sql.DB, directory, driver string) error {
if _, err := database.ExecContext(ctx, `
CREATE TABLE IF NOT EXISTS schema_migrations (
version TEXT PRIMARY KEY,
applied_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
)
`); err != nil {
return fmt.Errorf("create migration table: %w", err)
}
directory = migrationDirectory(directory, driver)
files, err := migrationFiles(directory)
if err != nil {
return err
}
for _, path := range files {
version := filepath.Base(path)
var applied bool
if err := database.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM schema_migrations WHERE version = $1)`, version).Scan(&applied); err != nil {
return fmt.Errorf("check migration %s: %w", version, err)
}
if applied {
continue
}
sql, err := os.ReadFile(path)
if err != nil {
return fmt.Errorf("read migration %s: %w", version, err)
}
tx, err := database.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin migration %s: %w", version, err)
}
if _, err := tx.ExecContext(ctx, string(sql)); err != nil {
_ = tx.Rollback()
return fmt.Errorf("apply migration %s: %w", version, err)
}
if _, err := tx.ExecContext(ctx, `INSERT INTO schema_migrations (version) VALUES ($1)`, version); err != nil {
_ = tx.Rollback()
return fmt.Errorf("record migration %s: %w", version, err)
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit migration %s: %w", version, err)
}
log.Printf("applied migration %s", version)
}
return nil
}
func migrationDirectory(directory, driver string) string {
if strings.EqualFold(driver, "sqlite") || strings.EqualFold(driver, "sqlite3") {
return filepath.Join(directory, "sqlite")
}
return directory
}
func migrationFiles(directory string) ([]string, error) {
entries, err := os.ReadDir(directory)
if err != nil {
return nil, fmt.Errorf("read migration directory: %w", err)
}
files := make([]string, 0, len(entries))
for _, entry := range entries {
if !entry.IsDir() && strings.HasSuffix(entry.Name(), ".sql") {
files = append(files, filepath.Join(directory, entry.Name()))
}
}
sort.Strings(files)
return files, nil
}