internal/store/store.go
191 lines · 4940 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.
28func Open(path string) (*Store, error) {
29 dsn := path + "?_pragma=journal_mode(WAL)&_pragma=foreign_keys(ON)&_pragma=busy_timeout(5000)"
30 if path == ":memory:" {
31 dsn = ":memory:?_pragma=foreign_keys(ON)"
32 }
33 db, err := sql.Open("sqlite", dsn)
34 if err != nil {
35 return nil, err
36 }
37 if err := db.Ping(); err != nil {
38 db.Close()
39 return nil, err
40 }
41 // SQLite creates the file 0666&~umask, so it lands 0644 by default. The
42 // directory above it is the real boundary, but the file holds token
43 // hashes, addresses and private repo names and has no business being
44 // world-readable on its own.
45 if path != ":memory:" {
46 if err := os.Chmod(path, 0o640); err != nil && !errors.Is(err, fs.ErrNotExist) {
47 db.Close()
48 return nil, err
49 }
50 }
51 return &Store{DB: db}, nil
52}
53
54func (s *Store) Close() error { return s.DB.Close() }
55
56type migration struct {
57 version int
58 name string
59 up string
60 down string
61}
62
63func loadMigrations() ([]migration, error) {
64 entries, err := fs.ReadDir(migrationFS, "migrations")
65 if err != nil {
66 return nil, err
67 }
68 byVersion := map[int]*migration{}
69 for _, e := range entries {
70 name := e.Name()
71 // <version>_<name>.<up|down>.sql
72 base, ok := strings.CutSuffix(name, ".sql")
73 if !ok {
74 return nil, fmt.Errorf("migration %q: not .sql", name)
75 }
76 var dir string
77 if b, ok := strings.CutSuffix(base, ".up"); ok {
78 base, dir = b, "up"
79 } else if b, ok := strings.CutSuffix(base, ".down"); ok {
80 base, dir = b, "down"
81 } else {
82 return nil, fmt.Errorf("migration %q: missing .up/.down", name)
83 }
84 verStr, rest, ok := strings.Cut(base, "_")
85 if !ok {
86 return nil, fmt.Errorf("migration %q: missing version prefix", name)
87 }
88 ver, err := strconv.Atoi(verStr)
89 if err != nil {
90 return nil, fmt.Errorf("migration %q: bad version: %w", name, err)
91 }
92 m := byVersion[ver]
93 if m == nil {
94 m = &migration{version: ver, name: rest}
95 byVersion[ver] = m
96 }
97 sqlBytes, err := migrationFS.ReadFile("migrations/" + name)
98 if err != nil {
99 return nil, err
100 }
101 if dir == "up" {
102 m.up = string(sqlBytes)
103 } else {
104 m.down = string(sqlBytes)
105 }
106 }
107 var ms []migration
108 for _, m := range byVersion {
109 if m.up == "" || m.down == "" {
110 return nil, fmt.Errorf("migration %d %q: missing up or down file", m.version, m.name)
111 }
112 ms = append(ms, *m)
113 }
114 sort.Slice(ms, func(i, j int) bool { return ms[i].version < ms[j].version })
115 for i, m := range ms {
116 if m.version != i+1 {
117 return nil, fmt.Errorf("migration versions not contiguous at %d", m.version)
118 }
119 }
120 return ms, nil
121}
122
123// Version returns the current schema version (0 = empty database).
124func (s *Store) Version() (int, error) {
125 var v int
126 err := s.DB.QueryRow("PRAGMA user_version").Scan(&v)
127 return v, err
128}
129
130// MigrateUp applies all pending migrations.
131func (s *Store) MigrateUp() error { return s.migrateTo(-1) }
132
133// MigrateTo migrates up or down to the given version. 0 empties the schema.
134func (s *Store) MigrateTo(target int) error { return s.migrateTo(target) }
135
136func (s *Store) migrateTo(target int) error {
137 ms, err := loadMigrations()
138 if err != nil {
139 return err
140 }
141 if target < 0 {
142 target = len(ms)
143 }
144 if target > len(ms) {
145 return fmt.Errorf("no such schema version %d (max %d)", target, len(ms))
146 }
147 cur, err := s.Version()
148 if err != nil {
149 return err
150 }
151 step := func(sqlText string, newVersion int) error {
152 tx, err := s.DB.Begin()
153 if err != nil {
154 return err
155 }
156 defer tx.Rollback()
157 if _, err := tx.Exec(sqlText); err != nil {
158 return err
159 }
160 if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", newVersion)); err != nil {
161 return err
162 }
163 return tx.Commit()
164 }
165 for cur < target {
166 m := ms[cur]
167 if err := step(m.up, m.version); err != nil {
168 return fmt.Errorf("migration %d up: %w", m.version, err)
169 }
170 cur = m.version
171 }
172 for cur > target {
173 m := ms[cur-1]
174 if err := step(m.down, m.version-1); err != nil {
175 return fmt.Errorf("migration %d down: %w", m.version, err)
176 }
177 cur = m.version - 1
178 }
179 return nil
180}
181
182// IsInternal reports whether err is the database or the I/O beneath it
183// failing, as opposed to a sentinel or a message about the caller's
184// input. Callers map it to a failure exit rather than a usage error.
185func IsInternal(err error) bool {
186 var sqlErr *sqlite.Error
187 var pathErr *fs.PathError
188 return errors.As(err, &sqlErr) || errors.As(err, &pathErr) ||
189 errors.Is(err, sql.ErrTxDone) || errors.Is(err, sql.ErrConnDone) ||
190 errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled)
191}