internal/store/store.go
269 lines · 7835 bytes
1// Package store owns SQLite access and schema migrations.
2package store
3
4import (
5 "context"
6 "database/sql"
7 "embed"
8 "errors"
9 "fmt"
10 "io/fs"
11 "os"
12 "sort"
13 "strconv"
14 "strings"
15
16 "modernc.org/sqlite"
17)
18
19//go:embed migrations/*.sql
20var migrationFS embed.FS
21
22type Store struct {
23 DB *sql.DB
24}
25
26// Open opens (creating if needed) the database at path with WAL mode and
27// foreign keys enforced. Use ":memory:" in tests.
28//
29// _txlock=immediate is what serialises writers. Every transaction in this
30// package writes, and a deferred one takes the write lock only when it
31// reaches its first write — by which point another writer may hold it.
32// SQLite answers that with SQLITE_BUSY and does not invoke the busy
33// handler, because waiting would deadlock two transactions each holding a
34// read lock the other needs; busy_timeout cannot help. Measured with
35// eight concurrent read-then-write transactions, 44% of them failed.
36// Beginning immediate takes the write lock up front, where busy_timeout
37// does apply, so a second writer waits its turn: the same load runs with
38// no failures, and readers, which WAL keeps out of the way, are
39// unaffected (#121).
40func Open(path string) (*Store, error) {
41 dsn := path + "?_txlock=immediate&_pragma=journal_mode(WAL)&_pragma=foreign_keys(ON)&_pragma=busy_timeout(5000)"
42 if path == ":memory:" {
43 dsn = ":memory:?_txlock=immediate&_pragma=foreign_keys(ON)"
44 }
45 db, err := sql.Open("sqlite", dsn)
46 if err != nil {
47 return nil, err
48 }
49 if err := db.Ping(); err != nil {
50 db.Close()
51 return nil, err
52 }
53 // SQLite creates the file 0666&~umask, so it lands 0644 by default. The
54 // directory above it is the real boundary, but the file holds token
55 // hashes, addresses and private repo names and has no business being
56 // world-readable on its own.
57 if path != ":memory:" {
58 if err := os.Chmod(path, 0o640); err != nil && !errors.Is(err, fs.ErrNotExist) {
59 db.Close()
60 return nil, err
61 }
62 }
63 return &Store{DB: db}, nil
64}
65
66func (s *Store) Close() error { return s.DB.Close() }
67
68type migration struct {
69 version int
70 name string
71 up string
72 down string
73 // upFKOff and downFKOff are true when the up/down script's first line
74 // is the directive "-- foreign_keys: off".
75 upFKOff bool
76 downFKOff bool
77}
78
79// fkOffDirective, as the first line of a migration script, opts that
80// direction out of foreign-key enforcement for its step.
81const fkOffDirective = "-- foreign_keys: off"
82
83func loadMigrations() ([]migration, error) {
84 entries, err := fs.ReadDir(migrationFS, "migrations")
85 if err != nil {
86 return nil, err
87 }
88 byVersion := map[int]*migration{}
89 for _, e := range entries {
90 name := e.Name()
91 // <version>_<name>.<up|down>.sql
92 base, ok := strings.CutSuffix(name, ".sql")
93 if !ok {
94 return nil, fmt.Errorf("migration %q: not .sql", name)
95 }
96 var dir string
97 if b, ok := strings.CutSuffix(base, ".up"); ok {
98 base, dir = b, "up"
99 } else if b, ok := strings.CutSuffix(base, ".down"); ok {
100 base, dir = b, "down"
101 } else {
102 return nil, fmt.Errorf("migration %q: missing .up/.down", name)
103 }
104 verStr, rest, ok := strings.Cut(base, "_")
105 if !ok {
106 return nil, fmt.Errorf("migration %q: missing version prefix", name)
107 }
108 ver, err := strconv.Atoi(verStr)
109 if err != nil {
110 return nil, fmt.Errorf("migration %q: bad version: %w", name, err)
111 }
112 m := byVersion[ver]
113 if m == nil {
114 m = &migration{version: ver, name: rest}
115 byVersion[ver] = m
116 }
117 sqlBytes, err := migrationFS.ReadFile("migrations/" + name)
118 if err != nil {
119 return nil, err
120 }
121 text := string(sqlBytes)
122 firstLine, _, _ := strings.Cut(text, "\n")
123 fkOff := strings.TrimSpace(firstLine) == fkOffDirective
124 if dir == "up" {
125 m.up = text
126 m.upFKOff = fkOff
127 } else {
128 m.down = text
129 m.downFKOff = fkOff
130 }
131 }
132 var ms []migration
133 for _, m := range byVersion {
134 if m.up == "" || m.down == "" {
135 return nil, fmt.Errorf("migration %d %q: missing up or down file", m.version, m.name)
136 }
137 ms = append(ms, *m)
138 }
139 sort.Slice(ms, func(i, j int) bool { return ms[i].version < ms[j].version })
140 for i, m := range ms {
141 if m.version != i+1 {
142 return nil, fmt.Errorf("migration versions not contiguous at %d", m.version)
143 }
144 }
145 return ms, nil
146}
147
148// Version returns the current schema version (0 = empty database).
149func (s *Store) Version() (int, error) {
150 var v int
151 err := s.DB.QueryRow("PRAGMA user_version").Scan(&v)
152 return v, err
153}
154
155// MigrateUp applies all pending migrations.
156func (s *Store) MigrateUp() error { return s.migrateTo(-1) }
157
158// MigrateTo migrates up or down to the given version. 0 empties the schema.
159func (s *Store) MigrateTo(target int) error { return s.migrateTo(target) }
160
161func (s *Store) migrateTo(target int) error {
162 ms, err := loadMigrations()
163 if err != nil {
164 return err
165 }
166 if target < 0 {
167 target = len(ms)
168 }
169 if target > len(ms) {
170 return fmt.Errorf("no such schema version %d (max %d)", target, len(ms))
171 }
172 cur, err := s.Version()
173 if err != nil {
174 return err
175 }
176 step := func(sqlText string, newVersion int, fkOff bool) error {
177 if !fkOff {
178 tx, err := s.DB.Begin()
179 if err != nil {
180 return err
181 }
182 defer tx.Rollback()
183 if _, err := tx.Exec(sqlText); err != nil {
184 return err
185 }
186 if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", newVersion)); err != nil {
187 return err
188 }
189 return tx.Commit()
190 }
191
192 // A script whose first line is "-- foreign_keys: off" rebuilds a
193 // table that other tables reference (labels, milestones): with
194 // foreign keys on, the rebuild-by-rename loses the children's
195 // rows. PRAGMA foreign_keys is a no-op inside a transaction, and
196 // the pool gives no guarantee that a pragma set on one connection
197 // is seen by the connection Begin() draws next, so the whole step
198 // — pragma off, transaction, pragma on, foreign_key_check — runs
199 // on a single pinned connection.
200 ctx := context.Background()
201 conn, err := s.DB.Conn(ctx)
202 if err != nil {
203 return err
204 }
205 defer conn.Close()
206 if _, err := conn.ExecContext(ctx, "PRAGMA foreign_keys = OFF"); err != nil {
207 return err
208 }
209 tx, err := conn.BeginTx(ctx, nil)
210 if err != nil {
211 return err
212 }
213 defer tx.Rollback()
214 if _, err := tx.Exec(sqlText); err != nil {
215 return err
216 }
217 if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", newVersion)); err != nil {
218 return err
219 }
220 if err := tx.Commit(); err != nil {
221 return err
222 }
223 if _, err := conn.ExecContext(ctx, "PRAGMA foreign_keys = ON"); err != nil {
224 return err
225 }
226 rows, err := conn.QueryContext(ctx, "PRAGMA foreign_key_check")
227 if err != nil {
228 return err
229 }
230 defer rows.Close()
231 if rows.Next() {
232 var table string
233 var rowid sql.NullInt64
234 var referredTable string
235 var fkid int
236 if err := rows.Scan(&table, &rowid, &referredTable, &fkid); err != nil {
237 return err
238 }
239 return fmt.Errorf("foreign_key_check failed after migration: %s", table)
240 }
241 return rows.Err()
242 }
243 for cur < target {
244 m := ms[cur]
245 if err := step(m.up, m.version, m.upFKOff); err != nil {
246 return fmt.Errorf("migration %d up: %w", m.version, err)
247 }
248 cur = m.version
249 }
250 for cur > target {
251 m := ms[cur-1]
252 if err := step(m.down, m.version-1, m.downFKOff); err != nil {
253 return fmt.Errorf("migration %d down: %w", m.version, err)
254 }
255 cur = m.version - 1
256 }
257 return nil
258}
259
260// IsInternal reports whether err is the database or the I/O beneath it
261// failing, as opposed to a sentinel or a message about the caller's
262// input. Callers map it to a failure exit rather than a usage error.
263func IsInternal(err error) bool {
264 var sqlErr *sqlite.Error
265 var pathErr *fs.PathError
266 return errors.As(err, &sqlErr) || errors.As(err, &pathErr) ||
267 errors.Is(err, sql.ErrTxDone) || errors.Is(err, sql.ErrConnDone) ||
268 errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled)
269}