Files
summercms/modules/lighthouse/broadcast.go
Jakub Zych 6f50b6c940 fix(11-07): enqueue broadcast jobs before GORM commits a single write
- lighthouse:after_create/update/delete also declare
  Before(gorm:commit_or_rollback_transaction); an After-only anchor put
  them past GORM's own commit, so a plain gdb.Create enqueued its
  broadcast job after the commit on the pool (deferred from 11-05)
- lighthouse gets the testcontainers Postgres harness and TestBroadcastTx
  (commit publishes once, rollback nothing, single-statement write
  enqueues on its own transaction, failed write enqueues nothing)
2026-09-30 14:05:17 +02:00

486 lines
14 KiB
Go

package lighthouse
import (
"bytes"
"context"
"database/sql"
"encoding/json"
"fmt"
"log/slog"
"reflect"
"strings"
"time"
"git.golem15.com/golem15/summercms/modules/wire"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
// Action is a model change that is broadcast.
type Action string
const (
ActionCreated Action = "created"
ActionUpdated Action = "updated"
ActionDeleted Action = "deleted"
)
// DefaultTTL is the ttl of the default broadcast payload, in seconds.
const DefaultTTL = 60
// Names of the GORM callbacks that broadcast model changes.
const (
CallbackSnapshot = "lighthouse:snapshot"
CallbackAfterCreate = "lighthouse:after_create"
CallbackAfterUpdate = "lighthouse:after_update"
CallbackAfterDelete = "lighthouse:after_delete"
)
const (
snapshotKey = "lighthouse:snapshot"
savepoint = "lighthouse_broadcast"
)
// Event describes the change being broadcast. Payload builders receive it.
type Event struct {
Action Action
Actor Actor
Timestamp wire.Time
TTL int
}
// Broadcastable is implemented (on the pointer receiver) by models that
// broadcast their creates, updates and deletes. An empty channel list means
// no broadcast.
type Broadcastable interface {
BroadcastChannels(ctx context.Context, tx *gorm.DB) ([]string, error)
}
// BroadcastPayloader replaces the default payload {model, actor, timestamp,
// ttl}.
type BroadcastPayloader interface {
BroadcastPayload(ctx context.Context, tx *gorm.DB, ev Event) (any, error)
}
// BroadcastAliaser replaces the default alias of the event name
// (<plugin>.<model>).
type BroadcastAliaser interface {
BroadcastAlias() string
}
// BroadcastFilter can veto the broadcast of an action.
type BroadcastFilter interface {
ShouldBroadcast(Action) bool
}
// BroadcastTTLer replaces the default ttl of 60 seconds.
type BroadcastTTLer interface {
BroadcastTTL() int
}
// Binding makes T broadcastable without methods on T, for payloads that
// need packages a model package must not import. Channels is required;
// every other field falls back to the defaults of Broadcastable.
type Binding[T any] struct {
// Alias replaces the default <plugin>.<model> of the event name.
Alias string
// Channels returns the channels of m; empty means no broadcast.
Channels func(ctx context.Context, tx *gorm.DB, m *T) ([]string, error)
// Payload replaces the default {model, actor, timestamp, ttl} payload.
Payload func(ctx context.Context, tx *gorm.DB, m *T, ev Event) (any, error)
// ShouldBroadcast can veto an action.
ShouldBroadcast func(Action) bool
// TTL replaces the default ttl of 60 seconds.
TTL int
}
// Bind registers the broadcast binding of model type T (a struct type).
// A nil Channels function or a second binding for T is an error. A bound
// type is broadcast through its binding even when it also implements
// Broadcastable.
func Bind[T any](svc *Service, b Binding[T]) error {
if svc == nil {
return fmt.Errorf("lighthouse: Bind on a nil service")
}
typ := reflect.TypeFor[T]()
if typ.Kind() != reflect.Struct {
return fmt.Errorf("lighthouse: Bind needs a struct model type, got %s", typ)
}
if b.Channels == nil {
return fmt.Errorf("lighthouse: Bind %s: Channels is nil", typ)
}
h := &handler{
alias: b.Alias,
ttl: b.TTL,
channels: func(ctx context.Context, tx *gorm.DB, m any) ([]string, error) {
return b.Channels(ctx, tx, m.(*T))
},
}
if h.alias == "" {
h.alias = defaultAlias(typ)
}
if h.ttl <= 0 {
h.ttl = DefaultTTL
}
if b.Payload != nil {
h.payload = func(ctx context.Context, tx *gorm.DB, m any, ev Event) (any, error) {
return b.Payload(ctx, tx, m.(*T), ev)
}
}
if b.ShouldBroadcast != nil {
h.should = func(_ any, a Action) bool { return b.ShouldBroadcast(a) }
}
svc.mu.Lock()
defer svc.mu.Unlock()
if svc.bindings == nil {
svc.bindings = map[reflect.Type]*handler{}
}
if _, dup := svc.bindings[typ]; dup {
return fmt.Errorf("lighthouse: %s is already bound", typ)
}
svc.bindings[typ] = h
return nil
}
// handler is the type-erased broadcast contract of one model type. The
// model argument is always a pointer to the model struct.
type handler struct {
alias string
ttl int
channels func(ctx context.Context, tx *gorm.DB, m any) ([]string, error)
payload func(ctx context.Context, tx *gorm.DB, m any, ev Event) (any, error)
should func(m any, a Action) bool
}
var broadcastableType = reflect.TypeFor[Broadcastable]()
// handlerFor returns the binding of typ, else a handler over the model's
// Broadcastable methods, else nil.
func (s *Service) handlerFor(typ reflect.Type) *handler {
s.mu.RLock()
h := s.bindings[typ]
s.mu.RUnlock()
if h != nil {
return h
}
ptr := reflect.PointerTo(typ)
if !ptr.Implements(broadcastableType) {
return nil
}
h = &handler{
alias: defaultAlias(typ),
ttl: DefaultTTL,
channels: func(ctx context.Context, tx *gorm.DB, m any) ([]string, error) {
return m.(Broadcastable).BroadcastChannels(ctx, tx)
},
should: func(m any, a Action) bool {
if f, ok := m.(BroadcastFilter); ok {
return f.ShouldBroadcast(a)
}
return true
},
}
zero := reflect.New(typ).Interface()
if a, ok := zero.(BroadcastAliaser); ok {
if alias := a.BroadcastAlias(); alias != "" {
h.alias = alias
}
}
if t, ok := zero.(BroadcastTTLer); ok {
if ttl := t.BroadcastTTL(); ttl > 0 {
h.ttl = ttl
}
}
if _, ok := zero.(BroadcastPayloader); ok {
h.payload = func(ctx context.Context, tx *gorm.DB, m any, ev Event) (any, error) {
return m.(BroadcastPayloader).BroadcastPayload(ctx, tx, ev)
}
}
return h
}
// defaultAlias is <plugin>.<model>: the last Go package path segment (the
// one before it when the last is "models") and the lowercased type name.
func defaultAlias(typ reflect.Type) string {
parts := strings.Split(typ.PkgPath(), "/")
plugin := "unknown"
if n := len(parts); n > 0 && parts[n-1] != "" {
plugin = parts[n-1]
if plugin == "models" && n > 1 {
plugin = parts[n-2]
}
}
return strings.ToLower(plugin + "." + typ.Name())
}
func (h *handler) eventName(a Action) string {
return strings.ToLower(string(a) + "." + h.alias)
}
func (h *handler) allows(m any, a Action) bool {
return h.should == nil || h.should(m, a)
}
// defaultPayload is the WinterCMS BroadcastableModel payload.
type defaultPayload struct {
Model any `json:"model"`
Actor Actor `json:"actor"`
Timestamp wire.Time `json:"timestamp"`
TTL int `json:"ttl"`
}
// commitCallback is GORM's commit of a transaction it opened itself.
const commitCallback = "gorm:commit_or_rollback_transaction"
// installCallbacks registers the broadcast callbacks on gdb, replacing
// earlier ones so a handle shared by several apps broadcasts through the
// most recent service. The after-write callbacks run after the model's own
// after hook and before GORM commits a single-statement write, so the job
// is enqueued on the write's transaction: GORM appends a callback that
// names only an After anchor to the end of the chain, past the commit.
func (s *Service) installCallbacks(gdb *gorm.DB) error {
cb := gdb.Callback()
if cb.Create().Get(CallbackAfterCreate) == nil {
if err := cb.Create().After("gorm:after_create").Before(commitCallback).Register(CallbackAfterCreate, s.afterCreate); err != nil {
return err
}
} else if err := cb.Create().Replace(CallbackAfterCreate, s.afterCreate); err != nil {
return err
}
if cb.Update().Get(CallbackAfterUpdate) == nil {
if err := cb.Update().After("gorm:after_update").Before(commitCallback).Register(CallbackAfterUpdate, s.afterUpdate); err != nil {
return err
}
} else if err := cb.Update().Replace(CallbackAfterUpdate, s.afterUpdate); err != nil {
return err
}
if cb.Delete().Get(CallbackSnapshot) == nil {
if err := cb.Delete().Before("gorm:before_delete").Register(CallbackSnapshot, s.snapshot); err != nil {
return err
}
} else if err := cb.Delete().Replace(CallbackSnapshot, s.snapshot); err != nil {
return err
}
if cb.Delete().Get(CallbackAfterDelete) == nil {
if err := cb.Delete().After("gorm:after_delete").Before(commitCallback).Register(CallbackAfterDelete, s.afterDelete); err != nil {
return err
}
} else if err := cb.Delete().Replace(CallbackAfterDelete, s.afterDelete); err != nil {
return err
}
return nil
}
// broadcasting reports whether the driver publishes anything. The null
// driver, or a driver whose Enabled reports false (Centrifugo without an
// API key), gets no broadcast jobs at all.
func (s *Service) broadcasting() bool {
switch d := s.driver.(type) {
case nil:
return false
case nullDriver:
return false
case interface{ Enabled() bool }:
return d.Enabled()
default:
return true
}
}
// target is one written model: its handler and a pointer to it.
type target struct {
h *handler
typ reflect.Type
model any
value reflect.Value
}
// targets returns the broadcastable, non-suppressed models of the
// statement that have a primary key. A batch update through an empty model
// has a zero key and is skipped: bulk paths suppress and Emit instead.
func (s *Service) targets(db *gorm.DB) []target {
if db.Error != nil || db.Statement == nil || db.Statement.Schema == nil || !s.broadcasting() {
return nil
}
sch := db.Statement.Schema
h := s.handlerFor(sch.ModelType)
if h == nil || suppressed(db.Statement.Context, sch.ModelType) {
return nil
}
pk := sch.PrioritizedPrimaryField
if pk == nil {
return nil
}
ctx := db.Statement.Context
var out []target
add := func(v reflect.Value) {
for v.Kind() == reflect.Pointer {
if v.IsNil() {
return
}
v = v.Elem()
}
if v.Kind() != reflect.Struct || v.Type() != sch.ModelType {
return
}
if _, zero := pk.ValueOf(ctx, v); zero {
return
}
if !v.CanAddr() {
c := reflect.New(v.Type()).Elem()
c.Set(v)
v = c
}
out = append(out, target{h: h, typ: sch.ModelType, model: v.Addr().Interface(), value: v})
}
rv := db.Statement.ReflectValue
switch rv.Kind() {
case reflect.Slice, reflect.Array:
for i := 0; i < rv.Len(); i++ {
add(rv.Index(i))
}
default:
add(rv)
}
return out
}
func (s *Service) afterCreate(db *gorm.DB) { s.afterWrite(db, ActionCreated) }
func (s *Service) afterUpdate(db *gorm.DB) { s.afterWrite(db, ActionUpdated) }
func (s *Service) afterWrite(db *gorm.DB, action Action) {
for _, t := range s.targets(db) {
if !t.h.allows(t.model, action) {
continue
}
s.inSavepoint(db, func(tx *gorm.DB) error {
args, ok, err := s.prepare(tx, t, action)
if err != nil || !ok {
return err
}
return s.enqueue(tx, args)
})
}
}
// snapshot runs before a delete: it reloads each row, computes its channels
// and deleted payload while the row still exists, and keeps them for
// afterDelete.
func (s *Service) snapshot(db *gorm.DB) {
var pending []BroadcastArgs
for _, t := range s.targets(db) {
if !t.h.allows(t.model, ActionDeleted) {
continue
}
s.inSavepoint(db, func(tx *gorm.DB) error {
full := reflect.New(t.typ)
pk, _ := db.Statement.Schema.PrioritizedPrimaryField.ValueOf(tx.Statement.Context, t.value)
if err := tx.Unscoped().Where(fmt.Sprintf("%s = ?", quoteColumn(db, db.Statement.Schema.PrioritizedPrimaryField)), pk).Take(full.Interface()).Error; err == nil {
t.model = full.Interface()
}
args, ok, err := s.prepare(tx, t, ActionDeleted)
if err == nil && ok {
pending = append(pending, args)
}
return err
})
}
if len(pending) > 0 {
db.InstanceSet(snapshotKey, pending)
}
}
// afterDelete enqueues the snapshots of a delete that succeeded.
func (s *Service) afterDelete(db *gorm.DB) {
if db.Error != nil {
return
}
v, ok := db.InstanceGet(snapshotKey)
if !ok {
return
}
pending, _ := v.([]BroadcastArgs)
for _, args := range pending {
s.inSavepoint(db, func(tx *gorm.DB) error { return s.enqueue(tx, args) })
}
}
func quoteColumn(db *gorm.DB, f *schema.Field) string {
return db.Statement.Quote(f.DBName)
}
// prepare computes the channels, event name and payload of one change.
// ok is false when the model has no channels.
func (s *Service) prepare(tx *gorm.DB, t target, action Action) (BroadcastArgs, bool, error) {
ctx := tx.Statement.Context
channels, err := t.h.channels(ctx, tx, t.model)
if err != nil {
return BroadcastArgs{}, false, fmt.Errorf("channels: %w", err)
}
if len(channels) == 0 {
return BroadcastArgs{}, false, nil
}
ev := Event{Action: action, Actor: s.Actor(ctx), Timestamp: wire.Time{Time: time.Now()}, TTL: t.h.ttl}
var payload any
if t.h.payload != nil {
payload, err = t.h.payload(ctx, tx, t.model, ev)
if err != nil {
return BroadcastArgs{}, false, fmt.Errorf("payload: %w", err)
}
} else {
payload = defaultPayload{Model: t.model, Actor: ev.Actor, Timestamp: ev.Timestamp, TTL: ev.TTL}
}
raw, err := marshalPayload(payload)
if err != nil {
return BroadcastArgs{}, false, err
}
return BroadcastArgs{Channels: channels, Event: t.h.eventName(action), Payload: raw}, true, nil
}
// inSavepoint runs fn on a fresh session of db's connection. Inside a
// transaction fn runs in a savepoint, so a failed query or enqueue is rolled
// back to it and never aborts the write. Failures are logged, not returned.
func (s *Service) inSavepoint(db *gorm.DB, fn func(tx *gorm.DB) error) {
ctx := db.Statement.Context
if ctx == nil {
ctx = context.Background()
}
tx := db.Session(&gorm.Session{NewDB: true, Context: ctx})
_, inTx := tx.Statement.ConnPool.(*sql.Tx)
if inTx {
if err := tx.SavePoint(savepoint).Error; err != nil {
s.Logger().Warn("realtime: broadcast skipped", slog.String("error", err.Error()))
return
}
}
var err error
func() {
defer func() {
if r := recover(); r != nil {
err = fmt.Errorf("panic: %v", r)
}
}()
err = fn(tx)
}()
if err != nil {
s.Logger().Warn("realtime: broadcast skipped", slog.String("error", err.Error()))
if inTx {
tx.RollbackTo(savepoint)
}
return
}
if inTx {
tx.Exec("RELEASE SAVEPOINT " + savepoint)
}
}
func marshalPayload(v any) (json.RawMessage, error) {
var buf bytes.Buffer
enc := json.NewEncoder(&buf)
enc.SetEscapeHTML(false)
if err := enc.Encode(v); err != nil {
return nil, fmt.Errorf("lighthouse: encode payload: %w", err)
}
return json.RawMessage(bytes.TrimSuffix(buf.Bytes(), []byte("\n"))), nil
}