From 9526b6b6380072c0258ab0237bbd3cce7430b4df Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Wed, 30 Sep 2026 22:24:04 +0200 Subject: [PATCH] fix(11-08): bind nested callbacks to parent transaction --- modules/lagoon/transaction.go | 18 ++++++++++------ modules/lagoon/transaction_test.go | 33 ++++++++++++++++++++++++++++++ 2 files changed, 45 insertions(+), 6 deletions(-) diff --git a/modules/lagoon/transaction.go b/modules/lagoon/transaction.go index 98f9ca6..0e892b1 100644 --- a/modules/lagoon/transaction.go +++ b/modules/lagoon/transaction.go @@ -19,8 +19,13 @@ type afterCommitKey struct{} // afterCommitBuffer holds the callbacks registered inside one transaction. type afterCommitBuffer struct { - mu sync.Mutex - fns []func(context.Context, *gorm.DB) + mu sync.Mutex + connPool gorm.ConnPool + fns []func(context.Context, *gorm.DB) +} + +func (b *afterCommitBuffer) owns(db *gorm.DB) bool { + return b != nil && b.connPool != nil && db != nil && db.Statement != nil && db.Statement.ConnPool == b.connPool } func (b *afterCommitBuffer) add(fns ...func(context.Context, *gorm.DB)) { @@ -55,12 +60,13 @@ func Transaction(ctx context.Context, gdb *gorm.DB, fn func(ctx context.Context, ctx = context.Background() } if parent, ok := ctx.Value(afterCommitKey{}).(*afterCommitBuffer); ok && parent != nil { - if !transactionalHandle(gdb) { + if !parent.owns(gdb) { return fmt.Errorf("lagoon: nested transaction requires the parent transaction handle") } child := &afterCommitBuffer{} childCtx := context.WithValue(ctx, afterCommitKey{}, child) err := gdb.WithContext(childCtx).Transaction(func(tx *gorm.DB) error { + child.connPool = tx.Statement.ConnPool return fn(childCtx, tx) }) if err != nil { @@ -72,6 +78,7 @@ func Transaction(ctx context.Context, gdb *gorm.DB, fn func(ctx context.Context, buf := &afterCommitBuffer{} txCtx := context.WithValue(ctx, afterCommitKey{}, buf) if err := gdb.WithContext(txCtx).Transaction(func(tx *gorm.DB) error { + buf.connPool = tx.Statement.ConnPool return fn(txCtx, tx) }); err != nil { return err @@ -90,9 +97,8 @@ func Transaction(ctx context.Context, gdb *gorm.DB, fn func(ctx context.Context, // Anywhere else, fn runs immediately on db's connection. // // The handle fn receives always has an empty statement on the connection -// the work belongs to (the pool after a commit, the open transaction -// inside a plain gorm Transaction), whatever handle AfterCommit was called -// with. Queries through it never continue from the written model's +// the work belongs to, whatever handle AfterCommit was called with. Queries +// through it never continue from the written model's // statement, even when they start with WithContext. func AfterCommit(ctx context.Context, db *gorm.DB, fn func(ctx context.Context, db *gorm.DB)) { if fn == nil { diff --git a/modules/lagoon/transaction_test.go b/modules/lagoon/transaction_test.go index 76b67a5..93cd4dc 100644 --- a/modules/lagoon/transaction_test.go +++ b/modules/lagoon/transaction_test.go @@ -146,6 +146,39 @@ func TestTransactionAfterCommit(t *testing.T) { } }) + 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 {