Files
silk/server-go/internal/database/migrate_test.go
T

112 lines
2.8 KiB
Go

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)
}
}