diff --git a/internal/db/postgres/discover_repo.go b/internal/db/postgres/discover_repo.go index 51053b8..36adb02 100644 --- a/internal/db/postgres/discover_repo.go +++ b/internal/db/postgres/discover_repo.go @@ -5,6 +5,7 @@ import ( "context" "database/sql" "fmt" + "time" ) type postgresDiscoverRepo struct { @@ -33,6 +34,9 @@ func NewDiscoverRepository(db *sql.DB, cursorSecret string) discover.Repository // GetDiscover retrieves posts from ALL communities (public feed) func (r *postgresDiscoverRepo) GetDiscover(ctx context.Context, req discover.GetDiscoverRequest) ([]*discover.FeedViewPost, *string, error) { + // Capture query time for stable cursor generation (used for hot sort pagination) + queryTime := time.Now() + // Build ORDER BY clause based on sort type orderBy, timeFilter := r.buildSortClause(req.Sort, req.Timeframe) @@ -119,7 +123,7 @@ func (r *postgresDiscoverRepo) GetDiscover(ctx context.Context, req discover.Get hotRanks = hotRanks[:req.Limit] lastPost := feedPosts[len(feedPosts)-1].Post lastHotRank := hotRanks[len(hotRanks)-1] - cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank) + cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank, queryTime) cursor = &cursorStr } diff --git a/internal/db/postgres/feed_repo.go b/internal/db/postgres/feed_repo.go index 0d7d114..9ab3de1 100644 --- a/internal/db/postgres/feed_repo.go +++ b/internal/db/postgres/feed_repo.go @@ -5,6 +5,7 @@ import ( "context" "database/sql" "fmt" + "time" ) type postgresFeedRepo struct { @@ -37,6 +38,9 @@ func NewCommunityFeedRepository(db *sql.DB, cursorSecret string) communityFeeds. // GetCommunityFeed retrieves posts from a community with sorting and pagination // Single query with JOINs for optimal performance func (r *postgresFeedRepo) GetCommunityFeed(ctx context.Context, req communityFeeds.GetCommunityFeedRequest) ([]*communityFeeds.FeedViewPost, *string, error) { + // Capture query time for stable cursor generation (used for hot sort pagination) + queryTime := time.Now() + // Build ORDER BY clause based on sort type orderBy, timeFilter := r.feedRepoBase.buildSortClause(req.Sort, req.Timeframe) @@ -125,7 +129,7 @@ func (r *postgresFeedRepo) GetCommunityFeed(ctx context.Context, req communityFe hotRanks = hotRanks[:req.Limit] lastPost := feedPosts[len(feedPosts)-1].Post lastHotRank := hotRanks[len(hotRanks)-1] - cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank) + cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank, queryTime) cursor = &cursorStr } diff --git a/internal/db/postgres/feed_repo_base.go b/internal/db/postgres/feed_repo_base.go index adbc36e..16a9dbf 100644 --- a/internal/db/postgres/feed_repo_base.go +++ b/internal/db/postgres/feed_repo_base.go @@ -192,15 +192,17 @@ func (r *feedRepoBase) parseCursor(cursor *string, sort string, paramOffset int) return filter, []interface{}{score, createdAt, uri}, nil case "hot": - // Cursor format: hot_rank::timestamp::uri - // CRITICAL: Must use computed hot_rank, not raw score, to prevent pagination bugs - if len(payloadParts) != 3 { + // Cursor format: hot_rank::post_created_at::uri::cursor_timestamp + // CRITICAL: cursor_timestamp is when the cursor was created, used for stable hot_rank comparison + // This prevents pagination bugs caused by hot_rank drift when NOW() changes between requests + if len(payloadParts) != 4 { return "", nil, fmt.Errorf("invalid cursor format for hot sort") } hotRankStr := payloadParts[0] - createdAt := payloadParts[1] + postCreatedAt := payloadParts[1] uri := payloadParts[2] + cursorTimestamp := payloadParts[3] // Validate hot_rank is numeric (float) hotRank := 0.0 @@ -208,9 +210,9 @@ func (r *feedRepoBase) parseCursor(cursor *string, sort string, paramOffset int) return "", nil, fmt.Errorf("invalid cursor hot rank") } - // Validate timestamp format - if _, err := time.Parse(time.RFC3339Nano, createdAt); err != nil { - return "", nil, fmt.Errorf("invalid cursor timestamp") + // Validate post timestamp format + if _, err := time.Parse(time.RFC3339Nano, postCreatedAt); err != nil { + return "", nil, fmt.Errorf("invalid cursor post timestamp") } // Validate URI format (must be AT-URI) @@ -218,13 +220,21 @@ func (r *feedRepoBase) parseCursor(cursor *string, sort string, paramOffset int) return "", nil, fmt.Errorf("invalid cursor URI") } - // CRITICAL: Compare against the computed hot_rank expression, not p.score - filter := fmt.Sprintf(`AND ((%s < $%d OR (%s = $%d AND p.created_at < $%d) OR (%s = $%d AND p.created_at = $%d AND p.uri < $%d)) AND p.uri != $%d)`, - r.hotRankExpression, paramOffset, - r.hotRankExpression, paramOffset, paramOffset+1, - r.hotRankExpression, paramOffset, paramOffset+1, paramOffset+2, + // Validate cursor timestamp format + if _, err := time.Parse(time.RFC3339Nano, cursorTimestamp); err != nil { + return "", nil, fmt.Errorf("invalid cursor timestamp") + } + + // CRITICAL: Use cursor_timestamp instead of NOW() for stable hot_rank comparison + // This ensures posts don't drift across page boundaries due to time passing + stableHotRankExpr := fmt.Sprintf( + `((p.score + 1) / POWER(EXTRACT(EPOCH FROM ($%d::timestamptz - p.created_at))/3600 + 2, 1.5))`, paramOffset+3) - return filter, []interface{}{hotRank, createdAt, uri, uri}, nil + + // Use tuple comparison for clean keyset pagination: (hot_rank, created_at, uri) < (cursor_values) + filter := fmt.Sprintf(`AND ((%s, p.created_at, p.uri) < ($%d, $%d, $%d))`, + stableHotRankExpr, paramOffset, paramOffset+1, paramOffset+2) + return filter, []interface{}{hotRank, postCreatedAt, uri, cursorTimestamp}, nil default: return "", nil, nil @@ -233,7 +243,8 @@ func (r *feedRepoBase) parseCursor(cursor *string, sort string, paramOffset int) // buildCursor creates HMAC-signed pagination cursor from last post // SECURITY: Cursor is signed with HMAC-SHA256 to prevent manipulation -func (r *feedRepoBase) buildCursor(post *posts.PostView, sort string, hotRank float64) string { +// queryTime is the timestamp when the query was executed, used for stable hot_rank comparison +func (r *feedRepoBase) buildCursor(post *posts.PostView, sort string, hotRank float64, queryTime time.Time) string { var payload string // Use :: as delimiter following Bluesky convention const delimiter = "::" @@ -252,10 +263,10 @@ func (r *feedRepoBase) buildCursor(post *posts.PostView, sort string, hotRank fl payload = fmt.Sprintf("%d%s%s%s%s", score, delimiter, post.CreatedAt.Format(time.RFC3339Nano), delimiter, post.URI) case "hot": - // Format: hot_rank::timestamp::uri - // CRITICAL: Use computed hot_rank with full precision + // Format: hot_rank::post_created_at::uri::cursor_timestamp + // CRITICAL: Include cursor_timestamp for stable hot_rank comparison across requests hotRankStr := strconv.FormatFloat(hotRank, 'g', -1, 64) - payload = fmt.Sprintf("%s%s%s%s%s", hotRankStr, delimiter, post.CreatedAt.Format(time.RFC3339Nano), delimiter, post.URI) + payload = fmt.Sprintf("%s%s%s%s%s%s%s", hotRankStr, delimiter, post.CreatedAt.Format(time.RFC3339Nano), delimiter, post.URI, delimiter, queryTime.Format(time.RFC3339Nano)) default: payload = post.URI diff --git a/internal/db/postgres/timeline_repo.go b/internal/db/postgres/timeline_repo.go index 47e7e7c..20ebed8 100644 --- a/internal/db/postgres/timeline_repo.go +++ b/internal/db/postgres/timeline_repo.go @@ -5,6 +5,7 @@ import ( "context" "database/sql" "fmt" + "time" ) type postgresTimelineRepo struct { @@ -35,6 +36,9 @@ func NewTimelineRepository(db *sql.DB, cursorSecret string) timeline.Repository // GetTimeline retrieves posts from all communities the user subscribes to // Single query with JOINs for optimal performance func (r *postgresTimelineRepo) GetTimeline(ctx context.Context, req timeline.GetTimelineRequest) ([]*timeline.FeedViewPost, *string, error) { + // Capture query time for stable cursor generation (used for hot sort pagination) + queryTime := time.Now() + // Build ORDER BY clause based on sort type orderBy, timeFilter := r.buildSortClause(req.Sort, req.Timeframe) @@ -125,7 +129,7 @@ func (r *postgresTimelineRepo) GetTimeline(ctx context.Context, req timeline.Get hotRanks = hotRanks[:req.Limit] lastPost := feedPosts[len(feedPosts)-1].Post lastHotRank := hotRanks[len(hotRanks)-1] - cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank) + cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank, queryTime) cursor = &cursorStr } diff --git a/tests/integration/feed_test.go b/tests/integration/feed_test.go index f41fdf7..8809dad 100644 --- a/tests/integration/feed_test.go +++ b/tests/integration/feed_test.go @@ -700,6 +700,116 @@ func TestGetCommunityFeed_HotCursorPrecision(t *testing.T) { t.Logf("SUCCESS: All posts with similar hot ranks preserved (precision bug fixed)") } +// TestGetCommunityFeed_HotCursorTimeDrift tests that hot sort pagination is stable across time drift. +// Regression test for a bug where posts would appear multiple times or be skipped when: +// 1. Time passes between page 1 and page 2 requests +// 2. Many posts have similar hot ranks +// +// Root cause: The cursor stored a hot_rank computed with NOW(), but the next query +// also used NOW() (which had advanced). This caused posts to drift across the cursor boundary. +// +// Fix: Store the cursor creation timestamp in the cursor and use it for subsequent comparisons, +// ensuring stable hot_rank computation across pagination requests. +func TestGetCommunityFeed_HotCursorTimeDrift(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + db := setupTestDB(t) + t.Cleanup(func() { _ = db.Close() }) + + // Setup services + feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") + communityRepo := postgres.NewCommunityRepository(db) + communityService := communities.NewCommunityService( + communityRepo, + "http://localhost:3001", + "did:web:test.coves.social", + "test.coves.social", + nil, + ) + feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) + handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) + + // Setup test data + ctx := context.Background() + testID := time.Now().UnixNano() + communityDID, err := createFeedTestCommunity(db, ctx, fmt.Sprintf("timedrift-%d", testID), fmt.Sprintf("timedrift-%d.test", testID)) + require.NoError(t, err) + + // Create 15 posts all with the SAME score and created at the SAME time + // This maximizes the chance of time drift causing duplicates: + // - All posts have nearly identical hot ranks + // - Any small change in NOW() could cause posts to swap order + baseTime := time.Now().Add(-1 * time.Hour) + var allPostURIs []string + for i := 0; i < 15; i++ { + // Add tiny offsets (1ms) to created_at for deterministic ordering + postURI := createTestPost(t, db, communityDID, fmt.Sprintf("did:plc:user%d", i), + fmt.Sprintf("Post %d", i), 10, baseTime.Add(time.Duration(i)*time.Millisecond)) + allPostURIs = append(allPostURIs, postURI) + } + + // Paginate through all posts with limit=5 + seenURIs := make(map[string]int) + var cursor *string + pageNum := 0 + + for { + pageNum++ + url := fmt.Sprintf("/xrpc/social.coves.communityFeed.getCommunity?community=%s&sort=hot&limit=5", communityDID) + if cursor != nil { + url += "&cursor=" + *cursor + } + + req := httptest.NewRequest(http.MethodGet, url, nil) + rec := httptest.NewRecorder() + handler.HandleGetCommunity(rec, req) + + require.Equal(t, http.StatusOK, rec.Code, "Page %d failed: %s", pageNum, rec.Body.String()) + + var page communityFeeds.FeedResponse + err = json.Unmarshal(rec.Body.Bytes(), &page) + require.NoError(t, err) + + if len(page.Feed) == 0 { + break + } + + for _, p := range page.Feed { + seenURIs[p.Post.URI]++ + if seenURIs[p.Post.URI] > 1 { + t.Errorf("DUPLICATE on page %d: %s (seen %d times)", pageNum, p.Post.URI, seenURIs[p.Post.URI]) + } + } + + cursor = page.Cursor + if cursor == nil { + break + } + + // Prevent infinite loops + if pageNum > 10 { + t.Fatal("Too many pages - possible infinite loop") + } + } + + // Verify we saw all posts exactly once + assert.Equal(t, 15, len(seenURIs), "Should see all 15 posts") + for uri, count := range seenURIs { + if count != 1 { + t.Errorf("Post %s seen %d times (expected 1)", uri, count) + } + } + + // Verify we saw all the posts we created + for _, uri := range allPostURIs { + assert.Contains(t, seenURIs, uri, "Missing post: %s", uri) + } + + t.Logf("SUCCESS: All 15 posts seen exactly once across %d pages (time drift bug fixed)", pageNum) +} + // TestGetCommunityFeed_BlobURLTransformation tests that blob refs are transformed to URLs func TestGetCommunityFeed_BlobURLTransformation(t *testing.T) { if testing.Short() {