package lagoon import ( "context" "fmt" "log/slog" "sync" "gorm.io/gorm" ) // AfterCommitCallback is the name of the GORM callback that runs the // AfterCommit work of a single-statement write once GORM commits it. const AfterCommitCallback = "lagoon:after_commit" const statementBufferKey = "lagoon:after_commit" type afterCommitKey struct{} // afterCommitBuffer holds the callbacks registered inside one transaction. type afterCommitBuffer struct { mu sync.Mutex fns []func(context.Context, *gorm.DB) } func (b *afterCommitBuffer) add(fns ...func(context.Context, *gorm.DB)) { b.mu.Lock() defer b.mu.Unlock() b.fns = append(b.fns, fns...) } func (b *afterCommitBuffer) take() []func(context.Context, *gorm.DB) { b.mu.Lock() defer b.mu.Unlock() fns := b.fns b.fns = nil return fns } // Transaction runs fn in a transaction whose AfterCommit callbacks run, in // registration order, only after the commit succeeds. Writes inside fn must // use the ctx and tx it receives. A Transaction nested in another becomes a // savepoint: its callbacks join the outer transaction's only when fn // succeeds, so work dropped with the savepoint never runs its callbacks. // A panicking callback is logged and never turns a committed write into an // error. func Transaction(ctx context.Context, gdb *gorm.DB, fn func(ctx context.Context, tx *gorm.DB) error) error { if gdb == nil { return fmt.Errorf("lagoon: gorm db is nil") } if ctx == nil { ctx = context.Background() } if parent, ok := ctx.Value(afterCommitKey{}).(*afterCommitBuffer); ok && parent != nil { child := &afterCommitBuffer{} childCtx := context.WithValue(ctx, afterCommitKey{}, child) err := gdb.WithContext(childCtx).Transaction(func(tx *gorm.DB) error { return fn(childCtx, tx) }) if err != nil { return err } parent.add(child.take()...) return nil } buf := &afterCommitBuffer{} txCtx := context.WithValue(ctx, afterCommitKey{}, buf) if err := gdb.WithContext(txCtx).Transaction(func(tx *gorm.DB) error { return fn(txCtx, tx) }); err != nil { return err } runAfterCommit(ctx, gdb.Session(&gorm.Session{NewDB: true, Context: ctx}), buf.take()) return nil } // AfterCommit registers fn to run after the surrounding transaction commits. // Inside Transaction it is buffered until that transaction commits. Inside a // single-statement write for which GORM opened its own transaction (for // example from a GORM create callback) it runs after that commit through the // AfterCommitCallback callback, and not at all when the write fails. // Anywhere else, including inside a plain gorm Transaction, fn runs // immediately with db. func AfterCommit(ctx context.Context, db *gorm.DB, fn func(ctx context.Context, db *gorm.DB)) { if fn == nil { return } if ctx == nil && db != nil && db.Statement != nil { ctx = db.Statement.Context } if ctx == nil { ctx = context.Background() } if buf := bufferFrom(ctx, db); buf != nil { buf.add(fn) return } if db != nil && db.Statement != nil { if _, started := db.InstanceGet("gorm:started_transaction"); started { buf := &afterCommitBuffer{} if existing, ok := db.InstanceGet(statementBufferKey); ok { if b, ok := existing.(*afterCommitBuffer); ok && b != nil { buf = b } } buf.add(fn) db.InstanceSet(statementBufferKey, buf) return } } runAfterCommit(ctx, db, []func(context.Context, *gorm.DB){fn}) } func bufferFrom(ctx context.Context, db *gorm.DB) *afterCommitBuffer { if buf, ok := ctx.Value(afterCommitKey{}).(*afterCommitBuffer); ok && buf != nil { return buf } if db != nil && db.Statement != nil && db.Statement.Context != nil { if buf, ok := db.Statement.Context.Value(afterCommitKey{}).(*afterCommitBuffer); ok && buf != nil { return buf } } return nil } // flushStatementAfterCommit is the AfterCommitCallback GORM callback: it // runs the statement's buffered callbacks once GORM's own transaction has // committed, on a fresh session over the pool. func flushStatementAfterCommit(db *gorm.DB) { raw, ok := db.InstanceGet(statementBufferKey) if !ok { return } buf, ok := raw.(*afterCommitBuffer) if !ok || buf == nil { return } fns := buf.take() if db.Error != nil || len(fns) == 0 { return } ctx := db.Statement.Context if ctx == nil { ctx = context.Background() } runAfterCommit(ctx, db.Session(&gorm.Session{NewDB: true, Context: ctx}), fns) } func runAfterCommit(ctx context.Context, db *gorm.DB, fns []func(context.Context, *gorm.DB)) { for _, fn := range fns { func() { defer func() { if r := recover(); r != nil { slog.Default().Warn("lagoon: after-commit callback panicked", "panic", fmt.Sprint(r)) } }() fn(ctx, db) }() } } // registerAfterCommit installs AfterCommitCallback on the create, update and // delete processors once per GORM handle. func registerAfterCommit(gdb *gorm.DB) error { cb := gdb.Callback() if cb.Create().Get(AfterCommitCallback) == nil { if err := cb.Create().After("gorm:commit_or_rollback_transaction").Register(AfterCommitCallback, flushStatementAfterCommit); err != nil { return err } } if cb.Update().Get(AfterCommitCallback) == nil { if err := cb.Update().After("gorm:commit_or_rollback_transaction").Register(AfterCommitCallback, flushStatementAfterCommit); err != nil { return err } } if cb.Delete().Get(AfterCommitCallback) == nil { if err := cb.Delete().After("gorm:commit_or_rollback_transaction").Register(AfterCommitCallback, flushStatementAfterCommit); err != nil { return err } } return nil }