diff --git a/tapclient/doc.go b/tapclient/doc.go new file mode 100644 index 0000000..de4ef7d --- /dev/null +++ b/tapclient/doc.go @@ -0,0 +1,37 @@ +// Package tap provides a client for consuming atproto events from a tap websocket. +// +// (this is jcalabro code from https://github.com/bluesky-social/indigo/pull/1241) +// +// The client handles connection management, automatic reconnection with backoff, +// and optional message acknowledgements. +// +// Basic usage: +// +// handler := func(ctx context.Context, ev *tap.Event) error { +// switch payload := ev.Payload().(type) { +// case *tap.RecordEvent: +// fmt.Printf("record.Action: %s\n", payload.Action) +// fmt.Printf("record.Collection: %s\n", payload.Collection) +// case *tap.IdentityEvent: +// fmt.Printf("identity.DID: %s\n", payload.DID) +// fmt.Printf("identity.Handle: %s\n", payload.Handle) +// } +// return nil +// } +// +// ws, err := tap.NewWebsocket("wss://example.com/tap", handler, +// tap.WithLogger(slog.Default()), +// tap.WithAcks(), +// ) +// if err != nil { +// // handle error... +// } +// +// if err := ws.Run(ctx); err != nil { +// // handle error... +// } +// +// Returning an error from the handler will cause the message to be retried with +// exponential backoff. To skip retries for permanent failures, wrap the error +// with [NewNonRetryableError]. +package tapclient diff --git a/tapclient/event.go b/tapclient/event.go new file mode 100644 index 0000000..5d1d49c --- /dev/null +++ b/tapclient/event.go @@ -0,0 +1,114 @@ +package tapclient + +import ( + "encoding/json" + "fmt" +) + +const ( + eventTypeACK = "ack" + eventTypeRecord = "record" + eventTypeIdentity = "identity" +) + +// Event represents an atproto event from tap. Use a type switch on the Payload() method to access event data. +type Event struct { + ID uint64 + Type string + + record *RecordEvent + identity *IdentityEvent +} + +// RecordEvent represents a record creation, update, or deletion in a repository +type RecordEvent struct { + DID string `json:"did"` + Collection string `json:"collection"` + Rkey string `json:"rkey"` + Action string `json:"action"` + CID string `json:"cid"` + Record json.RawMessage `json:"record"` + Live bool `json:"live"` +} + +// IdentityEvent represents an account status change +type IdentityEvent struct { + DID string `json:"did"` + Handle string `json:"handle"` + IsActive bool `json:"isActive"` + Status string `json:"status"` +} + +func (e *Event) UnmarshalJSON(data []byte) error { + event := struct { + ID uint64 `json:"id"` + Type string `json:"type"` + Record json.RawMessage `json:"record,omitempty"` + Identity json.RawMessage `json:"identity,omitempty"` + }{} + + if err := json.Unmarshal(data, &event); err != nil { + return fmt.Errorf("failed to unmarshal tap event: %w", err) + } + + e.ID = event.ID + e.Type = event.Type + + switch event.Type { + case eventTypeRecord: + e.record = &RecordEvent{} + if err := json.Unmarshal(event.Record, e.record); err != nil { + return fmt.Errorf("failed to unmarshal tap record event: %w", err) + } + case eventTypeIdentity: + e.identity = &IdentityEvent{} + if err := json.Unmarshal(event.Identity, e.identity); err != nil { + return fmt.Errorf("failed to unmarshal tap identity event: %w", err) + } + default: + return fmt.Errorf("unknown event type %q", event.Type) + } + + return nil +} + +func (e Event) MarshalJSON() ([]byte, error) { + event := struct { + ID uint64 `json:"id"` + Type string `json:"type"` + Record *RecordEvent `json:"record,omitempty"` + Identity *IdentityEvent `json:"identity,omitempty"` + }{ + ID: e.ID, + Type: e.Type, + Record: e.record, + Identity: e.identity, + } + + buf, err := json.Marshal(event) + if err != nil { + return nil, fmt.Errorf("failed to marshal tap event: %w", err) + } + + return buf, nil +} + +// Payload returns the typed event data as either *RecordEvent or *IdentityEvent. +func (e *Event) Payload() any { + switch e.Type { + case eventTypeRecord: + return e.record + case eventTypeIdentity: + return e.identity + } + + return nil // unreachable +} + +// Constructs a new ACK object to be serialized and sent back to tap +func NewACKPayload(id uint64) *Event { + return &Event{ + Type: eventTypeACK, + ID: id, + } +} diff --git a/tapclient/event_test.go b/tapclient/event_test.go new file mode 100644 index 0000000..972c97f --- /dev/null +++ b/tapclient/event_test.go @@ -0,0 +1,108 @@ +package tapclient + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestEventJSON(t *testing.T) { + t.Parallel() + + t.Run("marshal/unmarshal record", func(t *testing.T) { + t.Parallel() + require := require.New(t) + + original := Event{ + ID: 123, + Type: eventTypeRecord, + record: &RecordEvent{ + DID: "did:plc:test", + Collection: "app.bsky.feed.post", + Rkey: "abc123", + Action: "create", + CID: "bafytest", + Record: json.RawMessage(`{"text":"hello"}`), + Live: true, + }, + } + + buf, err := json.Marshal(original) + require.NoError(err) + + var decoded Event + require.NoError(json.Unmarshal(buf, &decoded)) + require.Equal(original.ID, decoded.ID) + require.Equal(original.Type, decoded.Type) + + payload, ok := decoded.Payload().(*RecordEvent) + require.True(ok) + require.Equal(original.record.DID, payload.DID) + require.Equal(original.record.Collection, payload.Collection) + require.Equal(original.record.Rkey, payload.Rkey) + require.Equal(original.record.Action, payload.Action) + require.Equal(original.record.CID, payload.CID) + require.Equal(original.record.Live, payload.Live) + require.JSONEq(string(original.record.Record), string(payload.Record)) + }) + + t.Run("marshal/unmarshal identity", func(t *testing.T) { + t.Parallel() + require := require.New(t) + + original := Event{ + ID: 456, + Type: eventTypeIdentity, + identity: &IdentityEvent{ + DID: "did:plc:user", + Handle: "test.bsky.social", + IsActive: true, + Status: "active", + }, + } + + buf, err := json.Marshal(original) + require.NoError(err) + + var decoded Event + require.NoError(json.Unmarshal(buf, &decoded)) + require.Equal(original.ID, decoded.ID) + require.Equal(original.Type, decoded.Type) + + payload, ok := decoded.Payload().(*IdentityEvent) + require.True(ok) + require.Equal(original.identity.DID, payload.DID) + require.Equal(original.identity.Handle, payload.Handle) + require.Equal(original.identity.IsActive, payload.IsActive) + require.Equal(original.identity.Status, payload.Status) + }) + + t.Run("unmarshal from raw json", func(t *testing.T) { + t.Parallel() + require := require.New(t) + + recordJSON := `{"id":1,"type":"record","record":{"did":"did:plc:abc","collection":"app.bsky.feed.like","rkey":"xyz","action":"create","cid":"mycid","record":{"subject":"at://did:plc:foo/app.bsky.feed.post/xyz"},"live":false}}` + var recordEvent Event + require.NoError(json.Unmarshal([]byte(recordJSON), &recordEvent)) + require.Equal(uint64(1), recordEvent.ID) + require.Equal(eventTypeRecord, recordEvent.Type) + + identityJSON := `{"id":2,"type":"identity","identity":{"did":"did:plc:def","handle":"foo.test","isActive":true,"status":"active"}}` + var identEvent Event + require.NoError(json.Unmarshal([]byte(identityJSON), &identEvent)) + require.Equal(uint64(2), identEvent.ID) + require.Equal(eventTypeIdentity, identEvent.Type) + }) + + t.Run("unmarshal unknown type", func(t *testing.T) { + t.Parallel() + require := require.New(t) + + badJSON := `{"id":1,"type":"unknown"}` + var ev Event + err := json.Unmarshal([]byte(badJSON), &ev) + require.Error(err) + require.Contains(err.Error(), "unknown event type") + }) +} diff --git a/tapclient/websocket.go b/tapclient/websocket.go new file mode 100644 index 0000000..0bde262 --- /dev/null +++ b/tapclient/websocket.go @@ -0,0 +1,320 @@ +package tapclient + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/url" + "sync" + "time" + + "github.com/gorilla/websocket" +) + +var ( + initialBackoff = 500 * time.Millisecond +) + +// A thin error wrapper that indicates to the tap client consumer loop that a message +// should not be retried (i.e. invalid user input that will surely fail again on retry). +type NonRetryableError struct { + err error +} + +func NewNonRetryableError(err error) *NonRetryableError { + return &NonRetryableError{err: err} +} + +func (err *NonRetryableError) Error() string { + if err.err != nil { + return err.err.Error() + } + return "" +} + +// Websocket implements a tap consumer that reads via a websocket +type Websocket struct { + log *slog.Logger + + addr string + sendAcks bool + maxErrs int + + connectTimeout time.Duration + readTimeout time.Duration + writeTimeout time.Duration + + handler WebsocketHandlerFunc +} + +// Defines an option for the tap websocket consumer +type WebsocketOption func(*Websocket) + +// Defines the log/slog logger to use throughout the lifecycle of the websocket +// consumer. Pass nil to disable logging. +func WithLogger(logger *slog.Logger) func(*Websocket) { + return func(ws *Websocket) { + ws.log = logger + + if ws.log == nil { + // write to io.Discard if a nil logger is passed + ws.log = slog.New(slog.NewTextHandler(io.Discard, nil)) + } + } +} + +// Sets the connect timeout for connecting to the websocket +func WithConnectTimeout(timeout time.Duration) func(*Websocket) { + return func(ws *Websocket) { + ws.connectTimeout = timeout + } +} + +// Sets the read timeout for reading data from the websocket +func WithReadTimeout(timeout time.Duration) func(*Websocket) { + return func(ws *Websocket) { + ws.readTimeout = timeout + } +} + +// Sets the write timeout for writing data to the websocket +func WithWriteTimeout(timeout time.Duration) func(*Websocket) { + return func(ws *Websocket) { + ws.writeTimeout = timeout + } +} + +// Controls how many times the loop will attempt to reconnect to the websocket in a row before giving up +func WithMaxConsecutiveErrors(numErrs int) func(*Websocket) { + return func(ws *Websocket) { + ws.maxErrs = numErrs + } +} + +// Turns on message acknowledgements +func WithAcks() func(*Websocket) { + return func(ws *Websocket) { + ws.sendAcks = true + } +} + +// Defines an option for the tap websocket consumer. A nil error indicates that an ACK will be sent to tap +// if WithAcks() is provided. +type WebsocketHandlerFunc func(context.Context, *Event) error + +// Initializes a tap websocket consumer +func NewWebsocket(addr string, handler WebsocketHandlerFunc, opts ...WebsocketOption) (*Websocket, error) { + u, err := url.Parse(addr) + if err != nil { + return nil, fmt.Errorf("failed to parse websocket url %q: %w", addr, err) + } + + switch u.Scheme { + case "ws", "wss": // ok + default: + return nil, fmt.Errorf("invalid websocket protocol scheme: wanted ws:// or wss://, got %q", u.Scheme) + } + + if handler == nil { + return nil, fmt.Errorf("a websocket message handler func is required") + } + + ws := &Websocket{ + log: slog.Default().WithGroup("tap"), + + addr: addr, + sendAcks: false, + maxErrs: 10, + + connectTimeout: 30 * time.Second, + readTimeout: 30 * time.Second, + writeTimeout: 30 * time.Second, + + handler: handler, + } + + for _, opt := range opts { + opt(ws) + } + + return ws, nil +} + +// Connects to and beings the main tap websocket consumer loop +func (ws *Websocket) Run(ctx context.Context) error { + for errCount := 0; ; { + select { + case <-ctx.Done(): + ws.log.Debug("websocket ingester shutting down") + return nil + default: + } + + err := ws.runOnce(ctx) + if errors.Is(err, context.Canceled) { + ws.log.Debug("websocket ingester shutting down") + return nil + } + + if err == nil { + errCount = 0 + ws.log.Debug("websocket connection closed normally, reconnecting") + continue + } + + errCount++ + ws.log.Error("websocket connection failed", "err", err, "consecutive_errors", errCount) + + if errCount >= ws.maxErrs { + return fmt.Errorf("websocket connection failed %d consecutive times: %w", errCount, err) + } + + ws.log.Warn("retrying websocket connection", "consecutive_errors", errCount) + if sleepMaybeExit(ctx, errCount) { + return nil + } + } +} + +func (ws *Websocket) close(conn *websocket.Conn) { + if err := conn.Close(); err != nil { + ws.log.Error("failed to close websocket connection", "err", err) + } +} + +func (ws *Websocket) runOnce(ctx context.Context) error { + dialer := websocket.Dialer{HandshakeTimeout: ws.connectTimeout} + conn, _, err := dialer.DialContext(ctx, ws.addr, nil) + if err != nil { + return fmt.Errorf("failed to connect to websocket at %q: %w", ws.addr, err) + } + + var closeOnce sync.Once + closeConn := func() { + closeOnce.Do(func() { + if err := conn.Close(); err != nil { + ws.log.Error("failed to close websocket connection", "err", err) + } + }) + } + defer closeConn() + + ws.log.Debug("connected to websocket", "addr", ws.addr) + + go func() { + <-ctx.Done() + closeConn() + }() + + for { + if done(ctx) { + return nil + } + + if err := conn.SetReadDeadline(time.Now().Add(ws.readTimeout)); err != nil { + return fmt.Errorf("failed to set websocket read deadline: %w", err) + } + + _, buf, err := conn.ReadMessage() + if err != nil { + if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) { + return nil // normal remote closure + } + + if ctx.Err() != nil { + return ctx.Err() + } + + return fmt.Errorf("failed to read websocket message: %w", err) + } + + var ev Event + if err := json.Unmarshal(buf, &ev); err != nil { + ws.log.Warn("failed to unmarshal event json", "err", err) + continue + } + + // indefinitely retry messages that failed to process unless a non-retryable error occurrs + for errCount := 0; ; errCount++ { + if done(ctx) { + break + } + + err := ws.handler(ctx, &ev) + if err == nil { + break + } + + ws.log.Error("failed to process event", "err", err) + if sleepMaybeExit(ctx, errCount) { + return nil + } + + var nr *NonRetryableError + if errors.As(err, &nr) { + ws.log.Error("handled non-retryable error", "id", ev.ID, "err", err) + break + } + } + + if ws.sendAcks { + ws.ack(ctx, conn, &ev) + } + } +} + +// Indefinitely tries acking the message with the tap server +func (ws *Websocket) ack(ctx context.Context, conn *websocket.Conn, ev *Event) { + for errCount := 0; ; errCount++ { + if done(ctx) { + return + } + + if err := conn.SetWriteDeadline(time.Now().Add(ws.writeTimeout)); err != nil { + ws.log.Warn("failed to set write deadline on ack", "err", err) + } + + err := conn.WriteJSON(NewACKPayload(ev.ID)) + if err == nil { + return + } + + if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) { + return // normal remote closure + } + + ws.log.Error("failed to send ack", "err", err) + + if sleepMaybeExit(ctx, errCount) { + return + } + } +} + +func done(ctx context.Context) bool { + select { + case <-ctx.Done(): + return true + default: + return false + } +} + +func sleepMaybeExit(ctx context.Context, errCount int) bool { + select { + case <-ctx.Done(): + return true // shutdown received during a backoff sleep means that we're done + case <-time.After(backoffDuration(errCount)): + return false + } +} + +func backoffDuration(errCount int) time.Duration { + multiplier := 1 << errCount + waitFor := initialBackoff * time.Duration(multiplier) + + return min(waitFor, 10*time.Second) +} diff --git a/tapclient/websocket_test.go b/tapclient/websocket_test.go new file mode 100644 index 0000000..95f986b --- /dev/null +++ b/tapclient/websocket_test.go @@ -0,0 +1,278 @@ +package tapclient + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/require" +) + +func init() { + initialBackoff = 0 +} + +var upgrader = websocket.Upgrader{} + +func TestWebsocket(t *testing.T) { + t.Parallel() + ctx := t.Context() + require := require.New(t) + + events := []Event{ + {ID: 1, Type: eventTypeRecord, record: &RecordEvent{DID: "did:plc:1", Collection: "app.bsky.feed.post"}}, + {ID: 2, Type: eventTypeRecord, record: &RecordEvent{DID: "did:plc:2", Collection: "app.bsky.feed.like"}}, + {ID: 3, Type: eventTypeIdentity, identity: &IdentityEvent{DID: "did:plc:3", Handle: "user3.test"}}, + } + + var received []*Event + var mu sync.Mutex + var wg sync.WaitGroup + wg.Add(len(events)) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + + for _, ev := range events { + buf, _ := json.Marshal(ev) + conn.WriteMessage(websocket.TextMessage, buf) + time.Sleep(10 * time.Millisecond) + } + + time.Sleep(50 * time.Millisecond) + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + })) + defer server.Close() + + wsURL := "ws://" + strings.TrimPrefix(server.URL, "http://") + + ws, err := NewWebsocket(wsURL, func(ctx context.Context, ev *Event) error { + mu.Lock() + received = append(received, ev) + mu.Unlock() + wg.Done() + return nil + }, WithLogger(nil)) + require.NoError(err) + + go ws.Run(ctx) + wg.Wait() + + require.Len(received, 3) + for i, ev := range received { + require.Equal(uint64(i+1), ev.ID) + + switch i { + case 0, 1: + switch pl := ev.Payload().(type) { + case *RecordEvent: + require.NotNil(events[i].record) + require.Equal(events[i].record.Collection, pl.Collection) + require.Equal(events[i].Type, eventTypeRecord) + default: + require.FailNow("incorrect payload type, want %T got %T", &RecordEvent{}, ev.Payload()) + } + + case 2: + switch pl := ev.Payload().(type) { + case *IdentityEvent: + require.NotNil(events[i].identity) + require.Equal(events[i].identity.Handle, pl.Handle) + require.Equal(events[i].Type, eventTypeIdentity) + default: + require.FailNow("incorrect payload type, want %T got %T", &IdentityEvent{}, ev.Payload()) + } + } + } +} + +func TestWebsocketWithAcks(t *testing.T) { + t.Parallel() + + t.Run("ack sent on success", func(t *testing.T) { + t.Parallel() + ctx := t.Context() + require := require.New(t) + + recordEvent := Event{ + ID: 42, + Type: eventTypeRecord, + record: &RecordEvent{ + DID: "did:plc:ack", + Collection: "app.bsky.feed.like", + Rkey: "ack", + Action: "create", + }, + } + + var receivedAck *Event + var wg sync.WaitGroup + wg.Add(1) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + + buf, _ := json.Marshal(recordEvent) + conn.WriteMessage(websocket.TextMessage, buf) + + _, ackBuf, err := conn.ReadMessage() + if err == nil { + receivedAck = &Event{} + json.Unmarshal(ackBuf, receivedAck) + } + wg.Done() + + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + })) + defer server.Close() + + wsURL := "ws://" + strings.TrimPrefix(server.URL, "http://") + + ws, err := NewWebsocket(wsURL, func(ctx context.Context, ev *Event) error { + return nil + }, WithLogger(nil), WithAcks()) + require.NoError(err) + + go ws.Run(ctx) + wg.Wait() + + require.NotNil(receivedAck) + require.Equal(eventTypeACK, receivedAck.Type) + require.Equal(recordEvent.ID, receivedAck.ID) + }) + + t.Run("ack not sent on error", func(t *testing.T) { + t.Parallel() + ctx := t.Context() + require := require.New(t) + + recordEvent := Event{ + ID: 99, + Type: eventTypeRecord, + record: &RecordEvent{ + DID: "did:plc:noack", + Collection: "app.bsky.feed.post", + Rkey: "noack", + Action: "create", + }, + } + + var receivedAck bool + var wg sync.WaitGroup + wg.Add(1) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + + buf, _ := json.Marshal(recordEvent) + conn.WriteMessage(websocket.TextMessage, buf) + + conn.SetReadDeadline(time.Now().Add(100 * time.Millisecond)) + _, _, err = conn.ReadMessage() + receivedAck = err == nil + wg.Done() + + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + })) + defer server.Close() + + wsURL := "ws://" + strings.TrimPrefix(server.URL, "http://") + + ws, err := NewWebsocket(wsURL, func(ctx context.Context, ev *Event) error { + return errors.New("processing failed") + }, WithLogger(nil), WithAcks()) + require.NoError(err) + + go ws.Run(ctx) + wg.Wait() + + require.False(receivedAck, "expected no ACK when handler returns error") + }) +} + +func TestWebsocketNonRetryableError(t *testing.T) { + t.Parallel() + ctx := t.Context() + require := require.New(t) + + events := []Event{ + {ID: 1, Type: eventTypeRecord, record: &RecordEvent{DID: "did:plc:1", Collection: "app.bsky.feed.post"}}, + {ID: 2, Type: eventTypeRecord, record: &RecordEvent{DID: "did:plc:2", Collection: "app.bsky.feed.post"}}, + {ID: 3, Type: eventTypeRecord, record: &RecordEvent{DID: "did:plc:3", Collection: "app.bsky.feed.post"}}, + } + + var callCounts sync.Map + var wg sync.WaitGroup + wg.Add(len(events)) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + + for _, ev := range events { + buf, _ := json.Marshal(ev) + conn.WriteMessage(websocket.TextMessage, buf) + time.Sleep(10 * time.Millisecond) + } + + time.Sleep(100 * time.Millisecond) + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + })) + defer server.Close() + + wsURL := "ws://" + strings.TrimPrefix(server.URL, "http://") + + ws, err := NewWebsocket(wsURL, func(ctx context.Context, ev *Event) error { + val, _ := callCounts.LoadOrStore(ev.ID, new(int)) + count := val.(*int) + *count++ + + if ev.ID == 2 { + if *count == 1 { + wg.Done() + } + return NewNonRetryableError(errors.New("bad input, do not retry")) + } + + wg.Done() + return nil + }, WithLogger(nil)) + require.NoError(err) + + go ws.Run(ctx) + wg.Wait() + + // event 1: should be called exactly once (success) + val1, _ := callCounts.Load(uint64(1)) + require.Equal(1, *val1.(*int), "event 1 should be processed once") + + // event 2: should be called exactly once (non-retryable error, no retry) + val2, _ := callCounts.Load(uint64(2)) + require.Equal(1, *val2.(*int), "event 2 with NonRetryableError should not be retried") + + // event 3: should be called exactly once (success, proving we moved on after non-retryable) + val3, _ := callCounts.Load(uint64(3)) + require.Equal(1, *val3.(*int), "event 3 should be processed after non-retryable error") +}