fix(11-08): bind nested callbacks to parent transaction
This commit is contained in:
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
Reference in New Issue
Block a user