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)
This commit is contained in:
Jakub Zych
2026-09-30 14:05:17 +02:00
parent c544319cec
commit 6f50b6c940
4 changed files with 521 additions and 5 deletions

View File

@@ -0,0 +1,362 @@
package lighthouse
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"log/slog"
"strings"
"sync"
"testing"
"time"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/compass"
"git.golem15.com/golem15/summercms/modules/conga"
"git.golem15.com/golem15/summercms/modules/lagoon"
"gorm.io/gorm"
)
// Widget broadcasts through its Broadcastable method.
type Widget struct {
ID uint `gorm:"column:id;primaryKey"`
Name string `gorm:"column:name"`
OwnerID uint `gorm:"column:owner_id"`
}
func (Widget) TableName() string { return "acme_widgets" }
// BroadcastChannels publishes to the owner's widget channel; the mixed
// case proves the job lowercases it.
func (w *Widget) BroadcastChannels(context.Context, *gorm.DB) ([]string, error) {
if w.OwnerID == 0 {
return nil, nil
}
return []string{fmt.Sprintf("Widgets:%d", w.OwnerID)}, nil
}
// Gadget is broadcast only through a Binding installed by the test.
type Gadget struct {
ID uint `gorm:"column:id;primaryKey"`
Name string `gorm:"column:name"`
OwnerID uint `gorm:"column:owner_id"`
}
func (Gadget) TableName() string { return "acme_gadgets" }
const broadcastTables = `
CREATE TABLE acme_widgets (id SERIAL PRIMARY KEY, name TEXT NOT NULL, owner_id INT NOT NULL DEFAULT 0);
CREATE TABLE acme_gadgets (id SERIAL PRIMARY KEY, name TEXT NOT NULL UNIQUE, owner_id INT NOT NULL DEFAULT 0);
CREATE TABLE acme_sprockets (id SERIAL PRIMARY KEY, name TEXT NOT NULL, channels TEXT NOT NULL DEFAULT '');
`
// logCapture records log messages with their attributes as strings.
type logCapture struct {
mu sync.Mutex
records []capturedLog
}
type capturedLog struct {
level slog.Level
msg string
attrs map[string]string
}
func (h *logCapture) Enabled(context.Context, slog.Level) bool { return true }
func (h *logCapture) WithAttrs([]slog.Attr) slog.Handler { return h }
func (h *logCapture) WithGroup(string) slog.Handler { return h }
func (h *logCapture) Handle(_ context.Context, r slog.Record) error {
rec := capturedLog{level: r.Level, msg: r.Message, attrs: map[string]string{}}
r.Attrs(func(a slog.Attr) bool {
rec.attrs[a.Key] = a.Value.String()
return true
})
h.mu.Lock()
h.records = append(h.records, rec)
h.mu.Unlock()
return nil
}
func (h *logCapture) count(msg string) int {
h.mu.Lock()
defer h.mu.Unlock()
n := 0
for _, r := range h.records {
if r.msg == msg {
n++
}
}
return n
}
func (h *logCapture) find(msg string) (capturedLog, bool) {
h.mu.Lock()
defer h.mu.Unlock()
for _, r := range h.records {
if r.msg == msg {
return r, true
}
}
return capturedLog{}, false
}
func (h *logCapture) all() string {
h.mu.Lock()
defer h.mu.Unlock()
var b strings.Builder
for _, r := range h.records {
fmt.Fprintf(&b, "%s %s %v\n", r.level, r.msg, r.attrs)
}
return b.String()
}
// lhEnv is an app with the memory driver (unless kv names another), the
// broadcast callbacks installed through the production boot order
// (From before the database is published) and a running worker.
type lhEnv struct {
app *backpack.App
svc *Service
db *sql.DB
gdb *gorm.DB
logs *logCapture
}
func newLHEnv(t *testing.T, kv map[string]any, setup func(t *testing.T, svc *Service)) lhEnv {
t.Helper()
db, dsn := migratedDB(t)
ctx := t.Context()
cfg, err := compass.Open(compass.Options{Dir: t.TempDir(), Env: "testing", Environ: []string{}})
if err != nil {
t.Fatal(err)
}
settings := map[string]any{"database.dsn": dsn, "realtime.driver": "memory"}
for k, v := range kv {
settings[k] = v
}
for k, v := range settings {
if err := cfg.Set(k, v); err != nil {
t.Fatal(err)
}
}
app := backpack.New(cfg)
logs := &logCapture{}
if err := app.Publish(slog.New(logs)); err != nil {
t.Fatal(err)
}
svc, err := From(app)
if err != nil {
t.Fatal(err)
}
if setup != nil {
setup(t, svc)
}
gdb, err := lagoon.Use(ctx, db)
if err != nil {
t.Fatal(err)
}
if err := gdb.Exec(broadcastTables).Error; err != nil {
t.Fatal(err)
}
if err := lagoon.Publish(app, db, gdb); err != nil {
t.Fatal(err)
}
w, err := conga.StartWorker(ctx, app, nil, conga.WorkerOptions{})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := w.Stop(stopCtx); err != nil {
t.Errorf("stop worker: %v", err)
}
})
return lhEnv{app: app, svc: svc, db: db, gdb: gdb, logs: logs}
}
func (e lhEnv) memory(t *testing.T) *MemoryDriver {
t.Helper()
m, ok := e.svc.Driver().(*MemoryDriver)
if !ok {
t.Fatalf("driver is %T, want the memory driver", e.svc.Driver())
}
return m
}
// broadcastJobs counts summer.broadcast jobs through a separate connection.
func (e lhEnv) broadcastJobs(t *testing.T) int {
t.Helper()
var n int
if err := e.db.QueryRowContext(t.Context(), `SELECT count(*) FROM river_job WHERE kind = 'summer.broadcast'`).Scan(&n); err != nil {
t.Fatal(err)
}
return n
}
// waitPublications waits for exactly n publications after the first from,
// then waits a little longer to catch extra ones.
func waitPublications(t *testing.T, m *MemoryDriver, from, n int) []Publication {
t.Helper()
deadline := time.Now().Add(10 * time.Second)
for len(m.Publications()) < from+n {
if time.Now().After(deadline) {
t.Fatalf("got %d publications, want %d", len(m.Publications())-from, n)
}
time.Sleep(20 * time.Millisecond)
}
time.Sleep(300 * time.Millisecond)
pubs := m.Publications()[from:]
if len(pubs) != n {
t.Fatalf("got %d publications, want exactly %d: %+v", len(pubs), n, pubs)
}
return pubs
}
// payloadKeys returns the top-level keys of a JSON object in order.
func payloadKeys(t *testing.T, raw json.RawMessage) []string {
t.Helper()
dec := json.NewDecoder(strings.NewReader(string(raw)))
if tok, err := dec.Token(); err != nil || tok != json.Delim('{') {
t.Fatalf("payload %s is not an object", raw)
}
var keys []string
for dec.More() {
tok, err := dec.Token()
if err != nil {
t.Fatal(err)
}
keys = append(keys, tok.(string))
var skip json.RawMessage
if err := dec.Decode(&skip); err != nil {
t.Fatal(err)
}
}
return keys
}
// TestBroadcastTx covers RT-03 and T-11-05: model broadcasts are River jobs
// enqueued in the write's transaction, published once after the commit and
// never after a rollback, including single-statement writes for which GORM
// opens its own transaction.
func TestBroadcastTx(t *testing.T) {
var (
mu sync.Mutex
inTxLog []bool
)
env := newLHEnv(t, nil, func(t *testing.T, svc *Service) {
err := Bind[Gadget](svc, Binding[Gadget]{
Alias: "acme.gadget",
Channels: func(_ context.Context, tx *gorm.DB, m *Gadget) ([]string, error) {
_, inTx := tx.Statement.ConnPool.(*sql.Tx)
mu.Lock()
inTxLog = append(inTxLog, inTx)
mu.Unlock()
return []string{fmt.Sprintf("gadgets:%d", m.OwnerID)}, nil
},
})
if err != nil {
t.Fatal(err)
}
})
mem := env.memory(t)
ctx := t.Context()
t.Run("commit_publishes_once_after_commit", func(t *testing.T) {
from := len(mem.Publications())
before := env.broadcastJobs(t)
err := lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error {
if err := tx.WithContext(ctx).Create(&Widget{Name: "w1", OwnerID: 5}).Error; err != nil {
return err
}
var inside int
if err := tx.Raw(`SELECT count(*) FROM river_job WHERE kind = 'summer.broadcast'`).Scan(&inside).Error; err != nil {
return err
}
if inside != before+1 {
t.Errorf("jobs inside the write transaction = %d, want %d", inside, before+1)
}
if outside := env.broadcastJobs(t); outside != before {
t.Errorf("job visible outside the transaction before commit: %d, want %d", outside, before)
}
return nil
})
if err != nil {
t.Fatal(err)
}
p := waitPublications(t, mem, from, 1)[0]
if p.Method != "publish" || len(p.Channels) != 1 || p.Channels[0] != "widgets:5" {
t.Fatalf("publication = %s %v, want publish [widgets:5]", p.Method, p.Channels)
}
if p.Event != "created.lighthouse.widget" {
t.Fatalf("event = %q", p.Event)
}
if got := strings.Join(payloadKeys(t, p.Payload), ","); got != "model,actor,timestamp,ttl" {
t.Fatalf("payload keys = %s", got)
}
var body struct {
Model Widget `json:"model"`
Actor Actor `json:"actor"`
TTL int `json:"ttl"`
}
if err := json.Unmarshal(p.Payload, &body); err != nil {
t.Fatal(err)
}
if body.Model.Name != "w1" || body.TTL != DefaultTTL || body.Actor.UserID != nil || body.Actor.Name == nil || *body.Actor.Name != "System" {
t.Fatalf("payload = %s", p.Payload)
}
})
t.Run("rollback_publishes_nothing", func(t *testing.T) {
from := len(mem.Publications())
before := env.broadcastJobs(t)
rollback := errors.New("rollback")
err := lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error {
if err := tx.WithContext(ctx).Create(&Widget{Name: "w-rolled-back", OwnerID: 5}).Error; err != nil {
return err
}
return rollback
})
if !errors.Is(err, rollback) {
t.Fatalf("err = %v", err)
}
if after := env.broadcastJobs(t); after != before {
t.Fatalf("broadcast jobs %d -> %d after a rollback", before, after)
}
time.Sleep(300 * time.Millisecond)
if n := len(mem.Publications()) - from; n != 0 {
t.Fatalf("%d publications after a rollback", n)
}
})
t.Run("single_statement_write_enqueues_in_its_own_transaction", func(t *testing.T) {
from := len(mem.Publications())
mu.Lock()
inTxLog = nil
mu.Unlock()
if err := env.gdb.WithContext(ctx).Create(&Gadget{Name: "g1", OwnerID: 3}).Error; err != nil {
t.Fatal(err)
}
mu.Lock()
got := append([]bool(nil), inTxLog...)
mu.Unlock()
if len(got) != 1 || !got[0] {
t.Fatalf("channels ran on the write's transaction: %v, want [true] (the job must be enqueued before GORM commits)", got)
}
p := waitPublications(t, mem, from, 1)[0]
if p.Event != "created.acme.gadget" || p.Channels[0] != "gadgets:3" {
t.Fatalf("publication = %+v", p)
}
})
t.Run("failed_single_statement_write_enqueues_nothing", func(t *testing.T) {
before := env.broadcastJobs(t)
if err := env.gdb.WithContext(ctx).Create(&Gadget{Name: "g1", OwnerID: 3}).Error; err == nil {
t.Fatal("duplicate gadget was created")
}
if after := env.broadcastJobs(t); after != before {
t.Fatalf("broadcast jobs %d -> %d after a failed write", before, after)
}
})
}