146 lines
3.3 KiB
Go
146 lines
3.3 KiB
Go
package db
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"database/sql"
|
|
"encoding/hex"
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"slices"
|
|
"strings"
|
|
)
|
|
|
|
const migrationsTableDDL = `
|
|
CREATE TABLE IF NOT EXISTS schema_migrations (
|
|
name TEXT PRIMARY KEY,
|
|
checksum_sha256 TEXT NOT NULL,
|
|
applied_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%SZ', 'now'))
|
|
);
|
|
`
|
|
|
|
func Migrate(ctx context.Context, database *sql.DB, migrationsDir string) error {
|
|
if migrationsDir == "" {
|
|
return fmt.Errorf("migrations dir is required")
|
|
}
|
|
|
|
if _, err := database.ExecContext(ctx, migrationsTableDDL); err != nil {
|
|
return fmt.Errorf("ensure schema_migrations table: %w", err)
|
|
}
|
|
|
|
applied, err := appliedMigrations(ctx, database)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
files, err := listMigrationFiles(migrationsDir)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
for _, filename := range files {
|
|
fullPath := filepath.Join(migrationsDir, filename)
|
|
|
|
contents, err := os.ReadFile(fullPath)
|
|
if err != nil {
|
|
return fmt.Errorf("read migration %s: %w", filename, err)
|
|
}
|
|
|
|
checksum := checksum(contents)
|
|
if appliedChecksum, ok := applied[filename]; ok {
|
|
if appliedChecksum != checksum {
|
|
return fmt.Errorf("migration %s checksum mismatch: applied=%s current=%s", filename, appliedChecksum, checksum)
|
|
}
|
|
|
|
continue
|
|
}
|
|
|
|
if err := applyMigration(ctx, database, filename, checksum, string(contents)); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func appliedMigrations(ctx context.Context, database *sql.DB) (map[string]string, error) {
|
|
rows, err := database.QueryContext(ctx, `SELECT name, checksum_sha256 FROM schema_migrations`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("load applied migrations: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
applied := make(map[string]string)
|
|
for rows.Next() {
|
|
var name string
|
|
var checksum string
|
|
|
|
if err := rows.Scan(&name, &checksum); err != nil {
|
|
return nil, fmt.Errorf("scan applied migration: %w", err)
|
|
}
|
|
|
|
applied[name] = checksum
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("iterate applied migrations: %w", err)
|
|
}
|
|
|
|
return applied, nil
|
|
}
|
|
|
|
func listMigrationFiles(migrationsDir string) ([]string, error) {
|
|
entries, err := os.ReadDir(migrationsDir)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("read migrations dir: %w", err)
|
|
}
|
|
|
|
files := make([]string, 0, len(entries))
|
|
for _, entry := range entries {
|
|
if entry.IsDir() || filepath.Ext(entry.Name()) != ".sql" {
|
|
continue
|
|
}
|
|
|
|
files = append(files, entry.Name())
|
|
}
|
|
|
|
slices.Sort(files)
|
|
|
|
return files, nil
|
|
}
|
|
|
|
func applyMigration(ctx context.Context, database *sql.DB, filename, checksum, sqlText string) error {
|
|
tx, err := database.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return fmt.Errorf("begin migration %s: %w", filename, err)
|
|
}
|
|
|
|
if strings.TrimSpace(sqlText) != "" {
|
|
if _, err := tx.ExecContext(ctx, sqlText); err != nil {
|
|
tx.Rollback()
|
|
return fmt.Errorf("execute migration %s: %w", filename, err)
|
|
}
|
|
}
|
|
|
|
if _, err := tx.ExecContext(
|
|
ctx,
|
|
`INSERT INTO schema_migrations (name, checksum_sha256) VALUES (?, ?)`,
|
|
filename,
|
|
checksum,
|
|
); err != nil {
|
|
tx.Rollback()
|
|
return fmt.Errorf("record migration %s: %w", filename, err)
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return fmt.Errorf("commit migration %s: %w", filename, err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func checksum(contents []byte) string {
|
|
sum := sha256.Sum256(contents)
|
|
return hex.EncodeToString(sum[:])
|
|
}
|