File
Blob: kv/sqlite3/migrations.go
| 1 | package sqlite3 |
| 2 | |
| 3 | import ( |
| 4 | "embed" |
| 5 | "fmt" |
| 6 | "io/fs" |
| 7 | "path" |
| 8 | "sort" |
| 9 | "strconv" |
| 10 | "strings" |
| 11 | ) |
| 12 | |
| 13 | //go:embed migrations/*.sql |
| 14 | var migrationFiles embed.FS |
| 15 | |
| 16 | type migration struct { |
| 17 | version int |
| 18 | name string |
| 19 | sql string |
| 20 | } |
| 21 | |
| 22 | func loadMigrations() ([]migration, error) { |
| 23 | entries, err := fs.ReadDir(migrationFiles, "migrations") |
| 24 | if err != nil { |
| 25 | return nil, err |
| 26 | } |
| 27 | |
| 28 | migrations := make([]migration, 0, len(entries)) |
| 29 | for _, entry := range entries { |
| 30 | if entry.IsDir() { |
| 31 | continue |
| 32 | } |
| 33 | |
| 34 | name := entry.Name() |
| 35 | version, err := parseMigrationVersion(name) |
| 36 | if err != nil { |
| 37 | return nil, err |
| 38 | } |
| 39 | |
| 40 | body, err := fs.ReadFile(migrationFiles, path.Join("migrations", name)) |
| 41 | if err != nil { |
| 42 | return nil, err |
| 43 | } |
| 44 | |
| 45 | migrations = append(migrations, migration{ |
| 46 | version: version, |
| 47 | name: name, |
| 48 | sql: string(body), |
| 49 | }) |
| 50 | } |
| 51 | |
| 52 | sort.Slice(migrations, func(i, j int) bool { |
| 53 | return migrations[i].version < migrations[j].version |
| 54 | }) |
| 55 | |
| 56 | if err := validateMigrationSequence(migrations); err != nil { |
| 57 | return nil, err |
| 58 | } |
| 59 | |
| 60 | return migrations, nil |
| 61 | } |
| 62 | |
| 63 | func validateMigrationSequence(migrations []migration) error { |
| 64 | for i, migration := range migrations { |
| 65 | expected := i + 1 |
| 66 | if migration.version != expected { |
| 67 | return fmt.Errorf("expected migration version %04d, got %04d (%s)", expected, migration.version, migration.name) |
| 68 | } |
| 69 | } |
| 70 | return nil |
| 71 | } |
| 72 | |
| 73 | func parseMigrationVersion(name string) (int, error) { |
| 74 | base := path.Base(name) |
| 75 | if !strings.HasSuffix(base, ".sql") { |
| 76 | return 0, fmt.Errorf("invalid migration name %q", name) |
| 77 | } |
| 78 | |
| 79 | trimmed := strings.TrimSuffix(base, ".sql") |
| 80 | prefix, _, ok := strings.Cut(trimmed, "-") |
| 81 | if !ok || prefix == "" { |
| 82 | return 0, fmt.Errorf("invalid migration name %q", name) |
| 83 | } |
| 84 | |
| 85 | version, err := strconv.Atoi(prefix) |
| 86 | if err != nil || version <= 0 { |
| 87 | return 0, fmt.Errorf("invalid migration version in %q", name) |
| 88 | } |
| 89 | return version, nil |
| 90 | } |