Extract database migration logic to migrations.go

This commit is contained in:
bronku 2025-03-23 21:45:24 +01:00
parent e5c5603ebd
commit 969dbfb1aa
2 changed files with 59 additions and 52 deletions

View file

@ -0,0 +1,58 @@
package store
import (
"embed"
_ "embed"
"errors"
"strconv"
"strings"
)
//go:embed migrations/*.sql
var migrations embed.FS
func (s *Store) loadFile(file string) error {
filename := strings.Split(file, ".")
version, err := strconv.Atoi(filename[0])
if err != nil {
return err
}
if version <= s.version() {
return nil
}
query, err := migrations.ReadFile("migrations/" + file)
if err != nil {
return err
}
_, err = s.db.Exec(string(query))
return err
}
func (s *Store) loadMigrations() error {
if s.db == nil {
return errors.New("database doesn't exist")
}
migration_files, err := migrations.ReadDir("migrations")
if err != nil {
return err
}
for _, e := range migration_files {
_ = s.loadFile(e.Name())
}
return nil
}
func (s *Store) version() int {
out := -1
row, err := s.db.Query("PRAGMA user_version;")
if err != nil {
return out
}
defer row.Close()
row.Next()
_ = row.Scan(&out)
return out
}

View file

@ -2,11 +2,6 @@ package store
import (
"database/sql"
"embed"
_ "embed"
"errors"
"strconv"
"strings"
_ "github.com/mattn/go-sqlite3"
)
@ -24,7 +19,7 @@ func OpenStore(filename string) (*Store, error) {
return &out, err
}
err = out.loadSchema()
err = out.loadMigrations()
return &out, err
}
@ -33,49 +28,3 @@ func (s *Store) Close() {
s.db.Close()
}
}
func (s *Store) version() int {
out := -1
row, err := s.db.Query("PRAGMA user_version;")
if err != nil {
return out
}
defer row.Close()
row.Next()
_ = row.Scan(&out)
return out
}
//go:embed migrations/*.sql
var migrations embed.FS
func (s *Store) loadSchema() error {
if s.db == nil {
return errors.New("database doesn't exist")
}
migration_files, err := migrations.ReadDir("migrations")
if err != nil {
return err
}
for _, e := range migration_files {
filename := strings.Split(e.Name(), ".")
version, err := strconv.Atoi(filename[0])
if err != nil {
continue
}
if version <= s.version() {
continue
}
query, err := migrations.ReadFile("migrations/" + e.Name())
if err != nil {
continue
}
_, err = s.db.Exec(string(query))
if err != nil {
return err
}
}
return nil
}