- OnDatabase runs a callback once the database is published (now, or when lagoon.Publish runs), so GORM callbacks registered at Boot also install under serve, where Boot runs before the database is opened - Transaction runs AfterCommit callbacks in order after a successful commit; nested calls are savepoints whose callbacks drop with them - the lagoon:after_commit GORM callback flushes single-statement AfterCommit work after GORM's own commit; outside a transaction it runs immediately
183 lines
5.4 KiB
Go
183 lines
5.4 KiB
Go
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
|
|
}
|