internal/store/store.go

8bae67052c75bf6469522c4b6501791f820de5e3
gitbay/internal/store/store.go history · blame · raw

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}