package knotfeed import ( "bytes" "context" "errors" "fmt" "io" "log/slog" "net/http" "net/http/httptest" "slices" "sync/atomic" "testing" "time" comatproto "github.com/bluesky-social/indigo/api/atproto" lexutil "github.com/bluesky-social/indigo/lex/util" "github.com/gorilla/websocket" ) func poisonTestLogger() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) } func commitFrameForSeq(t *testing.T, seq int64) []byte { t.Helper() var payload bytes.Buffer evt := comatproto.SyncSubscribeRepos_Commit{ Repo: "did:plc:scallop", Seq: seq, Rev: "3lb2xkw2qrs2j", Commit: lexutil.LexLink(testCid(t)), } if err := evt.MarshalCBOR(&payload); err != nil { t.Fatalf("MarshalCBOR: %v", err) } return append(headerFrame(t, "t", "#commit"), payload.Bytes()...) } func frameOnceServer(t *testing.T, frame []byte) *httptest.Server { t.Helper() upgrader := websocket.Upgrader{} srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } if err := conn.WriteMessage(websocket.BinaryMessage, frame); err != nil { return } _ = conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second)) conn.Close() })) t.Cleanup(srv.Close) return srv } func poisonConsumer(srv *httptest.Server, cursor *atomic.Int64, stored *atomic.Int64, now *atomic.Int64, grace time.Duration) *Consumer { return &Consumer{ Host: srv.Listener.Addr().String(), NoTLS: true, Logger: poisonTestLogger(), LoadCursor: func(context.Context) (Cursor, error) { return NewCursor(FeedAtproto, cursor.Load()), nil }, StoreCursor: func(_ context.Context, cur Cursor) error { cursor.Store(cur.Seq()) stored.Store(cur.Seq()) return nil }, Handle: func(context.Context, Message) error { return errors.New("boom") }, MaxHandlerAttempts: 3, PoisonGrace: grace, now: func() time.Time { return time.Unix(0, now.Load()) }, } } func poisonSequence(t *testing.T, c *Consumer, now, stored *atomic.Int64) []int64 { t.Helper() var seen []int64 for attempt := range 3 { now.Store(int64(attempt+1) * 6 * int64(time.Minute)) _, _, _ = c.session(context.Background()) seen = append(seen, stored.Load()) } return seen } func TestPoisonFrameSkipsAfterThreeFailuresAcrossTheGrace(t *testing.T) { srv := frameOnceServer(t, commitFrameForSeq(t, 100)) var cursor, stored, now atomic.Int64 c := poisonConsumer(srv, &cursor, &stored, &now, defaultPoisonGrace) if seen := poisonSequence(t, c, &now, &stored); !slices.Equal(seen, []int64{0, 0, 100}) { t.Fatalf("cursor stores per attempt = %v, want [0 0 100] once the poison lands", seen) } } func TestPoisonFrameHoldsWhileFailuresStayInsideTheGrace(t *testing.T) { srv := frameOnceServer(t, commitFrameForSeq(t, 100)) var cursor, stored, now atomic.Int64 c := poisonConsumer(srv, &cursor, &stored, &now, time.Hour) if seen := poisonSequence(t, c, &now, &stored); !slices.Equal(seen, []int64{0, 0, 0}) { t.Fatalf("cursor stores per attempt = %v, want all zero while inside the grace", seen) } } type pathLog struct{ paths chan string } func (p *pathLog) add(path string) { select { case p.paths <- path: default: } } func (p *pathLog) take(t *testing.T, n int) []string { t.Helper() var seen []string for range n { select { case path := <-p.paths: seen = append(seen, path) case <-time.After(10 * time.Second): t.Fatalf("the consumer stopped after %v", seen) } } return seen } func wsServer(t *testing.T, refuse string, serve func(conn *websocket.Conn, connect int64)) (*httptest.Server, *pathLog) { t.Helper() upgrader := websocket.Upgrader{} log := &pathLog{paths: make(chan string, 8)} var connects atomic.Int64 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { log.add(r.URL.Path) if r.URL.Path == refuse { w.WriteHeader(http.StatusNotFound) return } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } defer conn.Close() serve(conn, connects.Add(1)) _ = conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second)) })) t.Cleanup(srv.Close) return srv, log } func feedConsumer(t *testing.T, srv *httptest.Server, stored Cursor) *Consumer { t.Helper() c := &Consumer{ Host: srv.Listener.Addr().String(), NoTLS: true, Logger: poisonTestLogger(), LoadCursor: func(context.Context) (Cursor, error) { return stored, nil }, StoreCursor: func(context.Context, Cursor) error { return nil }, Handle: func(context.Context, Message) error { return nil }, } c.fill() return c } func runConsumer(t *testing.T, c *Consumer) { t.Helper() ctx, cancel := context.WithCancel(context.Background()) done := make(chan struct{}) go func() { defer close(done) _ = c.Run(ctx) }() t.Cleanup(func() { cancel() <-done }) } func TestWrongKindFrameEndsSessionAsMismatch(t *testing.T) { srv, _ := wsServer(t, "", func(conn *websocket.Conn, _ int64) { _ = conn.WriteMessage(websocket.TextMessage, []byte(`{"nsid":"x","created":1}`)) }) c := feedConsumer(t, srv, Cursor{}) _, _, err := c.session(context.Background()) var mismatch feedMismatch if !errors.As(err, &mismatch) { t.Fatalf("session error = %v, want a feed mismatch", err) } if mismatch.feed != FeedAtproto { t.Errorf("mismatch = %v, want the firehose we asked for", mismatch.feed) } } func TestMismatchTextSpellsOutBothFeeds(t *testing.T) { got := feedMismatch{feed: FeedAtproto}.Error() if got != "atproto firehose sent the knot event stream's frames" { t.Errorf("mismatch reads %q", got) } } func TestConsumerLeavesFeedOnlyWhenKnotSaysTo(t *testing.T) { for _, tt := range []struct { name string proven bool err error want Feed }{ {"mismatch on proven feed", true, feedMismatch{feed: FeedAtproto}, FeedLegacy}, {"refused connection on unproven feed", false, connectError{errors.New("no route")}, FeedLegacy}, {"refused connection on proven feed", true, connectError{errors.New("connection reset")}, FeedAtproto}, {"session that failed for its own reasons", true, errors.New("boom"), FeedAtproto}, } { t.Run(tt.name, func(t *testing.T) { c := &Consumer{Logger: poisonTestLogger()} c.fill() c.proven = tt.proven c.reactTo(tt.err) if c.current != tt.want { t.Errorf("feed = %v, want %v", c.current, tt.want) } if switched := tt.want != FeedAtproto; switched && c.proven { t.Error("new feed counts as proven before delivering anything") } }) } } func TestQuietKnotStillProvesFeedItAccepted(t *testing.T) { srv, _ := wsServer(t, "", func(*websocket.Conn, int64) {}) c := feedConsumer(t, srv, NewCursor(FeedLegacy, 11)) advanced, _, err := c.session(context.Background()) if advanced { t.Error("empty session reported progress") } if err == nil { t.Fatal("closed connection didn't error out") } if !c.proven { t.Fatal("knot accepted the upgrade, and the feed still reads unproven") } c.reactTo(connectError{errors.New("connection refused")}) if c.current != FeedLegacy { t.Errorf("feed = %v; we left a path the knot had already served after one refusal", c.current) } } func TestSwitchStaysUnprovenUntilNewFeedDelivers(t *testing.T) { srv, paths := wsServer(t, eventsPath, func(conn *websocket.Conn, _ int64) { _ = conn.WriteMessage(websocket.BinaryMessage, commitFrameForSeq(t, 500)) _ = conn.WriteMessage(websocket.TextMessage, []byte(`{"nsid":"x","created":1}`)) time.Sleep(time.Second) }) runConsumer(t, feedConsumer(t, srv, Cursor{})) tried := paths.take(t, 3) if tried[1] != eventsPath { t.Fatalf("stray text frame sent us to %q, want the legacy path", tried[1]) } if tried[2] == eventsPath { t.Error("we retried the refused legacy path, so the switch counted as proven too early") } } func TestStoredFeedIsAdoptedOnceAndNeverUndoesSwitch(t *testing.T) { c := &Consumer{ Host: "127.0.0.1:1", NoTLS: true, Logger: poisonTestLogger(), LoadCursor: func(context.Context) (Cursor, error) { return NewCursor(FeedAtproto, 500), nil }, StoreCursor: func(context.Context, Cursor) error { return nil }, } c.fill() c.current = FeedLegacy _, _, err := c.session(context.Background()) if c.current != FeedAtproto { t.Fatalf("feed = %v, want the stored cursor's feed on the first session", c.current) } c.reactTo(err) if c.current != FeedLegacy { t.Fatalf("feed = %v, want a switch after the refused connection", c.current) } for attempt := range 3 { if _, _, _ = c.session(context.Background()); c.current != FeedLegacy { t.Fatalf("session %d re-adopted the stored feed %v", attempt+2, c.current) } } } func TestSubscriptionAsksForItsOwnFeedAndPosition(t *testing.T) { frozen := time.Unix(0, 1788257313450261754) const legacySeq = 1788245839553422000 live := fmt.Sprint(frozen.UnixNano()) for _, tt := range []struct { name string noTLS bool feed Feed replay bool stored Cursor want string }{ {"firehose position on the firehose", true, FeedAtproto, false, NewCursor(FeedAtproto, 42), "ws://oyster.cafe/xrpc/com.atproto.sync.subscribeRepos?cursor=42"}, {"nothing stored on the firehose", true, FeedAtproto, false, Cursor{}, "ws://oyster.cafe/xrpc/com.atproto.sync.subscribeRepos"}, {"legacy position on the legacy feed", false, FeedLegacy, false, NewCursor(FeedLegacy, legacySeq), "wss://oyster.cafe/events?cursor=" + fmt.Sprint(legacySeq)}, {"firehose position on the legacy feed", false, FeedLegacy, false, NewCursor(FeedAtproto, 42), "wss://oyster.cafe/events?cursor=" + live}, {"legacy position on the firehose", false, FeedAtproto, false, NewCursor(FeedLegacy, legacySeq), "wss://oyster.cafe/xrpc/com.atproto.sync.subscribeRepos"}, {"nothing stored on the legacy feed", false, FeedLegacy, false, Cursor{}, "wss://oyster.cafe/events?cursor=" + live}, {"nothing stored while replaying from the start", false, FeedLegacy, true, Cursor{}, "wss://oyster.cafe/events?cursor=0"}, } { c := &Consumer{Host: "oyster.cafe", NoTLS: tt.noTLS, ReplayFromStart: tt.replay} c.fill() c.current = tt.feed c.now = func() time.Time { return frozen } if got := c.url(c.sessionCursor(tt.stored)); got != tt.want { t.Errorf("%s asks for %q, want %q", tt.name, got, tt.want) } } } func TestEveryFeedAndCursorRoundTripsItsText(t *testing.T) { for _, feed := range []Feed{FeedAtproto, FeedLegacy} { if parsed, err := ParseFeed(feed.Token()); err != nil || parsed != feed { t.Errorf("ParseFeed(%q) = (%v, %v), want %v", feed.Token(), parsed, err, feed) } } for _, want := range []Cursor{NewCursor(FeedAtproto, 500), NewCursor(FeedLegacy, 1788245839553422000), {}} { got, err := ParseCursor(want.Encode()) if err != nil || got != want { t.Errorf("ParseCursor(%q) = (%v, %v), want %v", want.Encode(), got, err, want) } } if got, err := ParseCursor("500"); err != nil || got != NewCursor(FeedAtproto, 500) { t.Errorf("bare sequence parsed as (%v, %v), want the firehose at 500", got, err) } for _, raw := range []string{"carrier pigeon:7", "legacy:soon"} { if _, err := ParseCursor(raw); err == nil { t.Errorf("ParseCursor(%q) parsed anyway", raw) } } if _, err := ParseFeed("carrier pigeon"); !errors.Is(err, ErrUnrecognizedFeed) { t.Errorf("unrecognized token gave %v, want something callers can classify", err) } }