376 lines
11 KiB
Go
376 lines
11 KiB
Go
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("nested_rejects_unrelated_transaction_handle", func(t *testing.T) {
|
|
var r recorder
|
|
err := Transaction(ctx, gdb, func(ctx context.Context, _ *gorm.DB) error {
|
|
foreign := gdb.WithContext(ctx).Begin()
|
|
if foreign.Error != nil {
|
|
return foreign.Error
|
|
}
|
|
defer foreign.Rollback()
|
|
entered := false
|
|
err := Transaction(ctx, foreign, func(ctx context.Context, tx *gorm.DB) error {
|
|
entered = true
|
|
AfterCommit(ctx, tx, func(context.Context, *gorm.DB) { r.add("foreign") })
|
|
return tx.Create(&acItem{Name: "foreign-nested"}).Error
|
|
})
|
|
if err == nil || !strings.Contains(err.Error(), "parent transaction handle") {
|
|
t.Fatalf("unrelated nested handle err = %v", err)
|
|
}
|
|
if entered {
|
|
t.Fatal("unrelated nested transaction entered fn")
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := r.names(); len(got) != 0 {
|
|
t.Fatalf("callbacks = %v, want none", got)
|
|
}
|
|
if committedCount(t, db, "foreign-nested") != 0 {
|
|
t.Fatal("unrelated nested transaction committed work")
|
|
}
|
|
})
|
|
|
|
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)
|
|
}
|
|
}
|