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