fix(11-08): bind nested callbacks to parent transaction

This commit is contained in:
Jakub Zych
2026-09-30 22:24:04 +02:00
parent 4d7823984e
commit 9526b6b638
2 changed files with 45 additions and 6 deletions

View File

@@ -20,9 +20,14 @@ type afterCommitKey struct{}
// afterCommitBuffer holds the callbacks registered inside one transaction. // afterCommitBuffer holds the callbacks registered inside one transaction.
type afterCommitBuffer struct { type afterCommitBuffer struct {
mu sync.Mutex mu sync.Mutex
connPool gorm.ConnPool
fns []func(context.Context, *gorm.DB) 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)) { func (b *afterCommitBuffer) add(fns ...func(context.Context, *gorm.DB)) {
b.mu.Lock() b.mu.Lock()
defer b.mu.Unlock() defer b.mu.Unlock()
@@ -55,12 +60,13 @@ func Transaction(ctx context.Context, gdb *gorm.DB, fn func(ctx context.Context,
ctx = context.Background() ctx = context.Background()
} }
if parent, ok := ctx.Value(afterCommitKey{}).(*afterCommitBuffer); ok && parent != nil { 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") return fmt.Errorf("lagoon: nested transaction requires the parent transaction handle")
} }
child := &afterCommitBuffer{} child := &afterCommitBuffer{}
childCtx := context.WithValue(ctx, afterCommitKey{}, child) childCtx := context.WithValue(ctx, afterCommitKey{}, child)
err := gdb.WithContext(childCtx).Transaction(func(tx *gorm.DB) error { err := gdb.WithContext(childCtx).Transaction(func(tx *gorm.DB) error {
child.connPool = tx.Statement.ConnPool
return fn(childCtx, tx) return fn(childCtx, tx)
}) })
if err != nil { if err != nil {
@@ -72,6 +78,7 @@ func Transaction(ctx context.Context, gdb *gorm.DB, fn func(ctx context.Context,
buf := &afterCommitBuffer{} buf := &afterCommitBuffer{}
txCtx := context.WithValue(ctx, afterCommitKey{}, buf) txCtx := context.WithValue(ctx, afterCommitKey{}, buf)
if err := gdb.WithContext(txCtx).Transaction(func(tx *gorm.DB) error { if err := gdb.WithContext(txCtx).Transaction(func(tx *gorm.DB) error {
buf.connPool = tx.Statement.ConnPool
return fn(txCtx, tx) return fn(txCtx, tx)
}); err != nil { }); err != nil {
return err 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. // Anywhere else, fn runs immediately on db's connection.
// //
// The handle fn receives always has an empty statement on the 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 // the work belongs to, whatever handle AfterCommit was called with. Queries
// inside a plain gorm Transaction), whatever handle AfterCommit was called // through it never continue from the written model's
// with. Queries through it never continue from the written model's
// statement, even when they start with WithContext. // statement, even when they start with WithContext.
func AfterCommit(ctx context.Context, db *gorm.DB, fn func(ctx context.Context, db *gorm.DB)) { func AfterCommit(ctx context.Context, db *gorm.DB, fn func(ctx context.Context, db *gorm.DB)) {
if fn == nil { if fn == nil {

View File

@@ -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) { t.Run("panicking_callback_does_not_fail_commit", func(t *testing.T) {
var r recorder var r recorder
err := Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error { err := Transaction(ctx, gdb, func(ctx context.Context, tx *gorm.DB) error {