krz/gitbay

A CLI-first git forge.

clone: git clone https://gitbay.org/krz/gitbay.git

repo-descriptions: internal/store/store.go · raw

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