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 }