package main import ( "context" "fmt" "net/http" "strconv" "time" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/gorilla/websocket" ) type LabelsMessage struct { Seq int64 Frame []byte } var ( wsUpgrader = websocket.Upgrader{} connMsgBufferDepth = 512 pingPeriod = 60 * time.Second pongWait = pingPeriod + 20*time.Second writeWait = 10 * time.Second ) func (sub *Subscriber) readLoop() { defer func() { sub.broker.unregister <- sub sub.conn.Close() }() sub.conn.SetReadLimit(32 * 1024) // only expecting ping/pong sub.conn.SetPongHandler(func(string) error { sub.conn.SetReadDeadline(time.Now().Add(pongWait)) return nil }) // TODO: also register a ping handler? for { if _, _, err := sub.conn.ReadMessage(); err != nil { return } } } func (sub *Subscriber) writeLoop() { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() sub.conn.Close() }() for { select { case msg, ok := <-sub.send: sub.conn.SetWriteDeadline(time.Now().Add(writeWait)) if !ok { // clean connection shutdown (via defered close) sub.conn.WriteMessage(websocket.CloseMessage, []byte{}) return } if msg.Seq <= sub.seq { sub.broker.logger.Warn("skipping dupe event", "msg.Seq", msg.Seq, "seq", sub.seq) continue } if err := sub.conn.WriteMessage(websocket.BinaryMessage, msg.Frame); err != nil { sub.broker.logger.Warn("failed to write websocket message frame", "err", err) return } sub.seq = msg.Seq case <-ticker.C: sub.conn.SetWriteDeadline(time.Now().Add(writeWait)) if err := sub.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } } func (srv *Server) SubscribeLabelsEndpoint(w http.ResponseWriter, r *http.Request) { ctx := r.Context() seq := int64(0) backfill := false params := r.URL.Query() if params.Has("cursor") { cursor, err := strconv.ParseInt(params.Get("cursor"), 10, 64) if err != nil { w.Header().Set("Content-Type", "application/json") http.Error(w, fmt.Sprintf(`{"error": "BadRequest", "message": "bad cursor param: %s"}`, err.Error()), http.StatusBadRequest) return } if cursor > 0 { high, err := srv.Store.LatestSeq(ctx) if err != nil { w.Header().Set("Content-Type", "application/json") http.Error(w, fmt.Sprintf(`{"error": "InternalServerError", "message": "%s"}`, err.Error()), http.StatusInternalServerError) return } // TODO: verify that this check is working (?) if cursor > high { w.Header().Set("Content-Type", "application/json") http.Error(w, fmt.Sprintf(`{"error": "FutureCursor", "message": "Current high cursor: %d"}`, high), http.StatusBadRequest) return } } seq = cursor backfill = true } conn, err := wsUpgrader.Upgrade(w, r, nil) if err != nil { w.Header().Set("Content-Type", "application/json") http.Error(w, fmt.Sprintf(`{"error": "BadRequest", "message": "%s"}`, err.Error()), http.StatusBadRequest) return } sub := &Subscriber{ broker: srv.Broker, conn: conn, send: make(chan LabelsMessage, connMsgBufferDepth), // TODO: better random identifier? include the IP address? id: syntax.NewTIDNow(0).String(), seq: seq, } // start to read (and discard) messages go sub.readLoop() // if cursor was provided, backfill from database first if backfill { if err = srv.backfillSubscriber(ctx, sub, 0); err != nil { srv.Logger.Warn("backfilling failed", "err", err) conn.Close() return } } // register subscriber (start receiving messages) srv.Broker.register <- sub // if we did backfill, do one last pass to catch any missed rows if backfill { if err = srv.backfillSubscriber(ctx, sub, 1); err != nil { srv.Logger.Warn("follow-up backfilling failed", "err", err) srv.Broker.unregister <- sub conn.Close() return } } // start enforcing keepalive ping/pong sub.conn.SetReadDeadline(time.Now().Add(pongWait)) // write broadcast messages go sub.writeLoop() } func (srv *Server) backfillSubscriber(ctx context.Context, sub *Subscriber, inc int) error { // increment is to move one forward in follow-up batch rows, err := srv.Store.ScanLabelRows(ctx, sub.seq+int64(inc)) if err != nil { return err } for { if len(rows) == 0 { break } for _, row := range rows { if sub.seq > int64(row.ID) { continue } sub.conn.SetWriteDeadline(time.Now().Add(writeWait)) msg, err := rowToMessage(row) if err != nil { return err } if err := sub.conn.WriteMessage(websocket.BinaryMessage, msg.Frame); err != nil { return err } sub.seq = msg.Seq } // next batch should be "plus one" to prevent looping on the same row rows, err = srv.Store.ScanLabelRows(ctx, sub.seq+1) if err != nil { return err } } return nil }