mirror of
https://github.com/R0m1k3/FlowReader.git
synced 2026-10-11 17:28:05 +02:00
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:
1 parent
d037e2be34
commit
03e57e4308
40 files changed
+2231
-1898
No files matched your search
+279
-462
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in new issue
Block a user