package database import ( "testing" "github.com/DATA-DOG/go-sqlmock" "silk-server-go/migrations" "github.com/golang-migrate/migrate/v4/source/iofs" "gorm.io/driver/postgres" "gorm.io/gorm" ) func newMockGormDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock) { t.Helper() sqlDB, mock, err := sqlmock.New() if err != nil { t.Fatalf("create sqlmock: %v", err) } gdb, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{}) if err != nil { t.Fatalf("open gorm with sqlmock: %v", err) } return gdb, mock } func TestCheckSchemaVersionRejectsMismatch(t *testing.T) { db, mock := newMockGormDB(t) defer func() { sqlDB, _ := db.DB() _ = sqlDB.Close() }() mock.ExpectQuery("SELECT version, dirty FROM schema_migrations"). WillReturnRows(sqlmock.NewRows([]string{"version", "dirty"}).AddRow(1, false)) if err := CheckSchemaVersion(db, "2"); err == nil { t.Fatal("expected schema mismatch error") } } func TestCheckSchemaVersionAcceptsMatch(t *testing.T) { db, mock := newMockGormDB(t) defer func() { sqlDB, _ := db.DB() _ = sqlDB.Close() }() mock.ExpectQuery("SELECT version, dirty FROM schema_migrations"). WillReturnRows(sqlmock.NewRows([]string{"version", "dirty"}).AddRow(1, false)) if err := CheckSchemaVersion(db, "1"); err != nil { t.Fatalf("expected schema version match, got %v", err) } } func TestCheckSchemaVersionRejectsMissingVersion(t *testing.T) { db, mock := newMockGormDB(t) defer func() { sqlDB, _ := db.DB() _ = sqlDB.Close() }() mock.ExpectQuery("SELECT version, dirty FROM schema_migrations"). WillReturnRows(sqlmock.NewRows([]string{"version", "dirty"})) if err := CheckSchemaVersion(db, "1"); err == nil { t.Fatal("expected missing schema version error") } } func TestCheckSchemaVersionRejectsDirty(t *testing.T) { db, mock := newMockGormDB(t) defer func() { sqlDB, _ := db.DB() _ = sqlDB.Close() }() mock.ExpectQuery("SELECT version, dirty FROM schema_migrations"). WillReturnRows(sqlmock.NewRows([]string{"version", "dirty"}).AddRow(1, true)) if err := CheckSchemaVersion(db, "1"); err == nil { t.Fatal("expected dirty migration error") } } func TestEmbeddedMigrationsIncludeBaseline(t *testing.T) { driver, err := iofs.New(migrations.FS, ".") if err != nil { t.Fatalf("load embedded migrations: %v", err) } version, err := driver.First() if err != nil { t.Fatalf("read first migration: %v", err) } if version != 1 { t.Fatalf("expected baseline migration version 1, got %d", version) } next, err := driver.Next(version) if err != nil || next != 2 { t.Fatalf("expected risk assessment migration version 2, got %d (err %v)", next, err) } next, err = driver.Next(next) if err != nil || next != 4 { t.Fatalf("expected notifications/outbox migration version 4, got %d (err %v)", next, err) } next, err = driver.Next(next) if err != nil || next != 5 { t.Fatalf("expected detection/disease migration version 5, got %d (err %v)", next, err) } next, err = driver.Next(next) if err != nil || next != 6 { t.Fatalf("expected biosecurity migration version 6, got %d (err %v)", next, err) } next, err = driver.Next(next) if err != nil || next != 7 { t.Fatalf("expected inspection idempotency migration version 7, got %d (err %v)", next, err) } }