krz/gitbay
A CLI-first git forge.
clone: git clone https://gitbay.org/krz/gitbay.git
main: internal/store/store.go · raw
1// Package store owns SQLite access and schema migrations.
2package store
3
4import (
5 "database/sql"
6 "embed"
7 "fmt"
8 "io/fs"
9 "sort"
10 "strconv"
11 "strings"
12
13 _ "modernc.org/sqlite"
14)
15
16//go:embed migrations/*.sql
17var migrationFS embed.FS
18
19type Store struct {
20 DB *sql.DB
21}
22
23// Open opens (creating if needed) the database at path with WAL mode and
24// foreign keys enforced. Use ":memory:" in tests.
25func Open(path string) (*Store, error) {
26 dsn := path + "?_pragma=journal_mode(WAL)&_pragma=foreign_keys(ON)&_pragma=busy_timeout(5000)"
27 if path == ":memory:" {
28 dsn = ":memory:?_pragma=foreign_keys(ON)"
29 }
30 db, err := sql.Open("sqlite", dsn)
31 if err != nil {
32 return nil, err
33 }
34 if err := db.Ping(); err != nil {
35 db.Close()
36 return nil, err
37 }
38 return &Store{DB: db}, nil
39}
40
41func (s *Store) Close() error { return s.DB.Close() }
42
43type migration struct {
44 version int
45 name string
46 up string
47 down string
48}
49
50func loadMigrations() ([]migration, error) {
51 entries, err := fs.ReadDir(migrationFS, "migrations")
52 if err != nil {
53 return nil, err
54 }
55 byVersion := map[int]*migration{}
56 for _, e := range entries {
57 name := e.Name()
58 // <version>_<name>.<up|down>.sql
59 base, ok := strings.CutSuffix(name, ".sql")
60 if !ok {
61 return nil, fmt.Errorf("migration %q: not .sql", name)
62 }
63 var dir string
64 if b, ok := strings.CutSuffix(base, ".up"); ok {
65 base, dir = b, "up"
66 } else if b, ok := strings.CutSuffix(base, ".down"); ok {
67 base, dir = b, "down"
68 } else {
69 return nil, fmt.Errorf("migration %q: missing .up/.down", name)
70 }
71 verStr, rest, ok := strings.Cut(base, "_")
72 if !ok {
73 return nil, fmt.Errorf("migration %q: missing version prefix", name)
74 }
75 ver, err := strconv.Atoi(verStr)
76 if err != nil {
77 return nil, fmt.Errorf("migration %q: bad version: %w", name, err)
78 }
79 m := byVersion[ver]
80 if m == nil {
81 m = &migration{version: ver, name: rest}
82 byVersion[ver] = m
83 }
84 sqlBytes, err := migrationFS.ReadFile("migrations/" + name)
85 if err != nil {
86 return nil, err
87 }
88 if dir == "up" {
89 m.up = string(sqlBytes)
90 } else {
91 m.down = string(sqlBytes)
92 }
93 }
94 var ms []migration
95 for _, m := range byVersion {
96 if m.up == "" || m.down == "" {
97 return nil, fmt.Errorf("migration %d %q: missing up or down file", m.version, m.name)
98 }
99 ms = append(ms, *m)
100 }
101 sort.Slice(ms, func(i, j int) bool { return ms[i].version < ms[j].version })
102 for i, m := range ms {
103 if m.version != i+1 {
104 return nil, fmt.Errorf("migration versions not contiguous at %d", m.version)
105 }
106 }
107 return ms, nil
108}
109
110// Version returns the current schema version (0 = empty database).
111func (s *Store) Version() (int, error) {
112 var v int
113 err := s.DB.QueryRow("PRAGMA user_version").Scan(&v)
114 return v, err
115}
116
117// MigrateUp applies all pending migrations.
118func (s *Store) MigrateUp() error { return s.migrateTo(-1) }
119
120// MigrateTo migrates up or down to the given version. 0 empties the schema.
121func (s *Store) MigrateTo(target int) error { return s.migrateTo(target) }
122
123func (s *Store) migrateTo(target int) error {
124 ms, err := loadMigrations()
125 if err != nil {
126 return err
127 }
128 if target < 0 {
129 target = len(ms)
130 }
131 if target > len(ms) {
132 return fmt.Errorf("no such schema version %d (max %d)", target, len(ms))
133 }
134 cur, err := s.Version()
135 if err != nil {
136 return err
137 }
138 step := func(sqlText string, newVersion int) error {
139 tx, err := s.DB.Begin()
140 if err != nil {
141 return err
142 }
143 defer tx.Rollback()
144 if _, err := tx.Exec(sqlText); err != nil {
145 return err
146 }
147 if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", newVersion)); err != nil {
148 return err
149 }
150 return tx.Commit()
151 }
152 for cur < target {
153 m := ms[cur]
154 if err := step(m.up, m.version); err != nil {
155 return fmt.Errorf("migration %d up: %w", m.version, err)
156 }
157 cur = m.version
158 }
159 for cur > target {
160 m := ms[cur-1]
161 if err := step(m.down, m.version-1); err != nil {
162 return fmt.Errorf("migration %d down: %w", m.version, err)
163 }
164 cur = m.version - 1
165 }
166 return nil
167}