fix(backend): harden security and speed up feeds and article API

Security:
- WebSocket events are routed to their owner only (no cross-user leak);
  hub close is idempotent (fixes double-close panic), adds ping/pong and
  write deadlines.
- Session tokens stored as SHA-256 (migration 008 keeps sessions valid);
  single-query auth middleware puts the user in the request context.
- Client IP only trusts X-Forwarded-For from TRUSTED_PROXIES; rate limiter
  map is bounded; per-user limit on AI summaries.
- Argon2id at OWASP minimum with a concurrency cap; constant-time login
  for unknown emails; atomic first-admin bootstrap; REGISTRATION_ENABLED.
- CSP/HSTS/COOP headers, same-origin guard on mutations, body size limits,
  wider SSRF denylist, bounded feed/page/AI response reads, generic errors.
- Upgrade chi, pgx, x/net, x/text, x/crypto (known CVEs); commit go.sum.

Performance:
- List endpoints return a plain-text excerpt and reading time instead of
  full HTML; content is sanitized once at ingest (legacy rows backfilled).
- Keyset pagination on (sort_at, id) with matching partial indexes;
  redundant indexes dropped (migration 007).
- Fetcher: bounded worker pool, conditional GET (ETag/Last-Modified),
  exponential backoff, dedupe before insert, column-safe truncation,
  retention-aware ingest, per-user refresh coalescing.
- Read/favorite/read-all are single ownership-scoped statements.
- gzip compression, immutable caching for hashed assets, path-safe SPA
  handler, server timeouts; expired sessions purged.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
Antigravity AgentandClaude Opus 5.5 committed 2026-10-09 07:34:08 +02:00
1 parent d037e2be34
commit 03e57e4308
40 files changed
+2231 -1898

No files matched your search

