109 lines
2.8 KiB
Go
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
|
|
}
|