diff --git a/internal/store/migrations.go b/internal/store/migrations.go new file mode 100644 index 0000000..815f791 --- /dev/null +++ b/internal/store/migrations.go @@ -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 +} diff --git a/internal/store/store.go b/internal/store/store.go index 7d02ce5..f2cdea8 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -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 -}