package lagoon import ( "context" "database/sql" "errors" "reflect" "sync" "testing" "gorm.io/gorm" ) type acItem struct { ID uint `gorm:"column:id;primaryKey"` Name string `gorm:"column:name"` } func (acItem) TableName() string { return "lagoon_ac_items" } // recorder collects after-commit callback names and whether the row they // were registered for was visible to another connection when they ran. type recorder struct { mu sync.Mutex got []string } func (r *recorder) add(name string) { r.mu.Lock() defer r.mu.Unlock() r.got = append(r.got, name) } func (r *recorder) names() []string { r.mu.Lock() defer r.mu.Unlock() return append([]string(nil), r.got...) } func committedCount(t *testing.T, db *sql.DB, name string) int { t.Helper() var n int if err := db.QueryRowContext(context.Background(), `SELECT count(*) FROM lagoon_ac_items WHERE name = $1`, name).Scan(&n); err != nil { t.Error(err) } return n } // TestTransactionAfterCommit covers the after-commit seam: callbacks run // only after a successful commit, in order, and never for rolled-back work. func TestTransactionAfterCommit(t *testing.T) { db, _ := dedicatedDB(t, "lagoon_after_commit") gdb, err := Use(t.Context(), db) if err != nil { t.Fatal(err) } if err := gdb.Exec(`CREATE TABLE lagoon_ac_items (id SERIAL PRIMARY KEY, name TEXT NOT NULL UNIQUE)`).Error; err != nil { t.Fatal(err) } ctx := t.Context() t.Run("commit_runs_in_order", func(t *testing.T) { var r recorder err := Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { if committedCount(t, db, "ordered") != 1 { t.Error("callback ran before the row was committed") } r.add("a") }) if err := tx.Create(&acItem{Name: "ordered"}).Error; err != nil { return err } AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("b") }) if len(r.names()) != 0 { t.Error("callback ran inside the transaction") } return nil }) if err != nil { t.Fatal(err) } if got := r.names(); !reflect.DeepEqual(got, []string{"a", "b"}) { t.Fatalf("callbacks = %v, want [a b]", got) } }) t.Run("rollback_runs_none", func(t *testing.T) { var r recorder rollback := errors.New("rollback") err := Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("x") }) if err := tx.Create(&acItem{Name: "rolled-back"}).Error; err != nil { return err } return rollback }) if !errors.Is(err, rollback) { t.Fatalf("err = %v", err) } if got := r.names(); len(got) != 0 { t.Fatalf("callbacks after rollback = %v", got) } if committedCount(t, db, "rolled-back") != 0 { t.Fatal("row survived the rollback") } }) t.Run("nested_failure_drops_inner", func(t *testing.T) { var r recorder err := Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("outer") }) inner := Transaction(ctx, tx, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("inner-failed") }) if err := tx.Create(&acItem{Name: "inner-failed"}).Error; err != nil { return err } return errors.New("inner fails") }) if inner == nil { t.Error("inner transaction did not fail") } return Transaction(ctx, tx, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("inner-ok") }) return tx.Create(&acItem{Name: "inner-ok"}).Error }) }) if err != nil { t.Fatal(err) } if got := r.names(); !reflect.DeepEqual(got, []string{"outer", "inner-ok"}) { t.Fatalf("callbacks = %v, want [outer inner-ok]", got) } if committedCount(t, db, "inner-failed") != 0 || committedCount(t, db, "inner-ok") != 1 { t.Fatal("savepoint work was not rolled back or kept as expected") } }) t.Run("panicking_callback_does_not_fail_commit", func(t *testing.T) { var r recorder err := Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { panic("boom") }) AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("after-panic") }) return nil }) if err != nil { t.Fatalf("committed transaction reported %v", err) } if got := r.names(); !reflect.DeepEqual(got, []string{"after-panic"}) { t.Fatalf("callbacks = %v", got) } }) t.Run("implicit_single_statement", func(t *testing.T) { // A separate handle so the test callback stays off the shared one. g2, err := Use(ctx, db) if err != nil { t.Fatal(err) } var r recorder if err := g2.Callback().Create().Before("gorm:create").Register("lagoon_test:after_commit", func(stmt *gorm.DB) { item, ok := stmt.Statement.Dest.(*acItem) if !ok { return } name := item.Name AfterCommit(stmt.Statement.Context, stmt, func(context.Context, *gorm.DB) { if committedCount(t, db, name) != 1 { t.Errorf("callback for %s ran before commit", name) } r.add(name) }) }); err != nil { t.Fatal(err) } if err := g2.WithContext(ctx).Create(&acItem{Name: "implicit"}).Error; err != nil { t.Fatal(err) } if err := g2.WithContext(ctx).Create(&acItem{Name: "implicit"}).Error; err == nil { t.Fatal("duplicate insert succeeded") } if got := r.names(); !reflect.DeepEqual(got, []string{"implicit"}) { t.Fatalf("callbacks = %v, want only the committed insert", got) } }) t.Run("plain_gorm_transaction_runs_now", func(t *testing.T) { var got *gorm.DB err := gdb.WithContext(ctx).Transaction(func(tx *gorm.DB) error { AfterCommit(ctx, tx, func(_ context.Context, d *gorm.DB) { got = d }) if got != tx { t.Error("callback did not run immediately with the tx handle") } return nil }) if err != nil { t.Fatal(err) } }) t.Run("outside_transaction_runs_now", func(t *testing.T) { ran := false AfterCommit(ctx, gdb, func(context.Context, *gorm.DB) { ran = true }) if !ran { t.Fatal("callback outside a transaction did not run immediately") } }) }