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 // (.). 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 . 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 .: 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).Error != nil { // A statement inside fn failed although fn did not report it (a // channel or payload function that treats a failed read as "no // broadcast", or the delete snapshot's reload): the transaction is // aborted and only a rollback to the savepoint keeps the write alive. tx.RollbackTo(savepoint) 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 }