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}