diff --git a/eventconsumer/consumer.go b/eventconsumer/consumer.go --- a/eventconsumer/consumer.go +++ b/eventconsumer/consumer.go @@ -4,26 +4,20 @@ "context" "encoding/json" "log/slog" - "math/rand" + "net/http" "net/url" "sync" "time" "tangled.org/core/eventconsumer/cursor" + "tangled.org/core/eventstream" "tangled.org/core/log" "github.com/avast/retry-go/v4" "github.com/gorilla/websocket" ) -type ProcessFunc func(ctx context.Context, source Source, message Message) error - -type Message struct { - Rkey string - Nsid string - Created int64 `json:"created"` - EventJson json.RawMessage `json:"event"` -} +type ProcessFunc func(ctx context.Context, source Source, event eventstream.Event) error type ConsumerConfig struct { Sources map[Source]struct{} @@ -34,8 +28,13 @@ WorkerCount int QueueSize int Logger *slog.Logger - Dev bool CursorStore cursor.Store + URLFunc func(Source, int64) (*url.URL, error) + + Dialer *websocket.Dialer + RequestHeader http.Header + MaxRetryAttempts uint + OnConnectExceeded func(Source, error) } func NewConsumerConfig() *ConsumerConfig { @@ -44,19 +43,12 @@ } } -type Source interface { - // url to start streaming events from - Url(cursor int64, dev bool) (*url.URL, error) - // cache key for cursor storage - Key() string -} - type Consumer struct { - wg sync.WaitGroup - dialer *websocket.Dialer - jobQueue chan job - logger *slog.Logger - randSource *rand.Rand + sourceWg sync.WaitGroup + workerWg sync.WaitGroup + dialer *websocket.Dialer + jobQueue chan job + logger *slog.Logger // sourcesMu guards sources. It must only be held for short, non-blocking // map operations; never across a blocking call (dial, read, close). @@ -69,6 +61,9 @@ type sourceState struct { cancel context.CancelFunc conn *websocket.Conn + + cursorMu sync.Mutex + cursorMax int64 } type job struct { @@ -98,48 +93,60 @@ if cfg.CursorStore == nil { cfg.CursorStore = &cursor.MemoryStore{} } + if cfg.URLFunc == nil { + cfg.URLFunc = DefaultURL(false) + } + dialer := cfg.Dialer + if dialer == nil { + dialer = websocket.DefaultDialer + } return &Consumer{ - cfg: cfg, - dialer: websocket.DefaultDialer, - jobQueue: make(chan job, cfg.QueueSize), // buffered job queue - logger: cfg.Logger, - randSource: rand.New(rand.NewSource(time.Now().UnixNano())), - sources: make(map[Source]*sourceState), + cfg: cfg, + dialer: dialer, + jobQueue: make(chan job, cfg.QueueSize), + logger: cfg.Logger, + sources: make(map[Source]*sourceState), } } func (c *Consumer) Start(ctx context.Context) { c.cfg.Logger.Info("starting consumer", "config", c.cfg) - // start workers for range c.cfg.WorkerCount { - c.wg.Add(1) + c.workerWg.Add(1) go c.worker(ctx) } - // start streaming for source := range c.cfg.Sources { c.AddSource(ctx, source) } } func (c *Consumer) Stop() { - // snapshot conns under lock so we don't hold sourcesMu across Close + // snapshot cancels and conns under lock so we don't hold sourcesMu across Close c.sourcesMu.Lock() + cancels := make([]context.CancelFunc, 0, len(c.sources)) conns := make([]*websocket.Conn, 0, len(c.sources)) for _, st := range c.sources { + if st.cancel != nil { + cancels = append(cancels, st.cancel) + } if st.conn != nil { conns = append(conns, st.conn) } } c.sourcesMu.Unlock() + for _, cancel := range cancels { + cancel() + } for _, conn := range conns { conn.Close() } - c.wg.Wait() + c.sourceWg.Wait() close(c.jobQueue) + c.workerWg.Wait() } func (c *Consumer) AddSource(ctx context.Context, s Source) { @@ -153,7 +160,7 @@ c.sources[s] = &sourceState{cancel: cancel} c.sourcesMu.Unlock() - c.wg.Add(1) + c.sourceWg.Add(1) go c.startConnectionLoop(srcCtx, s) } @@ -180,7 +187,7 @@ } func (c *Consumer) worker(ctx context.Context) { - defer c.wg.Done() + defer c.workerWg.Done() for { select { case <-ctx.Done(): @@ -190,28 +197,44 @@ return } - var msg Message - err := json.Unmarshal(j.message, &msg) + var ev eventstream.Event + err := json.Unmarshal(j.message, &ev) if err != nil { c.logger.Error("error deserializing message", "source", j.source.Key(), "err", err) - return + continue } - if err := c.cfg.ProcessFunc(ctx, j.source, msg); err != nil { + if err := c.cfg.ProcessFunc(ctx, j.source, ev); err != nil { c.logger.Error("error processing message", "source", j.source, "err", err) } - cursorVal := msg.Created - if cursorVal == 0 { - cursorVal = time.Now().UnixNano() - } - c.cfg.CursorStore.Set(j.source.Key(), cursorVal) + c.advanceCursor(j.source, ev.Created) } } } +func (c *Consumer) advanceCursor(s Source, newCursor int64) { + if newCursor == 0 { + return + } + c.sourcesMu.Lock() + st, ok := c.sources[s] + c.sourcesMu.Unlock() + if !ok { + return + } + + st.cursorMu.Lock() + defer st.cursorMu.Unlock() + if newCursor <= st.cursorMax { + return + } + st.cursorMax = newCursor + c.cfg.CursorStore.Set(s.Key(), newCursor) +} + func (c *Consumer) startConnectionLoop(ctx context.Context, source Source) { - defer c.wg.Done() + defer c.sourceWg.Done() // attempt connection initially err := c.runConnection(ctx, source) @@ -240,7 +263,7 @@ func (c *Consumer) runConnection(ctx context.Context, source Source) error { cursor := c.cfg.CursorStore.Get(source.Key()) - u, err := source.Url(cursor, c.cfg.Dev) + u, err := c.cfg.URLFunc(source, cursor) if err != nil { return err } @@ -248,7 +271,7 @@ c.logger.Info("connecting", "url", u.String()) retryOpts := []retry.Option{ - retry.Attempts(0), // infinite attempts + retry.Attempts(c.cfg.MaxRetryAttempts), retry.DelayType(retry.BackOffDelay), retry.Delay(c.cfg.RetryInterval), retry.MaxDelay(c.cfg.MaxRetryInterval), @@ -269,10 +292,13 @@ err = retry.Do(func() error { connCtx, cancel := context.WithTimeout(ctx, c.cfg.ConnectionTimeout) defer cancel() - conn, _, err = c.dialer.DialContext(connCtx, u.String(), nil) + conn, _, err = c.dialer.DialContext(connCtx, u.String(), c.cfg.RequestHeader) return err }, retryOpts...) if err != nil { + if c.cfg.OnConnectExceeded != nil { + c.cfg.OnConnectExceeded(source, err) + } return err } diff --git a/eventconsumer/consumer_test.go b/eventconsumer/consumer_test.go new file mode 100644 --- /dev/null +++ b/eventconsumer/consumer_test.go @@ -0,0 +1,282 @@ +package eventconsumer + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "tangled.org/core/eventconsumer/cursor" + "tangled.org/core/eventstream" + "tangled.org/core/notifier" +) + +type memSrc struct { + mu sync.Mutex + events []eventstream.Event +} + +func (s *memSrc) add(ev eventstream.Event) { + s.mu.Lock() + defer s.mu.Unlock() + s.events = append(s.events, ev) +} + +func (s *memSrc) GetEvents(cursor int64, limit int) ([]eventstream.Event, error) { + s.mu.Lock() + defer s.mu.Unlock() + out := []eventstream.Event{} + for _, ev := range s.events { + if ev.Created > cursor { + out = append(out, ev) + if len(out) == limit { + break + } + } + } + return out, nil +} + +func mkEv(i int) eventstream.Event { + return eventstream.Event{ + Rkey: fmt.Sprintf("rk-%04d", i), + Nsid: "sh.tangled.test", + EventJson: json.RawMessage(fmt.Sprintf(`{"i":%d}`, i)), + Created: int64(i + 1), + } +} + +func startEventServer(t *testing.T, src *memSrc) (Source, *notifier.Notifier) { + t.Helper() + n := notifier.New() + mux := http.NewServeMux() + mux.HandleFunc("/events", func(w http.ResponseWriter, r *http.Request) { + _ = eventstream.Stream(w, r, eventstream.StreamConfig{ + Backend: src, + Notifier: &n, + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + BatchSize: 5, + MaxBatchesPerDrain: 100, + }) + }) + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + addr := strings.TrimPrefix(srv.URL, "http://") + return Source{Kind: "test", Host: addr}, &n +} + +func TestConsumer_DrainAdvancesCursor(t *testing.T) { + src := &memSrc{} + for i := range 8 { + src.add(mkEv(i)) + } + + source, _ := startEventServer(t, src) + + store := &cursor.MemoryStore{} + seenMu := sync.Mutex{} + seen := []int64{} + + cfg := ConsumerConfig{ + ProcessFunc: func(ctx context.Context, _ Source, msg eventstream.Event) error { + seenMu.Lock() + seen = append(seen, msg.Created) + seenMu.Unlock() + return nil + }, + WorkerCount: 1, + QueueSize: 16, + ConnectionTimeout: 2 * time.Second, + CursorStore: store, + URLFunc: DefaultURL(true), + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + } + c := NewConsumer(cfg) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + c.Start(ctx) + c.AddSource(ctx, source) + + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + seenMu.Lock() + n := len(seen) + seenMu.Unlock() + if n >= 8 { + break + } + time.Sleep(20 * time.Millisecond) + } + + seenMu.Lock() + defer seenMu.Unlock() + if len(seen) != 8 { + t.Fatalf("processed %d events, want 8: %v", len(seen), seen) + } + for i, got := range seen { + if got != int64(i+1) { + t.Fatalf("event %d: got created=%d want %d", i, got, i+1) + } + } + + if final := store.Get(source.Key()); final != 8 { + t.Fatalf("cursor = %d, want 8", final) + } +} + +func TestConsumer_CursorMonotonic_OutOfOrderWorkers(t *testing.T) { + src := &memSrc{} + for i := range 4 { + src.add(mkEv(i)) + } + + source, _ := startEventServer(t, src) + + store := &cursor.MemoryStore{} + + releaseFirst := make(chan struct{}) + processed := make(chan int64, 4) + + cfg := ConsumerConfig{ + ProcessFunc: func(ctx context.Context, _ Source, msg eventstream.Event) error { + if msg.Created == 1 { + <-releaseFirst + } + processed <- msg.Created + return nil + }, + WorkerCount: 4, + QueueSize: 16, + ConnectionTimeout: 2 * time.Second, + CursorStore: store, + URLFunc: DefaultURL(true), + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + } + c := NewConsumer(cfg) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + c.Start(ctx) + c.AddSource(ctx, source) + + for range 3 { + select { + case <-processed: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for events 2-4 to be processed") + } + } + + if cur := store.Get(source.Key()); cur != 4 { + t.Fatalf("cursor before slow worker finished = %d, want 4", cur) + } + + close(releaseFirst) + select { + case <-processed: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for slow worker") + } + + if cur := store.Get(source.Key()); cur != 4 { + t.Fatalf("cursor regressed after slow worker: %d, want 4", cur) + } +} + +func TestConsumer_StopTerminatesWithoutCtxCancel(t *testing.T) { + src := &memSrc{} + source, _ := startEventServer(t, src) + + cfg := ConsumerConfig{ + ProcessFunc: func(ctx context.Context, _ Source, _ eventstream.Event) error { return nil }, + WorkerCount: 2, + QueueSize: 8, + ConnectionTimeout: 2 * time.Second, + CursorStore: &cursor.MemoryStore{}, + URLFunc: DefaultURL(true), + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + } + c := NewConsumer(cfg) + + c.Start(context.Background()) + c.AddSource(context.Background(), source) + + done := make(chan struct{}) + go func() { + c.Stop() + close(done) + }() + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("Stop did not return within 5s") + } +} + +func TestConsumer_ResumesFromStoredCursor(t *testing.T) { + src := &memSrc{} + for i := range 5 { + src.add(mkEv(i)) + } + + source, _ := startEventServer(t, src) + + store := &cursor.MemoryStore{} + store.Set(source.Key(), 3) + + seenMu := sync.Mutex{} + seen := []int64{} + + cfg := ConsumerConfig{ + ProcessFunc: func(ctx context.Context, _ Source, msg eventstream.Event) error { + seenMu.Lock() + seen = append(seen, msg.Created) + seenMu.Unlock() + return nil + }, + WorkerCount: 1, + QueueSize: 16, + ConnectionTimeout: 2 * time.Second, + CursorStore: store, + URLFunc: DefaultURL(true), + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + } + c := NewConsumer(cfg) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + c.Start(ctx) + c.AddSource(ctx, source) + + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + seenMu.Lock() + n := len(seen) + seenMu.Unlock() + if n >= 2 { + break + } + time.Sleep(20 * time.Millisecond) + } + + seenMu.Lock() + defer seenMu.Unlock() + if len(seen) < 2 { + t.Fatalf("processed %d events, want 2: %v", len(seen), seen) + } + if seen[0] != 4 || seen[1] != 5 { + t.Fatalf("resumed events = %v, want [4 5]", seen) + } +} diff --git a/eventconsumer/knot.go b/eventconsumer/knot.go deleted file mode 100644 --- a/eventconsumer/knot.go +++ /dev/null @@ -1,39 +0,0 @@ -package eventconsumer - -import ( - "fmt" - "net/url" -) - -type KnotSource struct { - Knot string -} - -func (k KnotSource) Key() string { - return k.Knot -} - -func (k KnotSource) Url(cursor int64, dev bool) (*url.URL, error) { - scheme := "wss" - if dev { - scheme = "ws" - } - - u, err := url.Parse(scheme + "://" + k.Knot + "/events") - if err != nil { - return nil, err - } - - if cursor != 0 { - query := url.Values{} - query.Add("cursor", fmt.Sprintf("%d", cursor)) - u.RawQuery = query.Encode() - } - return u, nil -} - -func NewKnotSource(knot string) KnotSource { - return KnotSource{ - Knot: knot, - } -} diff --git a/eventconsumer/source.go b/eventconsumer/source.go new file mode 100644 --- /dev/null +++ b/eventconsumer/source.go @@ -0,0 +1,42 @@ +package eventconsumer + +import ( + "net/url" + "strconv" +) + +type Kind string + +const ( + KindKnot Kind = "knot" + KindSpindle Kind = "spindle" +) + +type Source struct { + Kind Kind + Host string +} + +func NewKnotSource(host string) Source { return Source{Kind: KindKnot, Host: host} } +func NewSpindleSource(host string) Source { return Source{Kind: KindSpindle, Host: host} } + +func (s Source) Key() string { return string(s.Kind) + ":" + s.Host } + +func DefaultURL(dev bool) func(Source, int64) (*url.URL, error) { + scheme := "wss" + if dev { + scheme = "ws" + } + return func(s Source, cursor int64) (*url.URL, error) { + u, err := url.Parse(scheme + "://" + s.Host + "/events") + if err != nil { + return nil, err + } + if cursor != 0 { + q := url.Values{} + q.Add("cursor", strconv.FormatInt(cursor, 10)) + u.RawQuery = q.Encode() + } + return u, nil + } +} diff --git a/eventconsumer/spindle.go b/eventconsumer/spindle.go deleted file mode 100644 --- a/eventconsumer/spindle.go +++ /dev/null @@ -1,39 +0,0 @@ -package eventconsumer - -import ( - "fmt" - "net/url" -) - -type SpindleSource struct { - Spindle string -} - -func (s SpindleSource) Key() string { - return s.Spindle -} - -func (s SpindleSource) Url(cursor int64, dev bool) (*url.URL, error) { - scheme := "wss" - if dev { - scheme = "ws" - } - - u, err := url.Parse(scheme + "://" + s.Spindle + "/events") - if err != nil { - return nil, err - } - - if cursor != 0 { - query := url.Values{} - query.Add("cursor", fmt.Sprintf("%d", cursor)) - u.RawQuery = query.Encode() - } - return u, nil -} - -func NewSpindleSource(spindle string) SpindleSource { - return SpindleSource{ - Spindle: spindle, - } -} diff --git a/eventconsumer/cursor/memory.go b/eventconsumer/cursor/memory.go --- a/eventconsumer/cursor/memory.go +++ b/eventconsumer/cursor/memory.go @@ -8,16 +8,15 @@ store sync.Map } -func (m *MemoryStore) Set(knot string, cursor int64) { - m.store.Store(knot, cursor) +func (m *MemoryStore) Set(key string, cursor int64) { + m.store.Store(key, cursor) } -func (m *MemoryStore) Get(knot string) (cursor int64) { - if result, ok := m.store.Load(knot); ok { +func (m *MemoryStore) Get(key string) (cursor int64) { + if result, ok := m.store.Load(key); ok { if val, ok := result.(int64); ok { return val } } - return 0 } diff --git a/eventconsumer/cursor/redis.go b/eventconsumer/cursor/redis.go --- a/eventconsumer/cursor/redis.go +++ b/eventconsumer/cursor/redis.go @@ -22,22 +22,20 @@ } } -func (r *RedisStore) Set(knot string, cursor int64) { - key := fmt.Sprintf(cursorKey, knot) - r.rdb.Set(context.Background(), key, cursor, 0) +func (r *RedisStore) Set(key string, cursor int64) { + k := fmt.Sprintf(cursorKey, key) + r.rdb.Set(context.Background(), k, cursor, 0) } -func (r *RedisStore) Get(knot string) (cursor int64) { - key := fmt.Sprintf(cursorKey, knot) - val, err := r.rdb.Get(context.Background(), key).Result() +func (r *RedisStore) Get(key string) (cursor int64) { + k := fmt.Sprintf(cursorKey, key) + val, err := r.rdb.Get(context.Background(), k).Result() if err != nil { return 0 } - cursor, err = strconv.ParseInt(val, 10, 64) + parsed, err := strconv.ParseInt(val, 10, 64) if err != nil { - // TODO: log here return 0 } - - return cursor + return parsed } diff --git a/eventconsumer/cursor/sqlite.go b/eventconsumer/cursor/sqlite.go --- a/eventconsumer/cursor/sqlite.go +++ b/eventconsumer/cursor/sqlite.go @@ -2,7 +2,9 @@ import ( "database/sql" + "errors" "fmt" + "log/slog" _ "github.com/mattn/go-sqlite3" ) @@ -46,35 +48,33 @@ createTable := fmt.Sprintf(` create table if not exists %s ( knot text primary key, - cursor text + cursor integer );`, s.tableName) _, err := s.db.Exec(createTable) return err } -func (s *SqliteStore) Set(knot string, cursor int64) { +func (s *SqliteStore) Set(key string, cursor int64) { query := fmt.Sprintf(` insert into %s (knot, cursor) values (?, ?) on conflict(knot) do update set cursor=excluded.cursor; `, s.tableName) - _, err := s.db.Exec(query, knot, cursor) - - if err != nil { - // TODO: log here + if _, err := s.db.Exec(query, key, cursor); err != nil { + slog.Default().Error("cursor sqlite set failed", "key", key, "cursor", cursor, "err", err) } } -func (s *SqliteStore) Get(knot string) (cursor int64) { +func (s *SqliteStore) Get(key string) (cursor int64) { query := fmt.Sprintf(` select cursor from %s where knot = ?; `, s.tableName) - err := s.db.QueryRow(query, knot).Scan(&cursor) + err := s.db.QueryRow(query, key).Scan(&cursor) if err != nil { - if err != sql.ErrNoRows { - // TODO: log here + if !errors.Is(err, sql.ErrNoRows) { + slog.Default().Error("cursor sqlite get failed", "key", key, "err", err) } return 0 } diff --git a/eventconsumer/cursor/store.go b/eventconsumer/cursor/store.go --- a/eventconsumer/cursor/store.go +++ b/eventconsumer/cursor/store.go @@ -1,6 +1,6 @@ package cursor type Store interface { - Set(knot string, cursor int64) - Get(knot string) (cursor int64) + Set(key string, cursor int64) + Get(key string) (cursor int64) }