+279 -462
View File
@@ -4,12 +4,15 @@ import (
"context"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/michael/flowreader/internal/domain"
"github.com/michael/flowreader/internal/utils"
)
// ArticleRepository implements domain.ArticleRepository using PostgreSQL.
@@ -22,500 +25,265 @@ func NewArticleRepository(pool *pgxpool.Pool) *ArticleRepository {
return &ArticleRepository{pool: pool}
}
// Create inserts a new article into the database.
func (r *ArticleRepository) Create(article *domain.Article) error {
ctx := context.Background()
// listColumns are the light-weight columns sent for article lists: no full
// HTML content, only a plain-text excerpt and a word count.
const listColumns = `
a.id, a.feed_id, a.title, a.url, a.excerpt, a.ai_summary, a.author, a.image_url,
a.published_at, a.sort_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
a.word_count, f.title`
query := `
INSERT INTO articles (id, feed_id, guid, title, url, content, summary, ai_summary, author, image_url, published_at, created_at)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12)
ON CONFLICT (feed_id, guid) DO NOTHING
`
_, err := r.pool.Exec(ctx, query,
article.ID,
article.FeedID,
article.GUID,
article.Title,
nullString(article.URL),
nullString(article.Content),
nullString(article.Summary),
nullString(article.AISummary),
nullString(article.Author),
nullString(article.ImageURL),
article.PublishedAt,
article.CreatedAt,
)
if err != nil {
return fmt.Errorf("creating article: %w", err)
// ExistingGUIDs returns which of the given GUIDs are already stored for a feed.
func (r *ArticleRepository) ExistingGUIDs(ctx context.Context, feedID uuid.UUID, guids []string) (map[string]struct{}, error) {
out := make(map[string]struct{}, len(guids))
if len(guids) == 0 {
return out, nil
}
return nil
rows, err := r.pool.Query(ctx, `SELECT guid FROM articles WHERE feed_id = $1 AND guid = ANY($2)`, feedID, guids)
if err != nil {
return nil, fmt.Errorf("querying existing guids: %w", err)
}
defer rows.Close()
for rows.Next() {
var g string
if err := rows.Scan(&g); err != nil {
return nil, fmt.Errorf("scanning guid: %w", err)
}
out[g] = struct{}{}
}
return out, rows.Err()
}
// CreateBatch inserts multiple articles into the database.
func (r *ArticleRepository) CreateBatch(articles []*domain.Article) error {
ctx := context.Background()
// InsertNew inserts articles, ignoring GUIDs already present, and returns
// the number of rows actually inserted.
func (r *ArticleRepository) InsertNew(ctx context.Context, feedID uuid.UUID, articles []*domain.Article) (int, error) {
if len(articles) == 0 {
return 0, nil
}
const query = `
INSERT INTO articles (id, feed_id, guid, title, url, content, summary, excerpt, word_count,
author, image_url, published_at, created_at)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13)
ON CONFLICT (feed_id, guid) DO NOTHING`
batch := &pgx.Batch{}
query := `
INSERT INTO articles (id, feed_id, guid, title, url, content, summary, ai_summary, author, image_url, published_at, created_at)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12)
ON CONFLICT (feed_id, guid) DO NOTHING
`
for _, article := range articles {
for _, a := range articles {
batch.Queue(query,
article.ID,
article.FeedID,
article.GUID,
article.Title,
nullString(article.URL),
nullString(article.Content),
nullString(article.Summary),
nullString(article.AISummary),
nullString(article.Author),
nullString(article.ImageURL),
article.PublishedAt,
article.CreatedAt,
a.ID, feedID, a.GUID, a.Title,
nullString(a.URL), nullString(a.Content), nullString(a.Summary),
a.Excerpt, a.WordCount,
nullString(a.Author), nullString(a.ImageURL),
a.PublishedAt, a.CreatedAt,
)
}
results := r.pool.SendBatch(ctx, batch)
defer results.Close()
inserted := 0
for range articles {
if _, err := results.Exec(); err != nil {
return fmt.Errorf("batch insert: %w", err)
}
}
return nil
}
// GetByID retrieves an article by its ID.
func (r *ArticleRepository) GetByID(id uuid.UUID) (*domain.Article, error) {
ctx := context.Background()
query := `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE a.id = $1
`
article, err := r.scanArticle(r.pool.QueryRow(ctx, query, id))
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, fmt.Errorf("getting article by ID: %w", err)
}
return article, nil
}
// GetByFeedID retrieves articles for a specific feed.
func (r *ArticleRepository) GetByFeedID(feedID uuid.UUID, limit, offset int) ([]*domain.Article, error) {
ctx := context.Background()
query := `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE a.feed_id = $1
ORDER BY a.published_at DESC NULLS LAST, a.created_at DESC
LIMIT $2 OFFSET $3
`
rows, err := r.pool.Query(ctx, query, feedID, limit, offset)
if err != nil {
return nil, fmt.Errorf("querying articles: %w", err)
}
defer rows.Close()
return r.scanArticles(rows)
}
// GetByUserID retrieves articles for all user's feeds.
func (r *ArticleRepository) GetByUserID(userID uuid.UUID, limit, offset int, unreadOnly bool) ([]*domain.Article, error) {
ctx := context.Background()
var query string
if unreadOnly {
query = `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE f.user_id = $1 AND a.is_read = false
ORDER BY a.published_at DESC NULLS LAST, a.created_at DESC
LIMIT $2 OFFSET $3
`
} else {
query = `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE f.user_id = $1
ORDER BY a.published_at DESC NULLS LAST, a.created_at DESC
LIMIT $2 OFFSET $3
`
}
rows, err := r.pool.Query(ctx, query, userID, limit, offset)
if err != nil {
return nil, fmt.Errorf("querying articles: %w", err)
}
defer rows.Close()
return r.scanArticles(rows)
}
// GetByGUID retrieves an article by its GUID within a feed.
func (r *ArticleRepository) GetByGUID(feedID uuid.UUID, guid string) (*domain.Article, error) {
ctx := context.Background()
query := `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE a.feed_id = $1 AND a.guid = $2
`
article, err := r.scanArticle(r.pool.QueryRow(ctx, query, feedID, guid))
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, fmt.Errorf("getting article by GUID: %w", err)
}
return article, nil
}
// MarkAsRead marks an article as read.
func (r *ArticleRepository) MarkAsRead(id uuid.UUID) error {
ctx := context.Background()
query := `UPDATE articles SET is_read = true, read_at = $2 WHERE id = $1`
_, err := r.pool.Exec(ctx, query, id, time.Now())
if err != nil {
return fmt.Errorf("marking as read: %w", err)
}
return nil
}
// MarkAsUnread marks an article as unread.
func (r *ArticleRepository) MarkAsUnread(id uuid.UUID) error {
ctx := context.Background()
query := `UPDATE articles SET is_read = false, read_at = NULL WHERE id = $1`
_, err := r.pool.Exec(ctx, query, id)
if err != nil {
return fmt.Errorf("marking as unread: %w", err)
}
return nil
}
// MarkAllAsRead marks all articles in a feed as read.
func (r *ArticleRepository) MarkAllAsRead(feedID uuid.UUID) error {
ctx := context.Background()
query := `UPDATE articles SET is_read = true, read_at = $2 WHERE feed_id = $1 AND is_read = false`
_, err := r.pool.Exec(ctx, query, feedID, time.Now())
if err != nil {
return fmt.Errorf("marking all as read: %w", err)
}
return nil
}
// MarkAllAsReadGlobal marks all articles for a user as read.
func (r *ArticleRepository) MarkAllAsReadGlobal(userID uuid.UUID) error {
ctx := context.Background()
query := `
UPDATE articles
SET is_read = true, read_at = NOW()
WHERE feed_id IN (SELECT id FROM feeds WHERE user_id = $1) AND is_read = false
`
_, err := r.pool.Exec(ctx, query, userID)
if err != nil {
return fmt.Errorf("marking all articles as read globally: %w", err)
}
return nil
}
// ToggleFavorite toggles the favorite status of an article.
func (r *ArticleRepository) ToggleFavorite(id uuid.UUID) error {
ctx := context.Background()
query := `UPDATE articles SET is_favorite = NOT is_favorite WHERE id = $1`
_, err := r.pool.Exec(ctx, query, id)
if err != nil {
return fmt.Errorf("toggling favorite: %w", err)
}
return nil
}
// GetFavorites retrieves favorited articles for a user.
func (r *ArticleRepository) GetFavorites(userID uuid.UUID, limit, offset int) ([]*domain.Article, error) {
ctx := context.Background()
query := `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE f.user_id = $1 AND a.is_favorite = true
ORDER BY a.published_at DESC NULLS LAST, a.created_at DESC
LIMIT $2 OFFSET $3
`
rows, err := r.pool.Query(ctx, query, userID, limit, offset)
if err != nil {
return nil, fmt.Errorf("querying favorites: %w", err)
}
defer rows.Close()
return r.scanArticles(rows)
}
// CountUnread counts unread articles for a feed.
func (r *ArticleRepository) CountUnread(feedID uuid.UUID) (int, error) {
ctx := context.Background()
query := `SELECT COUNT(*) FROM articles WHERE feed_id = $1 AND is_read = false`
var count int
err := r.pool.QueryRow(ctx, query, feedID).Scan(&count)
if err != nil {
return 0, fmt.Errorf("counting unread: %w", err)
}
return count, nil
}
// scanArticle scans a single article row.
func (r *ArticleRepository) scanArticle(row pgx.Row) (*domain.Article, error) {
var article domain.Article
var url, content, summary, aiSummary, author, imageURL, feedTitle *string
var publishedAt, readAt *time.Time
err := row.Scan(
&article.ID,
&article.FeedID,
&article.GUID,
&article.Title,
&url,
&content,
&summary,
&aiSummary,
&author,
&imageURL,
&publishedAt,
&article.IsRead,
&article.IsFavorite,
&readAt,
&article.CreatedAt,
&feedTitle,
)
if err != nil {
return nil, err
}
if url != nil {
article.URL = *url
}
if content != nil {
article.Content = *content
}
if summary != nil {
article.Summary = *summary
}
if aiSummary != nil {
article.AISummary = *aiSummary
}
if author != nil {
article.Author = *author
}
if imageURL != nil {
article.ImageURL = *imageURL
}
if feedTitle != nil {
article.FeedTitle = *feedTitle
}
article.PublishedAt = publishedAt
article.ReadAt = readAt
return &article, nil
}
// scanArticles scans multiple article rows.
func (r *ArticleRepository) scanArticles(rows pgx.Rows) ([]*domain.Article, error) {
var articles []*domain.Article
for rows.Next() {
var article domain.Article
var url, content, summary, aiSummary, author, imageURL, feedTitle *string
var publishedAt, readAt *time.Time
err := rows.Scan(
&article.ID,
&article.FeedID,
&article.GUID,
&article.Title,
&url,
&content,
&summary,
&aiSummary,
&author,
&imageURL,
&publishedAt,
&article.IsRead,
&article.IsFavorite,
&readAt,
&article.CreatedAt,
&feedTitle,
)
tag, err := results.Exec()
if err != nil {
return nil, fmt.Errorf("scanning article: %w", err)
return inserted, fmt.Errorf("batch insert: %w", err)
}
if url != nil {
article.URL = *url
}
if content != nil {
article.Content = *content
}
if summary != nil {
article.Summary = *summary
}
if aiSummary != nil {
article.AISummary = *aiSummary
}
if author != nil {
article.Author = *author
}
if imageURL != nil {
article.ImageURL = *imageURL
}
if feedTitle != nil {
article.FeedTitle = *feedTitle
}
article.PublishedAt = publishedAt
article.ReadAt = readAt
articles = append(articles, &article)
inserted += int(tag.RowsAffected())
}
return articles, nil
return inserted, nil
}
// Search performs a full-text search on articles for a specific user.
func (r *ArticleRepository) Search(userID uuid.UUID, query string, limit, offset int) ([]*domain.Article, error) {
ctx := context.Background()
// Use plainto_tsquery or websearch_to_tsquery for natural language search
sql := `
SELECT a.id, a.feed_id, a.guid, a.title, a.url, a.content, a.summary, a.ai_summary, a.author,
a.image_url, a.published_at, a.is_read, a.is_favorite, a.read_at, a.created_at,
f.title as feed_title,
ts_rank_cd(a.tsv, websearch_to_tsquery('french', $2)) as rank
// GetForUser retrieves a full article (including content) owned by the user.
// Returns nil, nil when it doesn't exist or belongs to someone else.
func (r *ArticleRepository) GetForUser(ctx context.Context, id, userID uuid.UUID) (*domain.Article, error) {
const query = `
SELECT a.id, a.feed_id, a.title, a.url, a.content, a.summary, a.excerpt, a.ai_summary,
a.author, a.image_url, a.published_at, a.sort_at, a.is_read, a.is_favorite,
a.read_at, a.created_at, a.word_count, f.title
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE f.user_id = $1 AND a.tsv @@ websearch_to_tsquery('french', $2)
ORDER BY rank DESC, a.published_at DESC
LIMIT $3 OFFSET $4
`
WHERE a.id = $1 AND f.user_id = $2`
var a domain.Article
var url, content, summary, excerpt, aiSummary, author, imageURL, feedTitle *string
var wordCount *int
err := r.pool.QueryRow(ctx, query, id, userID).Scan(
&a.ID, &a.FeedID, &a.Title, &url, &content, &summary, &excerpt, &aiSummary,
&author, &imageURL, &a.PublishedAt, &a.SortAt, &a.IsRead, &a.IsFavorite,
&a.ReadAt, &a.CreatedAt, &wordCount, &feedTitle,
)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, fmt.Errorf("getting article: %w", err)
}
a.URL = deref(url)
a.Content = deref(content)
a.Summary = deref(summary)
a.Excerpt = deref(excerpt)
a.AISummary = deref(aiSummary)
a.Author = deref(author)
a.ImageURL = deref(imageURL)
a.FeedTitle = deref(feedTitle)
if wordCount != nil {
a.WordCount = *wordCount
} else {
a.WordCount = utils.WordCount(utils.PlainText(firstNonEmpty(a.Content, a.Summary)))
}
a.ReadingTime = utils.ReadingMinutes(a.WordCount)
return &a, nil
}
// List returns a page of the user's articles, newest first, using keyset
// pagination on (sort_at, id).
func (r *ArticleRepository) List(ctx context.Context, f domain.ArticleFilter) ([]*domain.Article, error) {
var sb strings.Builder
args := []any{f.UserID}
sb.WriteString(`SELECT ` + listColumns + `
FROM articles a
JOIN feeds f ON f.id = a.feed_id
WHERE f.user_id = $1`)
if f.FeedID != nil {
args = append(args, *f.FeedID)
sb.WriteString(` AND a.feed_id = $` + strconv.Itoa(len(args)))
}
if f.UnreadOnly {
sb.WriteString(` AND NOT a.is_read`)
}
if f.FavoritesOnly {
sb.WriteString(` AND a.is_favorite`)
}
if f.Cursor != nil {
args = append(args, f.Cursor.SortAt, f.Cursor.ID)
sb.WriteString(` AND (a.sort_at, a.id) < ($` + strconv.Itoa(len(args)-1) + `, $` + strconv.Itoa(len(args)) + `)`)
}
args = append(args, f.Limit)
sb.WriteString(` ORDER BY a.sort_at DESC, a.id DESC LIMIT $` + strconv.Itoa(len(args)))
rows, err := r.pool.Query(ctx, sb.String(), args...)
if err != nil {
return nil, fmt.Errorf("listing articles: %w", err)
}
defer rows.Close()
return scanListRows(rows, false)
}
// Search performs a full-text search on the user's articles.
func (r *ArticleRepository) Search(ctx context.Context, userID uuid.UUID, query string, limit, offset int) ([]*domain.Article, error) {
sql := `
SELECT ` + listColumns + `, ts_rank_cd(a.tsv, q) AS rank
FROM articles a
JOIN feeds f ON f.id = a.feed_id
CROSS JOIN websearch_to_tsquery('french', $2) q
WHERE f.user_id = $1 AND a.tsv @@ q
ORDER BY rank DESC, a.sort_at DESC
LIMIT $3 OFFSET $4`
rows, err := r.pool.Query(ctx, sql, userID, query, limit, offset)
if err != nil {
return nil, fmt.Errorf("searching articles: %w", err)
}
defer rows.Close()
return r.scanArticlesWithRank(rows)
return scanListRows(rows, true)
}
// scanArticlesWithRank scans multiple article rows with their rank.
func (r *ArticleRepository) scanArticlesWithRank(rows pgx.Rows) ([]*domain.Article, error) {
var articles []*domain.Article
func scanListRows(rows pgx.Rows, withRank bool) ([]*domain.Article, error) {
articles := make([]*domain.Article, 0, 32)
for rows.Next() {
var article domain.Article
var url, content, summary, aiSummary, author, imageURL, feedTitle *string
var publishedAt, readAt *time.Time
var a domain.Article
var url, excerpt, aiSummary, author, imageURL, feedTitle *string
var wordCount *int
var rank float32
err := rows.Scan(
&article.ID,
&article.FeedID,
&article.GUID,
&article.Title,
&url,
&content,
&summary,
&aiSummary,
&author,
&imageURL,
&publishedAt,
&article.IsRead,
&article.IsFavorite,
&readAt,
&article.CreatedAt,
&feedTitle,
&rank,
)
if err != nil {
return nil, fmt.Errorf("scanning article with rank: %w", err)
dest := []any{
&a.ID, &a.FeedID, &a.Title, &url, &excerpt, &aiSummary, &author, &imageURL,
&a.PublishedAt, &a.SortAt, &a.IsRead, &a.IsFavorite, &a.ReadAt, &a.CreatedAt,
&wordCount, &feedTitle,
}
if url != nil {
article.URL = *url
if withRank {
dest = append(dest, &rank)
}
if content != nil {
article.Content = *content
if err := rows.Scan(dest...); err != nil {
return nil, fmt.Errorf("scanning article: %w", err)
}
if summary != nil {
article.Summary = *summary
a.URL = deref(url)
a.Excerpt = deref(excerpt)
a.AISummary = deref(aiSummary)
a.Author = deref(author)
a.ImageURL = deref(imageURL)
a.FeedTitle = deref(feedTitle)
if wordCount != nil {
a.WordCount = *wordCount
}
if aiSummary != nil {
article.AISummary = *aiSummary
}
if author != nil {
article.Author = *author
}
if imageURL != nil {
article.ImageURL = *imageURL
}
if feedTitle != nil {
article.FeedTitle = *feedTitle
}
article.PublishedAt = publishedAt
article.ReadAt = readAt
articles = append(articles, &article)
a.ReadingTime = utils.ReadingMinutes(a.WordCount)
articles = append(articles, &a)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterating articles: %w", err)
}
return articles, nil
}
// UpdateAISummary updates the AI-generated summary of an article.
func (r *ArticleRepository) UpdateAISummary(id uuid.UUID, summary string) error {
ctx := context.Background()
query := `UPDATE articles SET ai_summary = $2 WHERE id = $1`
_, err := r.pool.Exec(ctx, query, id, summary)
// SetRead marks an owned article as read or unread in a single statement.
func (r *ArticleRepository) SetRead(ctx context.Context, id, userID uuid.UUID, read bool) (bool, error) {
const query = `
UPDATE articles a
SET is_read = $3, read_at = CASE WHEN $3 THEN NOW() ELSE NULL END
FROM feeds f
WHERE a.id = $1 AND f.id = a.feed_id AND f.user_id = $2`
tag, err := r.pool.Exec(ctx, query, id, userID, read)
if err != nil {
return false, fmt.Errorf("setting read state: %w", err)
}
return tag.RowsAffected() > 0, nil
}
// ToggleFavorite flips the favorite flag of an owned article atomically.
func (r *ArticleRepository) ToggleFavorite(ctx context.Context, id, userID uuid.UUID) (bool, bool, error) {
const query = `
UPDATE articles a
SET is_favorite = NOT a.is_favorite
FROM feeds f
WHERE a.id = $1 AND f.id = a.feed_id AND f.user_id = $2
RETURNING a.is_favorite`
var fav bool
err := r.pool.QueryRow(ctx, query, id, userID).Scan(&fav)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return false, false, nil
}
return false, false, fmt.Errorf("toggling favorite: %w", err)
}
return fav, true, nil
}
// MarkFeedRead marks every unread article of an owned feed as read.
func (r *ArticleRepository) MarkFeedRead(ctx context.Context, feedID, userID uuid.UUID) (int64, error) {
const query = `
UPDATE articles a SET is_read = true, read_at = NOW()
FROM feeds f
WHERE a.feed_id = $1 AND f.id = a.feed_id AND f.user_id = $2 AND NOT a.is_read`
tag, err := r.pool.Exec(ctx, query, feedID, userID)
if err != nil {
return 0, fmt.Errorf("marking feed read: %w", err)
}
return tag.RowsAffected(), nil
}
// MarkAllRead marks all of a user's articles as read.
func (r *ArticleRepository) MarkAllRead(ctx context.Context, userID uuid.UUID) (int64, error) {
const query = `
UPDATE articles SET is_read = true, read_at = NOW()
WHERE feed_id IN (SELECT id FROM feeds WHERE user_id = $1) AND NOT is_read`
tag, err := r.pool.Exec(ctx, query, userID)
if err != nil {
return 0, fmt.Errorf("marking all articles read: %w", err)
}
return tag.RowsAffected(), nil
}
// UpdateAISummary updates the AI-generated summary of an article.
func (r *ArticleRepository) UpdateAISummary(ctx context.Context, id uuid.UUID, summary string) error {
if _, err := r.pool.Exec(ctx, `UPDATE articles SET ai_summary = $2 WHERE id = $1`, id, summary); err != nil {
return fmt.Errorf("updating AI summary: %w", err)
}
return nil
@@ -524,18 +292,51 @@ func (r *ArticleRepository) UpdateAISummary(id uuid.UUID, summary string) error
// DeleteOldArticles removes articles older than the specified duration, except for favorites.
func (r *ArticleRepository) DeleteOldArticles(ctx context.Context, olderThan time.Duration) (int64, error) {
threshold := time.Now().Add(-olderThan)
query := `
DELETE FROM articles
WHERE created_at < $1 AND is_favorite = false
`
result, err := r.pool.Exec(ctx, query, threshold)
tag, err := r.pool.Exec(ctx, `DELETE FROM articles WHERE created_at < $1 AND NOT is_favorite`, threshold)
if err != nil {
return 0, fmt.Errorf("deleting old articles: %w", err)
}
return tag.RowsAffected(), nil
}
return result.RowsAffected(), nil
// BackfillRow is a legacy article that still needs sanitization and derived fields.
type BackfillRow struct {
ID uuid.UUID
Content string
Summary string
}
// PendingBackfill returns up to limit articles whose derived columns are missing.
func (r *ArticleRepository) PendingBackfill(ctx context.Context, limit int) ([]BackfillRow, error) {
rows, err := r.pool.Query(ctx,
`SELECT id, content, summary FROM articles WHERE word_count IS NULL LIMIT $1`, limit)
if err != nil {
return nil, fmt.Errorf("querying backfill rows: %w", err)
}
defer rows.Close()
var out []BackfillRow
for rows.Next() {
var row BackfillRow
var content, summary *string
if err := rows.Scan(&row.ID, &content, &summary); err != nil {
return nil, fmt.Errorf("scanning backfill row: %w", err)
}
row.Content, row.Summary = deref(content), deref(summary)
out = append(out, row)
}
return out, rows.Err()
}
// UpdateDerived stores sanitized content and derived list fields for one article.
func (r *ArticleRepository) UpdateDerived(ctx context.Context, id uuid.UUID, content, summary, excerpt string, words int) error {
_, err := r.pool.Exec(ctx,
`UPDATE articles SET content = $2, summary = $3, excerpt = $4, word_count = $5 WHERE id = $1`,
id, nullString(content), nullString(summary), excerpt, words)
if err != nil {
return fmt.Errorf("updating derived fields: %w", err)
}
return nil
}
// nullString returns nil if string is empty.
@@ -545,3 +346,19 @@ func nullString(s string) *string {
}
return &s
}
func deref(s *string) string {
if s == nil {
return ""
}
return *s
}
func firstNonEmpty(values ...string) string {
for _, v := range values {
if v != "" {
return v
}
}
return ""
}
+86 -191
View File
@@ -22,6 +22,34 @@ func NewFeedRepository(pool *pgxpool.Pool) *FeedRepository {
return &FeedRepository{pool: pool}
}
const feedColumns = `
f.id, f.user_id, f.url, f.title, f.description, f.site_url, f.image_url,
f.last_fetched_at, f.fetch_error, f.created_at, f.updated_at,
f.etag, f.last_modified, f.error_count, f.next_fetch_at`
// scanFeed scans feedColumns (plus optional extra destinations).
func scanFeed(row pgx.Row, extra ...any) (*domain.Feed, error) {
var feed domain.Feed
var title, description, siteURL, imageURL, fetchError, etag, lastModified *string
dest := []any{
&feed.ID, &feed.UserID, &feed.URL, &title, &description, &siteURL, &imageURL,
&feed.LastFetchedAt, &fetchError, &feed.CreatedAt, &feed.UpdatedAt,
&etag, &lastModified, &feed.ErrorCount, &feed.NextFetchAt,
}
if err := row.Scan(append(dest, extra...)...); err != nil {
return nil, err
}
feed.Title = deref(title)
feed.Description = deref(description)
feed.SiteURL = deref(siteURL)
feed.ImageURL = deref(imageURL)
feed.FetchError = deref(fetchError)
feed.ETag = deref(etag)
feed.LastModified = deref(lastModified)
return &feed, nil
}
// Create inserts a new feed into the database.
func (r *FeedRepository) Create(feed *domain.Feed) error {
ctx := context.Background()
@@ -52,69 +80,36 @@ func (r *FeedRepository) Create(feed *domain.Feed) error {
// GetByID retrieves a feed by its ID.
func (r *FeedRepository) GetByID(id uuid.UUID) (*domain.Feed, error) {
ctx := context.Background()
query := `
SELECT id, user_id, url, title, description, site_url, image_url,
last_fetched_at, fetch_error, created_at, updated_at
FROM feeds
WHERE id = $1
`
var feed domain.Feed
var description, siteURL, imageURL, fetchError *string
var lastFetchedAt *time.Time
err := r.pool.QueryRow(ctx, query, id).Scan(
&feed.ID,
&feed.UserID,
&feed.URL,
&feed.Title,
&description,
&siteURL,
&imageURL,
&lastFetchedAt,
&fetchError,
&feed.CreatedAt,
&feed.UpdatedAt,
)
query := `SELECT ` + feedColumns + ` FROM feeds f WHERE f.id = $1`
feed, err := scanFeed(r.pool.QueryRow(context.Background(), query, id))
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, fmt.Errorf("getting feed by ID: %w", err)
}
if description != nil {
feed.Description = *description
}
if siteURL != nil {
feed.SiteURL = *siteURL
}
if imageURL != nil {
feed.ImageURL = *imageURL
}
if fetchError != nil {
feed.FetchError = *fetchError
}
feed.LastFetchedAt = lastFetchedAt
return &feed, nil
return feed, nil
}
// GetByUserID retrieves all feeds for a user.
// GetByUserID retrieves all feeds for a user with their unread counts.
func (r *FeedRepository) GetByUserID(userID uuid.UUID) ([]*domain.Feed, error) {
ctx := context.Background()
// One aggregate over the partial unread index instead of a correlated
// COUNT(*) per feed.
query := `
SELECT f.id, f.user_id, f.url, f.title, f.description, f.site_url, f.image_url,
f.last_fetched_at, f.fetch_error, f.created_at, f.updated_at,
COALESCE((SELECT COUNT(*) FROM articles a WHERE a.feed_id = f.id AND a.is_read = false), 0) as unread_count
SELECT ` + feedColumns + `, COALESCE(u.unread, 0)
FROM feeds f
LEFT JOIN (
SELECT a.feed_id, COUNT(*) AS unread
FROM articles a
JOIN feeds uf ON uf.id = a.feed_id AND uf.user_id = $1
WHERE NOT a.is_read
GROUP BY a.feed_id
) u ON u.feed_id = f.id
WHERE f.user_id = $1
ORDER BY f.title ASC
`
ORDER BY lower(f.title) ASC`
rows, err := r.pool.Query(ctx, query, userID)
if err != nil {
@@ -122,101 +117,31 @@ func (r *FeedRepository) GetByUserID(userID uuid.UUID) ([]*domain.Feed, error) {
}
defer rows.Close()
var feeds []*domain.Feed
feeds := make([]*domain.Feed, 0, 16)
for rows.Next() {
var feed domain.Feed
var description, siteURL, imageURL, fetchError *string
var lastFetchedAt *time.Time
err := rows.Scan(
&feed.ID,
&feed.UserID,
&feed.URL,
&feed.Title,
&description,
&siteURL,
&imageURL,
&lastFetchedAt,
&fetchError,
&feed.CreatedAt,
&feed.UpdatedAt,
&feed.UnreadCount,
)
var unread int
feed, err := scanFeed(rows, &unread)
if err != nil {
return nil, fmt.Errorf("scanning feed: %w", err)
}
if description != nil {
feed.Description = *description
}
if siteURL != nil {
feed.SiteURL = *siteURL
}
if imageURL != nil {
feed.ImageURL = *imageURL
}
if fetchError != nil {
feed.FetchError = *fetchError
}
feed.LastFetchedAt = lastFetchedAt
feeds = append(feeds, &feed)
feed.UnreadCount = unread
feeds = append(feeds, feed)
}
return feeds, nil
return feeds, rows.Err()
}
// GetByURL retrieves a feed by its URL for a specific user.
func (r *FeedRepository) GetByURL(userID uuid.UUID, url string) (*domain.Feed, error) {
ctx := context.Background()
query := `
SELECT id, user_id, url, title, description, site_url, image_url,
last_fetched_at, fetch_error, created_at, updated_at
FROM feeds
WHERE user_id = $1 AND url = $2
`
var feed domain.Feed
var description, siteURL, imageURL, fetchError *string
var lastFetchedAt *time.Time
err := r.pool.QueryRow(ctx, query, userID, url).Scan(
&feed.ID,
&feed.UserID,
&feed.URL,
&feed.Title,
&description,
&siteURL,
&imageURL,
&lastFetchedAt,
&fetchError,
&feed.CreatedAt,
&feed.UpdatedAt,
)
query := `SELECT ` + feedColumns + ` FROM feeds f WHERE f.user_id = $1 AND f.url = $2`
feed, err := scanFeed(r.pool.QueryRow(context.Background(), query, userID, url))
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, fmt.Errorf("getting feed by URL: %w", err)
}
if description != nil {
feed.Description = *description
}
if siteURL != nil {
feed.SiteURL = *siteURL
}
if imageURL != nil {
feed.ImageURL = *imageURL
}
if fetchError != nil {
feed.FetchError = *fetchError
}
feed.LastFetchedAt = lastFetchedAt
return &feed, nil
return feed, nil
}
// Update updates a feed in the database.
@@ -258,21 +183,16 @@ func (r *FeedRepository) Delete(id uuid.UUID) error {
return nil
}
// GetFeedsToFetch returns feeds that need to be fetched.
// GetFeedsToFetch returns feeds whose next fetch is due, never-fetched first.
func (r *FeedRepository) GetFeedsToFetch(limit int) ([]*domain.Feed, error) {
ctx := context.Background()
query := `
SELECT id, user_id, url, title, description, site_url, image_url,
last_fetched_at, fetch_error, created_at, updated_at
FROM feeds
WHERE last_fetched_at IS NULL
OR last_fetched_at < NOW() - INTERVAL '15 minutes'
ORDER BY last_fetched_at ASC NULLS FIRST
LIMIT $1
`
SELECT ` + feedColumns + `
FROM feeds f
WHERE f.next_fetch_at IS NULL OR f.next_fetch_at <= NOW()
ORDER BY f.next_fetch_at ASC NULLS FIRST
LIMIT $1`
rows, err := r.pool.Query(ctx, query, limit)
rows, err := r.pool.Query(context.Background(), query, limit)
if err != nil {
return nil, fmt.Errorf("querying feeds to fetch: %w", err)
}
@@ -280,66 +200,41 @@ func (r *FeedRepository) GetFeedsToFetch(limit int) ([]*domain.Feed, error) {
var feeds []*domain.Feed
for rows.Next() {
var feed domain.Feed
var description, siteURL, imageURL, fetchError *string
var lastFetchedAt *time.Time
err := rows.Scan(
&feed.ID,
&feed.UserID,
&feed.URL,
&feed.Title,
&description,
&siteURL,
&imageURL,
&lastFetchedAt,
&fetchError,
&feed.CreatedAt,
&feed.UpdatedAt,
)
feed, err := scanFeed(rows)
if err != nil {
return nil, fmt.Errorf("scanning feed: %w", err)
}
if description != nil {
feed.Description = *description
}
if siteURL != nil {
feed.SiteURL = *siteURL
}
if imageURL != nil {
feed.ImageURL = *imageURL
}
if fetchError != nil {
feed.FetchError = *fetchError
}
feed.LastFetchedAt = lastFetchedAt
feeds = append(feeds, &feed)
feeds = append(feeds, feed)
}
return feeds, nil
return feeds, rows.Err()
}
// UpdateFetchStatus updates the fetch status of a feed.
func (r *FeedRepository) UpdateFetchStatus(id uuid.UUID, fetchedAt time.Time, fetchError string) error {
ctx := context.Background()
// SaveFetchResult records the outcome of a fetch (status, schedule, HTTP
// validators and, when parsed, feed metadata) in a single statement.
// The title is only replaced while it is still the placeholder URL.
func (r *FeedRepository) SaveFetchResult(id uuid.UUID, res domain.FetchResult) error {
const query = `
UPDATE feeds SET
last_fetched_at = $2,
next_fetch_at = $3,
fetch_error = NULLIF($4, ''),
error_count = $5,
etag = COALESCE(NULLIF($6, ''), etag),
last_modified = COALESCE(NULLIF($7, ''), last_modified),
title = CASE WHEN $8::text IS NOT NULL AND $8 <> '' AND (title IS NULL OR title = '' OR title = url)
THEN $8 ELSE title END,
description = COALESCE($9, description),
site_url = COALESCE($10, site_url),
image_url = COALESCE($11, image_url)
WHERE id = $1`
var query string
var args []interface{}
if fetchError == "" {
query = `UPDATE feeds SET last_fetched_at = $2, fetch_error = NULL WHERE id = $1`
args = []interface{}{id, fetchedAt}
} else {
query = `UPDATE feeds SET last_fetched_at = $2, fetch_error = $3 WHERE id = $1`
args = []interface{}{id, fetchedAt, fetchError}
}
_, err := r.pool.Exec(ctx, query, args...)
_, err := r.pool.Exec(context.Background(), query, id,
res.FetchedAt, res.NextFetchAt, res.Error, res.ErrorCount,
res.ETag, res.LastModified,
res.Title, res.Description, res.SiteURL, res.ImageURL,
)
if err != nil {
return fmt.Errorf("updating fetch status: %w", err)
return fmt.Errorf("saving fetch result: %w", err)
}
return nil
}
+41 -13
View File
@@ -2,6 +2,8 @@ package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
@@ -12,6 +14,9 @@ import (
)
// SessionRepository implements domain.SessionRepository using PostgreSQL.
//
// Only the SHA-256 of a session token is stored, so a database leak does not
// hand out live sessions. Callers always pass the raw token.
type SessionRepository struct {
pool *pgxpool.Pool
}
@@ -21,6 +26,12 @@ func NewSessionRepository(pool *pgxpool.Pool) *SessionRepository {
return &SessionRepository{pool: pool}
}
// hashToken returns the hex SHA-256 of a raw session token.
func hashToken(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}
// Create inserts a new session into the database.
func (r *SessionRepository) Create(session *domain.Session) error {
ctx := context.Background()
@@ -33,7 +44,7 @@ func (r *SessionRepository) Create(session *domain.Session) error {
_, err := r.pool.Exec(ctx, query,
session.ID,
session.UserID,
session.Token,
hashToken(session.Token),
session.ExpiresAt,
session.CreatedAt,
session.UserAgent,
@@ -47,22 +58,21 @@ func (r *SessionRepository) Create(session *domain.Session) error {
return nil
}
// GetByToken retrieves a session by its token.
// GetByToken retrieves a non-expired session by its raw token.
func (r *SessionRepository) GetByToken(token string) (*domain.Session, error) {
ctx := context.Background()
query := `
SELECT id, user_id, token, expires_at, created_at, user_agent, ip_address
SELECT id, user_id, expires_at, created_at, user_agent, ip_address
FROM sessions
WHERE token = $1 AND expires_at > NOW()
`
var session domain.Session
var userAgent, ipAddress *string
err := r.pool.QueryRow(ctx, query, token).Scan(
err := r.pool.QueryRow(ctx, query, hashToken(token)).Scan(
&session.ID,
&session.UserID,
&session.Token,
&session.ExpiresAt,
&session.CreatedAt,
&userAgent,
@@ -76,22 +86,40 @@ func (r *SessionRepository) GetByToken(token string) (*domain.Session, error) {
return nil, fmt.Errorf("getting session by token: %w", err)
}
if userAgent != nil {
session.UserAgent = *userAgent
}
if ipAddress != nil {
session.IPAddress = *ipAddress
}
session.Token = token
session.UserAgent = deref(userAgent)
session.IPAddress = deref(ipAddress)
return &session, nil
}
// Delete removes a session by its token.
// GetUserByToken resolves a raw session token to its user in one round trip.
// Returns nil, nil when the session is unknown or expired.
func (r *SessionRepository) GetUserByToken(ctx context.Context, token string) (*domain.User, error) {
const query = `
SELECT u.id, u.email, u.created_at, u.role
FROM sessions s
JOIN users u ON u.id = s.user_id
WHERE s.token = $1 AND s.expires_at > NOW()`
var u domain.User
err := r.pool.QueryRow(ctx, query, hashToken(token)).Scan(&u.ID, &u.Email, &u.CreatedAt, &u.Role)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, fmt.Errorf("resolving session: %w", err)
}
u.IsAdmin = u.Role == domain.RoleAdmin
return &u, nil
}
// Delete removes a session by its raw token.
func (r *SessionRepository) Delete(token string) error {
ctx := context.Background()
query := `DELETE FROM sessions WHERE token = $1`
_, err := r.pool.Exec(ctx, query, token)
_, err := r.pool.Exec(ctx, query, hashToken(token))
if err != nil {
return fmt.Errorf("deleting session: %w", err)
}
+37 -58
View File
@@ -22,98 +22,76 @@ func NewUserRepository(pool *pgxpool.Pool) *UserRepository {
return &UserRepository{pool: pool}
}
// Create inserts a new user into the database.
// Create inserts a new user into the database. The very first account becomes
// admin; the decision is made atomically under an advisory lock so two
// concurrent first registrations can't both be promoted.
func (r *UserRepository) Create(user *domain.User) error {
ctx := context.Background()
// Check if this is the first user
var count int
err := r.pool.QueryRow(ctx, "SELECT COUNT(*) FROM users").Scan(&count)
tx, err := r.pool.Begin(ctx)
if err != nil {
return fmt.Errorf("counting users: %w", err)
return fmt.Errorf("starting transaction: %w", err)
}
defer tx.Rollback(ctx)
// Arbitrary constant key serialising user creation.
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(7314001)`); err != nil {
return fmt.Errorf("locking user creation: %w", err)
}
if count == 0 {
user.Role = domain.RoleAdmin
} else if user.Role == "" {
user.Role = domain.RoleUser
role := user.Role
if role == "" {
role = domain.RoleUser
}
query := `
const query = `
INSERT INTO users (id, email, password_hash, created_at, role)
VALUES ($1, $2, $3, $4, $5)
`
VALUES ($1, $2, $3, $4,
CASE WHEN NOT EXISTS (SELECT 1 FROM users) THEN 'admin' ELSE $5 END)
RETURNING role`
_, err = r.pool.Exec(ctx, query,
user.ID,
user.Email,
user.PasswordHash,
user.CreatedAt,
user.Role,
)
if err != nil {
if err := tx.QueryRow(ctx, query,
user.ID, user.Email, user.PasswordHash, user.CreatedAt, role,
).Scan(&user.Role); err != nil {
return fmt.Errorf("creating user: %w", err)
}
user.IsAdmin = user.Role == domain.RoleAdmin
return nil
return tx.Commit(ctx)
}
// GetByEmail retrieves a user by their email address.
// GetByEmail retrieves a user by their email address (case-insensitive).
func (r *UserRepository) GetByEmail(email string) (*domain.User, error) {
ctx := context.Background()
query := `
return r.getOne(`
SELECT id, email, password_hash, created_at, role
FROM users
WHERE email = $1
`
var user domain.User
err := r.pool.QueryRow(ctx, query, email).Scan(
&user.ID,
&user.Email,
&user.PasswordHash,
&user.CreatedAt,
&user.Role,
)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil // User not found
}
return nil, fmt.Errorf("getting user by email: %w", err)
}
return &user, nil
WHERE lower(email) = lower($1)`, email)
}
// GetByID retrieves a user by their ID.
func (r *UserRepository) GetByID(id uuid.UUID) (*domain.User, error) {
ctx := context.Background()
query := `
return r.getOne(`
SELECT id, email, password_hash, created_at, role
FROM users
WHERE id = $1
`
WHERE id = $1`, id)
}
func (r *UserRepository) getOne(query string, arg any) (*domain.User, error) {
var user domain.User
err := r.pool.QueryRow(ctx, query, id).Scan(
err := r.pool.QueryRow(context.Background(), query, arg).Scan(
&user.ID,
&user.Email,
&user.PasswordHash,
&user.CreatedAt,
&user.Role,
)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil // User not found
}
return nil, fmt.Errorf("getting user by ID: %w", err)
return nil, fmt.Errorf("getting user: %w", err)
}
user.IsAdmin = user.Role == domain.RoleAdmin
return &user, nil
}
@@ -134,10 +112,11 @@ func (r *UserRepository) List() ([]*domain.User, error) {
if err := rows.Scan(&u.ID, &u.Email, &u.CreatedAt, &u.Role); err != nil {
return nil, fmt.Errorf("scanning user: %w", err)
}
u.IsAdmin = u.Role == domain.RoleAdmin
users = append(users, &u)
}
return users, nil
return users, rows.Err()
}
// Delete removes a user and their data (cascaded by DB).
@@ -150,11 +129,11 @@ func (r *UserRepository) Delete(id uuid.UUID) error {
return nil
}
// Exists checks if an email already exists.
// Exists checks if an email already exists (case-insensitive).
func (r *UserRepository) Exists(email string) (bool, error) {
ctx := context.Background()
query := `SELECT EXISTS(SELECT 1 FROM users WHERE email = $1)`
query := `SELECT EXISTS(SELECT 1 FROM users WHERE lower(email) = lower($1))`
var exists bool
err := r.pool.QueryRow(ctx, query, email).Scan(&exists)