Extract database migration logic to migrations.go
This commit is contained in:
parent
e5c5603ebd
commit
969dbfb1aa
2 changed files with 59 additions and 52 deletions
58
internal/store/migrations.go
Normal file
58
internal/store/migrations.go
Normal 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
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue