package lagoon import ( "strings" "testing" "gorm.io/gorm" ) type hookRow struct { ID uint `gorm:"column:id;primaryKey"` Name string `gorm:"column:name"` Flag bool `gorm:"column:flag"` } func (hookRow) TableName() string { return "lagoon_hook_rows" } func (h *hookRow) BeforeCreate(tx *gorm.DB) error { h.Flag = true return nil } type cascadeParent struct { ID uint `gorm:"column:id;primaryKey"` Name string `gorm:"column:name"` DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"` } func (cascadeParent) TableName() string { return "lagoon_cascade_parents" } func (p *cascadeParent) BeforeDelete(tx *gorm.DB) error { return WithSoftDeleteCascade(tx, func(tx *gorm.DB) error { return tx.Where("parent_id = ?", p.ID).Delete(&cascadeChild{}).Error }) } type cascadeChild struct { ID uint `gorm:"column:id;primaryKey"` ParentID uint `gorm:"column:parent_id"` Name string `gorm:"column:name"` DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"` } func (cascadeChild) TableName() string { return "lagoon_cascade_children" } type failingParent struct { ID uint `gorm:"column:id;primaryKey"` Name string `gorm:"column:name"` DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"` } func (failingParent) TableName() string { return "lagoon_cascade_parents" } func (p *failingParent) BeforeDelete(tx *gorm.DB) error { return WithSoftDeleteCascade(tx, func(tx *gorm.DB) error { if err := tx.Where("parent_id = ?", p.ID).Delete(&cascadeChild{}).Error; err != nil { return err } return gorm.ErrInvalidData }) } func TestLifecycleBeforeCreate(t *testing.T) { sqlDB, _ := dedicatedDB(t, "lagoon_lifecycle_create") gdb, err := Use(t.Context(), sqlDB) if err != nil { t.Fatal(err) } if err := gdb.Exec(` CREATE TABLE lagoon_hook_rows ( id SERIAL PRIMARY KEY, name TEXT NOT NULL, flag BOOLEAN NOT NULL DEFAULT FALSE )`).Error; err != nil { t.Fatal(err) } row := hookRow{Name: "created"} if err := gdb.Create(&row).Error; err != nil { t.Fatal(err) } if !row.Flag { t.Fatal("BeforeCreate did not run") } var stored hookRow if err := gdb.Take(&stored, row.ID).Error; err != nil { t.Fatal(err) } if !stored.Flag { t.Fatal("BeforeCreate flag was not persisted") } } func TestWithSoftDeleteCascadeSameTransaction(t *testing.T) { gdb := cascadeSchema(t, "lagoon_lifecycle_cascade") parent := cascadeParent{Name: "p"} if err := gdb.Create(&parent).Error; err != nil { t.Fatal(err) } child := cascadeChild{ParentID: parent.ID, Name: "c"} if err := gdb.Create(&child).Error; err != nil { t.Fatal(err) } if err := gdb.Delete(&parent).Error; err != nil { t.Fatal(err) } var parents, children int64 if err := gdb.Unscoped().Model(&cascadeParent{}).Where("id = ? AND deleted_at IS NOT NULL", parent.ID).Count(&parents).Error; err != nil { t.Fatal(err) } if err := gdb.Unscoped().Model(&cascadeChild{}).Where("id = ? AND deleted_at IS NOT NULL", child.ID).Count(&children).Error; err != nil { t.Fatal(err) } if parents != 1 || children != 1 { t.Fatalf("want both soft-deleted, parents=%d children=%d", parents, children) } } func TestWithSoftDeleteCascadeRollback(t *testing.T) { gdb := cascadeSchema(t, "lagoon_lifecycle_rollback") parent := failingParent{Name: "p"} if err := gdb.Create(&parent).Error; err != nil { t.Fatal(err) } child := cascadeChild{ParentID: parent.ID, Name: "c"} if err := gdb.Create(&child).Error; err != nil { t.Fatal(err) } err := gdb.Delete(&parent).Error if err == nil { t.Fatal("want cascade error") } if !strings.Contains(err.Error(), gorm.ErrInvalidData.Error()) { t.Fatalf("got %v", err) } var parents, children int64 if err := gdb.Model(&cascadeParent{}).Where("id = ?", parent.ID).Count(&parents).Error; err != nil { t.Fatal(err) } if err := gdb.Model(&cascadeChild{}).Where("id = ?", child.ID).Count(&children).Error; err != nil { t.Fatal(err) } if parents != 1 || children != 1 { t.Fatalf("cascade error must roll back both, parents=%d children=%d", parents, children) } } func TestWithSoftDeleteCascadeNilTx(t *testing.T) { err := WithSoftDeleteCascade(nil, func(tx *gorm.DB) error { return nil }) if err == nil { t.Fatal("want nil tx error") } if !strings.Contains(err.Error(), "nil") { t.Fatalf("got %v", err) } } func cascadeSchema(t *testing.T, dbName string) *gorm.DB { t.Helper() sqlDB, _ := dedicatedDB(t, dbName) gdb, err := Use(t.Context(), sqlDB) if err != nil { t.Fatal(err) } stmts := []string{ `CREATE TABLE lagoon_cascade_parents ( id SERIAL PRIMARY KEY, name TEXT NOT NULL, deleted_at TIMESTAMPTZ )`, `CREATE TABLE lagoon_cascade_children ( id SERIAL PRIMARY KEY, parent_id INTEGER NOT NULL, name TEXT NOT NULL, deleted_at TIMESTAMPTZ )`, } for _, stmt := range stmts { if err := gdb.Exec(stmt).Error; err != nil { t.Fatal(err) } } return gdb }