diff --git a/internal/db/postgres/discover_repo.go b/internal/db/postgres/discover_repo.go index 8e7d3c1..edee2e9 100644 --- a/internal/db/postgres/discover_repo.go +++ b/internal/db/postgres/discover_repo.go @@ -48,7 +48,7 @@ func (r *postgresDiscoverRepo) GetDiscover(ctx context.Context, req discover.Get SELECT p.uri, p.cid, p.rkey, p.author_did, u.handle as author_handle, - p.community_did, c.name as community_name, c.avatar_cid as community_avatar, + p.community_did, c.handle as community_handle, c.name as community_name, c.avatar_cid as community_avatar, p.title, p.content, p.content_facets, p.embed, p.content_labels, p.created_at, p.edited_at, p.indexed_at, p.upvote_count, p.downvote_count, p.score, p.comment_count, @@ -59,7 +59,7 @@ func (r *postgresDiscoverRepo) GetDiscover(ctx context.Context, req discover.Get SELECT p.uri, p.cid, p.rkey, p.author_did, u.handle as author_handle, - p.community_did, c.name as community_name, c.avatar_cid as community_avatar, + p.community_did, c.handle as community_handle, c.name as community_name, c.avatar_cid as community_avatar, p.title, p.content, p.content_facets, p.embed, p.content_labels, p.created_at, p.edited_at, p.indexed_at, p.upvote_count, p.downvote_count, p.score, p.comment_count, diff --git a/internal/db/postgres/feed_repo.go b/internal/db/postgres/feed_repo.go index 7c3584d..b50e282 100644 --- a/internal/db/postgres/feed_repo.go +++ b/internal/db/postgres/feed_repo.go @@ -3,24 +3,18 @@ package postgres import ( "context" "database/sql" - "encoding/base64" - "encoding/json" "fmt" - "strconv" - "strings" - "time" "Coves/internal/core/communityFeeds" - "Coves/internal/core/posts" ) type postgresFeedRepo struct { - db *sql.DB + *feedRepoBase } // sortClauses maps sort types to safe SQL ORDER BY clauses // This whitelist prevents SQL injection via dynamic ORDER BY construction -var sortClauses = map[string]string{ +var communityFeedSortClauses = map[string]string{ "hot": `(p.score / POWER(EXTRACT(EPOCH FROM (NOW() - p.created_at))/3600 + 2, 1.5)) DESC, p.created_at DESC, p.uri DESC`, "top": `p.score DESC, p.created_at DESC, p.uri DESC`, "new": `p.created_at DESC, p.uri DESC`, @@ -30,21 +24,24 @@ var sortClauses = map[string]string{ // NOTE: Uses NOW() which means hot_rank changes over time - this is expected behavior // for hot sorting (posts naturally age out). Slight time drift between cursor creation // and usage may cause minor reordering but won't drop posts entirely (unlike using raw score). -const hotRankExpression = `(p.score / POWER(EXTRACT(EPOCH FROM (NOW() - p.created_at))/3600 + 2, 1.5))` +const communityFeedHotRankExpression = `(p.score / POWER(EXTRACT(EPOCH FROM (NOW() - p.created_at))/3600 + 2, 1.5))` // NewCommunityFeedRepository creates a new PostgreSQL feed repository -func NewCommunityFeedRepository(db *sql.DB) communityFeeds.Repository { - return &postgresFeedRepo{db: db} +func NewCommunityFeedRepository(db *sql.DB, cursorSecret string) communityFeeds.Repository { + return &postgresFeedRepo{ + feedRepoBase: newFeedRepoBase(db, communityFeedHotRankExpression, communityFeedSortClauses, cursorSecret), + } } // 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) { // Build ORDER BY clause based on sort type - orderBy, timeFilter := r.buildSortClause(req.Sort, req.Timeframe) + orderBy, timeFilter := r.feedRepoBase.buildSortClause(req.Sort, req.Timeframe) // Build cursor filter for pagination - cursorFilter, cursorValues, err := r.parseCursor(req.Cursor, req.Sort) + // Community feed uses $3+ for cursor params (after $1=community and $2=limit) + cursorFilter, cursorValues, err := r.feedRepoBase.parseCursor(req.Cursor, req.Sort, 3) if err != nil { return nil, nil, communityFeeds.ErrInvalidCursor } @@ -57,18 +54,18 @@ func (r *postgresFeedRepo) GetCommunityFeed(ctx context.Context, req communityFe SELECT p.uri, p.cid, p.rkey, p.author_did, u.handle as author_handle, - p.community_did, c.name as community_name, c.avatar_cid as community_avatar, + p.community_did, c.handle as community_handle, c.name as community_name, c.avatar_cid as community_avatar, p.title, p.content, p.content_facets, p.embed, p.content_labels, p.created_at, p.edited_at, p.indexed_at, p.upvote_count, p.downvote_count, p.score, p.comment_count, %s as hot_rank - FROM posts p`, hotRankExpression) + FROM posts p`, communityFeedHotRankExpression) } else { selectClause = ` SELECT p.uri, p.cid, p.rkey, p.author_did, u.handle as author_handle, - p.community_did, c.name as community_name, c.avatar_cid as community_avatar, + p.community_did, c.handle as community_handle, c.name as community_name, c.avatar_cid as community_avatar, p.title, p.content, p.content_facets, p.embed, p.content_labels, p.created_at, p.edited_at, p.indexed_at, p.upvote_count, p.downvote_count, p.score, p.comment_count, @@ -108,11 +105,11 @@ func (r *postgresFeedRepo) GetCommunityFeed(ctx context.Context, req communityFe var feedPosts []*communityFeeds.FeedViewPost var hotRanks []float64 // Store hot ranks for cursor building for rows.Next() { - feedPost, hotRank, err := r.scanFeedViewPost(rows) + postView, hotRank, err := r.feedRepoBase.scanFeedPost(rows) if err != nil { return nil, nil, fmt.Errorf("failed to scan feed post: %w", err) } - feedPosts = append(feedPosts, feedPost) + feedPosts = append(feedPosts, &communityFeeds.FeedViewPost{Post: postView}) hotRanks = append(hotRanks, hotRank) } @@ -127,322 +124,9 @@ 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.buildCursor(lastPost, req.Sort, lastHotRank) + cursorStr := r.feedRepoBase.buildCursor(lastPost, req.Sort, lastHotRank) cursor = &cursorStr } return feedPosts, cursor, nil } - -// buildSortClause returns the ORDER BY SQL and optional time filter -func (r *postgresFeedRepo) buildSortClause(sort, timeframe string) (string, string) { - // Use whitelist map for ORDER BY clause (defense-in-depth against SQL injection) - orderBy := sortClauses[sort] - if orderBy == "" { - orderBy = sortClauses["hot"] // safe default - } - - // Add time filter for "top" sort - var timeFilter string - if sort == "top" { - timeFilter = r.buildTimeFilter(timeframe) - } - - return orderBy, timeFilter -} - -// buildTimeFilter returns SQL filter for timeframe -func (r *postgresFeedRepo) buildTimeFilter(timeframe string) string { - if timeframe == "" || timeframe == "all" { - return "" - } - - var interval string - switch timeframe { - case "hour": - interval = "1 hour" - case "day": - interval = "1 day" - case "week": - interval = "1 week" - case "month": - interval = "1 month" - case "year": - interval = "1 year" - default: - return "" - } - - return fmt.Sprintf("AND p.created_at > NOW() - INTERVAL '%s'", interval) -} - -// parseCursor decodes pagination cursor -func (r *postgresFeedRepo) parseCursor(cursor *string, sort string) (string, []interface{}, error) { - if cursor == nil || *cursor == "" { - return "", nil, nil - } - - // Decode base64 cursor - decoded, err := base64.StdEncoding.DecodeString(*cursor) - if err != nil { - return "", nil, fmt.Errorf("invalid cursor encoding") - } - - // Parse cursor based on sort type using :: delimiter (Bluesky convention) - parts := strings.Split(string(decoded), "::") - - switch sort { - case "new": - // Cursor format: timestamp::uri - if len(parts) != 2 { - return "", nil, fmt.Errorf("invalid cursor format") - } - - createdAt := parts[0] - uri := parts[1] - - // Validate timestamp format - if _, err := time.Parse(time.RFC3339Nano, createdAt); err != nil { - return "", nil, fmt.Errorf("invalid cursor timestamp") - } - - // Validate URI format (must be AT-URI) - if !strings.HasPrefix(uri, "at://") { - return "", nil, fmt.Errorf("invalid cursor URI") - } - - filter := `AND (p.created_at < $3 OR (p.created_at = $3 AND p.uri < $4))` - return filter, []interface{}{createdAt, uri}, nil - - case "top": - // Cursor format: score::timestamp::uri - if len(parts) != 3 { - return "", nil, fmt.Errorf("invalid cursor format for %s sort", sort) - } - - scoreStr := parts[0] - createdAt := parts[1] - uri := parts[2] - - // Validate score is numeric - score := 0 - if _, err := fmt.Sscanf(scoreStr, "%d", &score); err != nil { - return "", nil, fmt.Errorf("invalid cursor score") - } - - // Validate timestamp format - if _, err := time.Parse(time.RFC3339Nano, createdAt); err != nil { - return "", nil, fmt.Errorf("invalid cursor timestamp") - } - - // Validate URI format (must be AT-URI) - if !strings.HasPrefix(uri, "at://") { - return "", nil, fmt.Errorf("invalid cursor URI") - } - - filter := `AND (p.score < $3 OR (p.score = $3 AND p.created_at < $4) OR (p.score = $3 AND p.created_at = $4 AND p.uri < $5))` - 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(parts) != 3 { - return "", nil, fmt.Errorf("invalid cursor format for hot sort") - } - - hotRankStr := parts[0] - createdAt := parts[1] - uri := parts[2] - - // Validate hot_rank is numeric (float) - hotRank := 0.0 - if _, err := fmt.Sscanf(hotRankStr, "%f", &hotRank); err != nil { - 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 URI format (must be AT-URI) - if !strings.HasPrefix(uri, "at://") { - return "", nil, fmt.Errorf("invalid cursor URI") - } - - // CRITICAL: Compare against the computed hot_rank expression, not p.score - // This prevents dropping posts with higher raw scores but lower hot ranks - // - // NOTE: We exclude the exact cursor post by URI to handle time drift in hot_rank - // (hot_rank changes with NOW(), so the same post may have different ranks over time) - filter := fmt.Sprintf(`AND ((%s < $3 OR (%s = $3 AND p.created_at < $4) OR (%s = $3 AND p.created_at = $4 AND p.uri < $5)) AND p.uri != $6)`, - hotRankExpression, hotRankExpression, hotRankExpression) - return filter, []interface{}{hotRank, createdAt, uri, uri}, nil - - default: - return "", nil, nil - } -} - -// buildCursor creates pagination cursor from last post -func (r *postgresFeedRepo) buildCursor(post *posts.PostView, sort string, hotRank float64) string { - var cursorStr string - // Use :: as delimiter following Bluesky convention - // Safe because :: doesn't appear in ISO timestamps or AT-URIs - const delimiter = "::" - - switch sort { - case "new": - // Format: timestamp::uri (following Bluesky pattern) - cursorStr = fmt.Sprintf("%s%s%s", post.CreatedAt.Format(time.RFC3339Nano), delimiter, post.URI) - - case "top": - // Format: score::timestamp::uri - score := 0 - if post.Stats != nil { - score = post.Stats.Score - } - cursorStr = 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 to prevent pagination bugs - // Using 'g' format with -1 precision gives us full float64 precision without trailing zeros - // This prevents posts being dropped when hot ranks differ by <1e-6 - hotRankStr := strconv.FormatFloat(hotRank, 'g', -1, 64) - cursorStr = fmt.Sprintf("%s%s%s%s%s", hotRankStr, delimiter, post.CreatedAt.Format(time.RFC3339Nano), delimiter, post.URI) - - default: - cursorStr = post.URI - } - - return base64.StdEncoding.EncodeToString([]byte(cursorStr)) -} - -// scanFeedViewPost scans a row into FeedViewPost -// Alpha: No viewer state - basic community feed only -func (r *postgresFeedRepo) scanFeedViewPost(rows *sql.Rows) (*communityFeeds.FeedViewPost, float64, error) { - var ( - postView posts.PostView - authorView posts.AuthorView - communityRef posts.CommunityRef - title, content sql.NullString - facets, embed sql.NullString - labelsJSON sql.NullString - editedAt sql.NullTime - communityAvatar sql.NullString - hotRank sql.NullFloat64 - ) - - err := rows.Scan( - &postView.URI, &postView.CID, &postView.RKey, - &authorView.DID, &authorView.Handle, - &communityRef.DID, &communityRef.Name, &communityAvatar, - &title, &content, &facets, &embed, &labelsJSON, - &postView.CreatedAt, &editedAt, &postView.IndexedAt, - &postView.UpvoteCount, &postView.DownvoteCount, &postView.Score, &postView.CommentCount, - &hotRank, - ) - if err != nil { - return nil, 0, err - } - - // Build author view (no display_name or avatar in users table yet) - postView.Author = &authorView - - // Build community ref - communityRef.Avatar = nullStringPtr(communityAvatar) - postView.Community = &communityRef - - // Set optional fields - postView.Title = nullStringPtr(title) - postView.Text = nullStringPtr(content) - - // Parse facets JSON - if facets.Valid { - var facetArray []interface{} - if err := json.Unmarshal([]byte(facets.String), &facetArray); err == nil { - postView.TextFacets = facetArray - } - } - - // Parse embed JSON - if embed.Valid { - var embedData interface{} - if err := json.Unmarshal([]byte(embed.String), &embedData); err == nil { - postView.Embed = embedData - } - } - - // Build stats - postView.Stats = &posts.PostStats{ - Upvotes: postView.UpvoteCount, - Downvotes: postView.DownvoteCount, - Score: postView.Score, - CommentCount: postView.CommentCount, - } - - // Alpha: No viewer state for basic feed - // TODO(feed-generator): Implement viewer state (saved, voted, blocked) in feed generator skeleton - - // Build the record (required by lexicon - social.coves.community.post structure) - record := map[string]interface{}{ - "$type": "social.coves.community.post", - "community": communityRef.DID, - "author": authorView.DID, - "createdAt": postView.CreatedAt.Format(time.RFC3339), - } - - // Add optional fields to record if present - if title.Valid { - record["title"] = title.String - } - if content.Valid { - record["content"] = content.String - } - if facets.Valid { - var facetArray []interface{} - if err := json.Unmarshal([]byte(facets.String), &facetArray); err == nil { - record["facets"] = facetArray - } - } - if embed.Valid { - var embedData interface{} - if err := json.Unmarshal([]byte(embed.String), &embedData); err == nil { - record["embed"] = embedData - } - } - if labelsJSON.Valid { - // Labels are stored as JSONB containing full com.atproto.label.defs#selfLabels structure - // Deserialize and include in record - var selfLabels posts.SelfLabels - if err := json.Unmarshal([]byte(labelsJSON.String), &selfLabels); err == nil { - record["labels"] = selfLabels - } - } - - postView.Record = record - - // Wrap in FeedViewPost - feedPost := &communityFeeds.FeedViewPost{ - Post: &postView, - // Reason: nil, // TODO(feed-generator): Implement pinned posts - // Reply: nil, // TODO(feed-generator): Implement reply context - } - - // Return the computed hot_rank (0.0 if NULL for non-hot sorts) - hotRankValue := 0.0 - if hotRank.Valid { - hotRankValue = hotRank.Float64 - } - - return feedPost, hotRankValue, nil -} - -// Helper function to convert sql.NullString to *string -func nullStringPtr(ns sql.NullString) *string { - if !ns.Valid { - return nil - } - return &ns.String -} diff --git a/internal/db/postgres/feed_repo_base.go b/internal/db/postgres/feed_repo_base.go index e2923e3..aec6e23 100644 --- a/internal/db/postgres/feed_repo_base.go +++ b/internal/db/postgres/feed_repo_base.go @@ -284,6 +284,7 @@ func (r *feedRepoBase) scanFeedPost(rows *sql.Rows) (*posts.PostView, float64, e facets, embed sql.NullString labelsJSON sql.NullString editedAt sql.NullTime + communityHandle sql.NullString communityAvatar sql.NullString hotRank sql.NullFloat64 ) @@ -291,7 +292,7 @@ func (r *feedRepoBase) scanFeedPost(rows *sql.Rows) (*posts.PostView, float64, e err := rows.Scan( &postView.URI, &postView.CID, &postView.RKey, &authorView.DID, &authorView.Handle, - &communityRef.DID, &communityRef.Name, &communityAvatar, + &communityRef.DID, &communityHandle, &communityRef.Name, &communityAvatar, &title, &content, &facets, &embed, &labelsJSON, &postView.CreatedAt, &editedAt, &postView.IndexedAt, &postView.UpvoteCount, &postView.DownvoteCount, &postView.Score, &postView.CommentCount, @@ -305,6 +306,9 @@ func (r *feedRepoBase) scanFeedPost(rows *sql.Rows) (*posts.PostView, float64, e postView.Author = &authorView // Build community ref + if communityHandle.Valid { + communityRef.Handle = communityHandle.String + } communityRef.Avatar = nullStringPtr(communityAvatar) postView.Community = &communityRef @@ -382,3 +386,12 @@ func (r *feedRepoBase) scanFeedPost(rows *sql.Rows) (*posts.PostView, float64, e return &postView, hotRankValue, nil } + +// nullStringPtr converts sql.NullString to *string +// Helper function used by feed scanning logic across all feed types +func nullStringPtr(ns sql.NullString) *string { + if !ns.Valid { + return nil + } + return &ns.String +} diff --git a/internal/db/postgres/timeline_repo.go b/internal/db/postgres/timeline_repo.go index 28b797e..f2618cf 100644 --- a/internal/db/postgres/timeline_repo.go +++ b/internal/db/postgres/timeline_repo.go @@ -52,7 +52,7 @@ func (r *postgresTimelineRepo) GetTimeline(ctx context.Context, req timeline.Get SELECT p.uri, p.cid, p.rkey, p.author_did, u.handle as author_handle, - p.community_did, c.name as community_name, c.avatar_cid as community_avatar, + p.community_did, c.handle as community_handle, c.name as community_name, c.avatar_cid as community_avatar, p.title, p.content, p.content_facets, p.embed, p.content_labels, p.created_at, p.edited_at, p.indexed_at, p.upvote_count, p.downvote_count, p.score, p.comment_count, @@ -63,7 +63,7 @@ func (r *postgresTimelineRepo) GetTimeline(ctx context.Context, req timeline.Get SELECT p.uri, p.cid, p.rkey, p.author_did, u.handle as author_handle, - p.community_did, c.name as community_name, c.avatar_cid as community_avatar, + p.community_did, c.handle as community_handle, c.name as community_name, c.avatar_cid as community_avatar, p.title, p.content, p.content_facets, p.embed, p.content_labels, p.created_at, p.edited_at, p.indexed_at, p.upvote_count, p.downvote_count, p.score, p.comment_count,