internal/store/store.go

61564b7c32807deb349e28f2b5c6909cb4143870
gitbay/internal/store/store.go history · blame · raw

203 lines · 5718 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}
 74
 75func loadMigrations() ([]migration, error) {
 76	entries, err := fs.ReadDir(migrationFS, "migrations")
 77	if err != nil {
 78		return nil, err
 79	}
 80	byVersion := map[int]*migration{}
 81	for _, e := range entries {
 82		name := e.Name()
 83		// <version>_<name>.<up|down>.sql
 84		base, ok := strings.CutSuffix(name, ".sql")
 85		if !ok {
 86			return nil, fmt.Errorf("migration %q: not .sql", name)
 87		}
 88		var dir string
 89		if b, ok := strings.CutSuffix(base, ".up"); ok {
 90			base, dir = b, "up"
 91		} else if b, ok := strings.CutSuffix(base, ".down"); ok {
 92			base, dir = b, "down"
 93		} else {
 94			return nil, fmt.Errorf("migration %q: missing .up/.down", name)
 95		}
 96		verStr, rest, ok := strings.Cut(base, "_")
 97		if !ok {
 98			return nil, fmt.Errorf("migration %q: missing version prefix", name)
 99		}
100		ver, err := strconv.Atoi(verStr)
101		if err != nil {
102			return nil, fmt.Errorf("migration %q: bad version: %w", name, err)
103		}
104		m := byVersion[ver]
105		if m == nil {
106			m = &migration{version: ver, name: rest}
107			byVersion[ver] = m
108		}
109		sqlBytes, err := migrationFS.ReadFile("migrations/" + name)
110		if err != nil {
111			return nil, err
112		}
113		if dir == "up" {
114			m.up = string(sqlBytes)
115		} else {
116			m.down = string(sqlBytes)
117		}
118	}
119	var ms []migration
120	for _, m := range byVersion {
121		if m.up == "" || m.down == "" {
122			return nil, fmt.Errorf("migration %d %q: missing up or down file", m.version, m.name)
123		}
124		ms = append(ms, *m)
125	}
126	sort.Slice(ms, func(i, j int) bool { return ms[i].version < ms[j].version })
127	for i, m := range ms {
128		if m.version != i+1 {
129			return nil, fmt.Errorf("migration versions not contiguous at %d", m.version)
130		}
131	}
132	return ms, nil
133}
134
135// Version returns the current schema version (0 = empty database).
136func (s *Store) Version() (int, error) {
137	var v int
138	err := s.DB.QueryRow("PRAGMA user_version").Scan(&v)
139	return v, err
140}
141
142// MigrateUp applies all pending migrations.
143func (s *Store) MigrateUp() error { return s.migrateTo(-1) }
144
145// MigrateTo migrates up or down to the given version. 0 empties the schema.
146func (s *Store) MigrateTo(target int) error { return s.migrateTo(target) }
147
148func (s *Store) migrateTo(target int) error {
149	ms, err := loadMigrations()
150	if err != nil {
151		return err
152	}
153	if target < 0 {
154		target = len(ms)
155	}
156	if target > len(ms) {
157		return fmt.Errorf("no such schema version %d (max %d)", target, len(ms))
158	}
159	cur, err := s.Version()
160	if err != nil {
161		return err
162	}
163	step := func(sqlText string, newVersion int) error {
164		tx, err := s.DB.Begin()
165		if err != nil {
166			return err
167		}
168		defer tx.Rollback()
169		if _, err := tx.Exec(sqlText); err != nil {
170			return err
171		}
172		if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", newVersion)); err != nil {
173			return err
174		}
175		return tx.Commit()
176	}
177	for cur < target {
178		m := ms[cur]
179		if err := step(m.up, m.version); err != nil {
180			return fmt.Errorf("migration %d up: %w", m.version, err)
181		}
182		cur = m.version
183	}
184	for cur > target {
185		m := ms[cur-1]
186		if err := step(m.down, m.version-1); err != nil {
187			return fmt.Errorf("migration %d down: %w", m.version, err)
188		}
189		cur = m.version - 1
190	}
191	return nil
192}
193
194// IsInternal reports whether err is the database or the I/O beneath it
195// failing, as opposed to a sentinel or a message about the caller's
196// input. Callers map it to a failure exit rather than a usage error.
197func IsInternal(err error) bool {
198	var sqlErr *sqlite.Error
199	var pathErr *fs.PathError
200	return errors.As(err, &sqlErr) || errors.As(err, &pathErr) ||
201		errors.Is(err, sql.ErrTxDone) || errors.Is(err, sql.ErrConnDone) ||
202		errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled)
203}