package lagoon import ( "bytes" "context" "database/sql" "errors" "log/slog" "reflect" "strings" "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" } type acOther struct { ID uint `gorm:"column:id;primaryKey"` Label string `gorm:"column:label"` } func (acOther) TableName() string { return "lagoon_ac_others" } // 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_is_refused", func(t *testing.T) { var logs bytes.Buffer previous := slog.Default() slog.SetDefault(slog.New(slog.NewTextHandler(&logs, &slog.HandlerOptions{Level: slog.LevelWarn}))) defer slog.SetDefault(previous) ran := false err := gdb.WithContext(ctx).Transaction(func(tx *gorm.DB) error { AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { ran = true }) return nil }) if err != nil { t.Fatal(err) } if ran { t.Fatal("callback ran inside an unmanaged transaction") } if !strings.Contains(logs.String(), "after-commit callback skipped inside unmanaged transaction") { t.Fatalf("warning = %q", logs.String()) } }) // The handle a callback receives has an empty statement on the write's // connection: a query through it (with WithContext, as application code // writes it) must not continue from the written model's statement. t.Run("callback_handle_has_a_clean_statement", func(t *testing.T) { if err := gdb.Exec(`CREATE TABLE lagoon_ac_others (id SERIAL PRIMARY KEY, label TEXT NOT NULL)`).Error; err != nil { t.Fatal(err) } if err := gdb.Exec(`INSERT INTO lagoon_ac_others (label) VALUES ('other')`).Error; err != nil { t.Fatal(err) } g3, err := Use(ctx, db) if err != nil { t.Fatal(err) } var mu sync.Mutex labels := map[string]string{} if err := g3.Callback().Create().Before("gorm:create").Register("lagoon_test:clean_handle", func(stmt *gorm.DB) { item, ok := stmt.Statement.Dest.(*acItem) if !ok { return } name := item.Name AfterCommit(stmt.Statement.Context, stmt, func(ctx context.Context, d *gorm.DB) { var o acOther label := "" if err := d.WithContext(ctx).Take(&o).Error; err != nil { label = "error: " + err.Error() } else { label = o.Label } mu.Lock() labels[name] = label mu.Unlock() }) }); err != nil { t.Fatal(err) } if err := g3.WithContext(ctx).Create(&acItem{Name: "clean-implicit"}).Error; err != nil { t.Fatal(err) } if err := g3.WithContext(ctx).Transaction(func(tx *gorm.DB) error { return tx.Create(&acItem{Name: "clean-plain-tx"}).Error }); err != nil { t.Fatal(err) } if err := Transaction(ctx, g3, func(ctx context.Context, tx *gorm.DB) error { return tx.WithContext(ctx).Create(&acItem{Name: "clean-lagoon-tx"}).Error }); err != nil { t.Fatal(err) } mu.Lock() defer mu.Unlock() for _, name := range []string{"clean-implicit", "clean-lagoon-tx"} { if got := labels[name]; got != "other" { t.Errorf("%s: query through the callback handle read %q, want the lagoon_ac_others row \"other\"", name, got) } } if _, ok := labels["clean-plain-tx"]; ok { t.Fatal("unmanaged transaction callback ran") } }) 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") } }) } // TestTransactionEdges covers the argument edges of Transaction and // AfterCommit. func TestTransactionEdges(t *testing.T) { if err := Transaction(context.Background(), nil, func(context.Context, *gorm.DB) error { return nil }); err == nil { t.Fatal("Transaction accepted a nil handle") } AfterCommit(context.Background(), nil, nil) // a nil fn is a no-op ran := false AfterCommit(nil, nil, func(ctx context.Context, db *gorm.DB) { if ctx == nil || db != nil { t.Error("nil ctx and handle were not normalized") } ran = true }) if !ran { t.Fatal("AfterCommit without a handle did not run immediately") } if cleanHandle(nil, context.Background()) != nil { t.Fatal("cleanHandle(nil) is not nil") } db, _ := dedicatedDB(t, "lagoon_after_commit_edges") gdb, err := Use(t.Context(), db) if err != nil { t.Fatal(err) } var got context.Context if err := Transaction(nil, gdb, func(ctx context.Context, tx *gorm.DB) error { AfterCommit(ctx, tx, func(ctx context.Context, _ *gorm.DB) { got = ctx }) return nil }); err != nil { t.Fatal(err) } if got == nil { t.Fatal("callback of a Transaction with a nil ctx did not run with a context") } var nestedRan bool var nestedEntered bool err = Transaction(t.Context(), gdb, func(ctx context.Context, _ *gorm.DB) error { return Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error { nestedEntered = true AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { nestedRan = true }) return nil }) }) if err == nil || !strings.Contains(err.Error(), "parent transaction handle") { t.Fatalf("nested Transaction with root handle err = %v", err) } if nestedEntered || nestedRan { t.Fatal("nested root handle entered work or ran callbacks") } if err := registerAfterCommit(gdb); err != nil { t.Fatalf("registering the after-commit callback twice: %v", err) } }