diff --git a/api/database/database.go b/api/database/database.go index 6402e954..823b08fd 100644 --- a/api/database/database.go +++ b/api/database/database.go @@ -2,6 +2,7 @@ package database import ( "context" + "errors" "fmt" "log" "net/url" @@ -204,50 +205,16 @@ func MigrateDatabase(db *gorm.DB) error { } func ClearDatabase(db *gorm.DB) error { - return db.Transaction(func(tx *gorm.DB) error { - - dbDriver := drivers.DatabaseDriverFromEnv() - - if dbDriver == drivers.MYSQL { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0;").Error; err != nil { - return err - } - } - - if err := clearTables(tx, dbDriver); err != nil { - return err - } - - if dbDriver == drivers.MYSQL { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 1;").Error; err != nil { - return err - } - } - - return nil - }) -} - -func clearTables(tx *gorm.DB, dbDriver drivers.DatabaseDriverType) error { - dryRun := tx.Session(&gorm.Session{DryRun: true}) + var errs []error for _, model := range database_models { - // get table name of model structure - table := dryRun.Find(model).Statement.Table - - switch dbDriver { - case drivers.POSTGRES: - if err := tx.Exec(fmt.Sprintf("TRUNCATE TABLE %s CASCADE", table)).Error; err != nil { - return err - } - case drivers.MYSQL: - if err := tx.Exec(fmt.Sprintf("TRUNCATE TABLE %s", table)).Error; err != nil { - return err - } - case drivers.SQLITE: - if err := tx.Exec(fmt.Sprintf("DELETE FROM %s", table)).Error; err != nil { - return err - } + if err := db.Migrator().DropTable(model); err != nil { + errs = append(errs, err) } } + + if err := errors.Join(errs...); err != nil { + return fmt.Errorf("drop tables error: %w", err) + } + return nil } diff --git a/api/graphql/models/album.go b/api/graphql/models/album.go index 4bcf6a80..f2c05e96 100644 --- a/api/graphql/models/album.go +++ b/api/graphql/models/album.go @@ -99,6 +99,7 @@ func (a *Album) Thumbnail(db *gorm.DB) (*Media, error) { ) SELECT * FROM media WHERE media.album_id IN (SELECT id FROM sub_albums) + ORDER BY media.id LIMIT 1 ` diff --git a/api/scanner/media_type/media_type_test.go b/api/scanner/media_type/media_type_test.go index 34118fe5..9085ff3b 100644 --- a/api/scanner/media_type/media_type_test.go +++ b/api/scanner/media_type/media_type_test.go @@ -1,9 +1,15 @@ package media_type import ( + "flag" "testing" ) +func init() { + // Avoid panic with providing flags in `test_utils/integration_setup.go`. + flag.CommandLine.Init("media_type", flag.ContinueOnError) +} + type boolImage bool const isImage boolImage = true diff --git a/api/test_utils/integration_setup.go b/api/test_utils/integration_setup.go index e6e2e5a6..46718f1f 100644 --- a/api/test_utils/integration_setup.go +++ b/api/test_utils/integration_setup.go @@ -72,7 +72,7 @@ func DatabaseTest(t *testing.T) *gorm.DB { t.Skip("Database integration tests disabled") } - if err := test_dbm.SetupOrReset(); err != nil { + if err := test_dbm.SetupAndReset(); err != nil { t.Fatalf("failed to setup or reset test database: %v", err) } diff --git a/api/test_utils/test_db_manager.go b/api/test_utils/test_db_manager.go index 24ae58a0..09f8faa2 100644 --- a/api/test_utils/test_db_manager.go +++ b/api/test_utils/test_db_manager.go @@ -1,8 +1,9 @@ package test_utils import ( + "fmt" + "github.com/photoview/photoview/api/database" - "github.com/pkg/errors" "gorm.io/gorm" "gorm.io/gorm/logger" ) @@ -11,12 +12,14 @@ type TestDBManager struct { DB *gorm.DB } -func (dbm *TestDBManager) SetupOrReset() error { +func (dbm *TestDBManager) SetupAndReset() error { if dbm.DB == nil { - return dbm.setup() - } else { - return dbm.reset() + if err := dbm.setup(); err != nil { + return fmt.Errorf("setup db error: %w", err) + } } + + return dbm.reset() } func (dbm *TestDBManager) Close() error { @@ -24,13 +27,9 @@ func (dbm *TestDBManager) Close() error { return nil } - if err := dbm.reset(); err != nil { - return err - } - sqlDB, err := dbm.DB.DB() if err != nil { - return errors.Wrap(err, "get db instance when closing test database") + return fmt.Errorf("get db instance when closing test database error: %w", err) } sqlDB.Close() @@ -45,26 +44,21 @@ func (dbm *TestDBManager) setup() error { } db, err := database.ConfigureDatabase(&config) if err != nil { - return errors.Wrap(err, "configure test database") - } - - if err := database.MigrateDatabase(db); err != nil { - return errors.Wrap(err, "migrate test database") + return fmt.Errorf("configure test database error: %w", err) } dbm.DB = db - if err := dbm.reset(); err != nil { - return err - } - return nil } func (dbm *TestDBManager) reset() error { - if err := database.ClearDatabase(dbm.DB); err != nil { - return errors.Wrap(err, "reset test database") + return fmt.Errorf("clean database error: %w", err) + } + + if err := database.MigrateDatabase(dbm.DB); err != nil { + return fmt.Errorf("migrate database error: %w", err) } return nil