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

+6 -14
View File
@@ -44,11 +44,9 @@ func (h *AdminHandler) DeleteUser(w http.ResponseWriter, r *http.Request) {
// Prevent an admin from deleting their own account (lockout / accidental
// self-removal). The AdminOnly middleware has already verified admin rights.
if cookie, cErr := r.Cookie("session_id"); cErr == nil {
if current, _ := h.authService.GetUserByToken(cookie.Value); current != nil && current.ID == userID {
respondError(w, http.StatusForbidden, "You cannot delete your own account")
return
}
if current := currentUser(r); current != nil && current.ID == userID {
respondError(w, http.StatusForbidden, "You cannot delete your own account")
return
}
// Logic to delete user and all associated data
@@ -60,21 +58,15 @@ func (h *AdminHandler) DeleteUser(w http.ResponseWriter, r *http.Request) {
respondJSON(w, http.StatusOK, map[string]string{"message": "User deleted successfully"})
}
// AdminOnly middleware restricts access to admins.
// AdminOnly middleware restricts access to admins. Must run after RequireAuth.
func (h *AdminHandler) AdminOnly(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
cookie, err := r.Cookie("session_id")
if err != nil {
user := currentUser(r)
if user == nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
user, err := h.authService.GetUserByToken(cookie.Value)
if err != nil || user == nil {
respondError(w, http.StatusUnauthorized, "Invalid session")
return
}
if user.Role != domain.RoleAdmin {
respondError(w, http.StatusForbidden, "Admin access required")
return
+193 -377
View File
@@ -1,9 +1,10 @@
package handler
import (
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/go-chi/chi/v5"
"github.com/google/uuid"
@@ -13,11 +14,10 @@ import (
"github.com/michael/flowreader/internal/ws"
)
// ArticleHandler handles article-related HTTP requests.
// ArticleHandler handles article-related HTTP requests. All routes run
// behind RequireAuth; ownership is enforced inside each SQL statement.
type ArticleHandler struct {
articleRepo domain.ArticleRepository
feedService *service.FeedService
authService *service.AuthService
aiService *service.AIService
sanitizer *utils.ContentSanitizer
extractor *utils.ContentExtractor
@@ -25,11 +25,9 @@ type ArticleHandler struct {
}
// NewArticleHandler creates a new article handler.
func NewArticleHandler(articleRepo domain.ArticleRepository, feedService *service.FeedService, authService *service.AuthService, aiService *service.AIService, hub *ws.Hub) *ArticleHandler {
func NewArticleHandler(articleRepo domain.ArticleRepository, aiService *service.AIService, hub *ws.Hub) *ArticleHandler {
return &ArticleHandler{
articleRepo: articleRepo,
feedService: feedService,
authService: authService,
aiService: aiService,
sanitizer: utils.NewContentSanitizer(),
extractor: utils.NewContentExtractor(),
@@ -37,409 +35,232 @@ func NewArticleHandler(articleRepo domain.ArticleRepository, feedService *servic
}
}
// sanitizeArticle cleans the user-facing HTML fields of a single article to
// prevent stored XSS from malicious feeds. AISummary is rendered as plain text
// by the client, so only Content and Summary need sanitization.
func (h *ArticleHandler) sanitizeArticle(a *domain.Article) {
if a == nil {
// notify pushes an event to the acting user's other tabs/devices only.
func (h *ArticleHandler) notify(userID uuid.UUID, eventType string, payload any) {
if h.hub != nil {
h.hub.SendToUser(userID, eventType, payload)
}
}
func parseLimit(r *http.Request) int {
limit, _ := strconv.Atoi(r.URL.Query().Get("limit"))
if limit <= 0 || limit > 100 {
limit = 30
}
return limit
}
// parseCursor reads ?cursor=<RFC3339 sort_at>,<uuid> (the last item of the
// previous page). A missing cursor means the first page.
func parseCursor(r *http.Request) (*domain.ArticleCursor, bool) {
raw := r.URL.Query().Get("cursor")
if raw == "" {
return nil, true
}
ts, id, ok := strings.Cut(raw, ",")
if !ok {
return nil, false
}
sortAt, err := time.Parse(time.RFC3339Nano, ts)
if err != nil {
return nil, false
}
uid, err := uuid.Parse(id)
if err != nil {
return nil, false
}
return &domain.ArticleCursor{SortAt: sortAt, ID: uid}, true
}
func (h *ArticleHandler) list(w http.ResponseWriter, r *http.Request, f domain.ArticleFilter) {
cursor, ok := parseCursor(r)
if !ok {
respondError(w, http.StatusBadRequest, "Invalid cursor")
return
}
if a.Content != "" {
a.Content = h.sanitizer.Sanitize(a.Content)
}
if a.Summary != "" {
a.Summary = h.sanitizer.Sanitize(a.Summary)
}
}
f.UserID = currentUser(r).ID
f.Cursor = cursor
f.Limit = parseLimit(r)
f.UnreadOnly = f.UnreadOnly || r.URL.Query().Get("unread") == "true"
// sanitizeArticles cleans a slice of articles in place.
func (h *ArticleHandler) sanitizeArticles(articles []*domain.Article) {
for _, a := range articles {
h.sanitizeArticle(a)
}
}
// getUserFromRequest extracts the authenticated user from the request.
func (h *ArticleHandler) getUserFromRequest(r *http.Request) (uuid.UUID, error) {
cookie, err := r.Cookie("session_id")
articles, err := h.articleRepo.List(r.Context(), f)
if err != nil {
return uuid.Nil, err
respondError(w, http.StatusInternalServerError, "Failed to get articles")
return
}
user, err := h.authService.GetUserByToken(cookie.Value)
if err != nil || user == nil {
return uuid.Nil, err
}
return user.ID, nil
respondJSON(w, http.StatusOK, articles)
}
// List handles GET /api/v1/articles
func (h *ArticleHandler) List(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
// Parse query parameters
limit, _ := strconv.Atoi(r.URL.Query().Get("limit"))
if limit <= 0 || limit > 100 {
limit = 50
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
if offset < 0 {
offset = 0
}
unreadOnly := r.URL.Query().Get("unread") == "true"
articles, err := h.articleRepo.GetByUserID(userID, limit, offset, unreadOnly)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to get articles")
return
}
h.sanitizeArticles(articles)
respondJSON(w, http.StatusOK, articles)
h.list(w, r, domain.ArticleFilter{})
}
// ListByFeed handles GET /api/v1/feeds/{id}/articles
func (h *ArticleHandler) ListByFeed(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
feedID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid feed ID")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(feedID, userID)
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
return
}
// Parse query parameters
limit, _ := strconv.Atoi(r.URL.Query().Get("limit"))
if limit <= 0 || limit > 100 {
limit = 50
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
if offset < 0 {
offset = 0
}
articles, err := h.articleRepo.GetByFeedID(feedID, limit, offset)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to get articles")
return
}
h.sanitizeArticles(articles)
respondJSON(w, http.StatusOK, articles)
}
// Get handles GET /api/v1/articles/{id}
func (h *ArticleHandler) Get(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
article, err := h.articleRepo.GetByID(articleID)
if err != nil || article == nil {
respondError(w, http.StatusNotFound, "Article not found")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(article.FeedID, userID)
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
return
}
// Sanitize user-facing HTML content (defense against stored XSS).
h.sanitizeArticle(article)
respondJSON(w, http.StatusOK, article)
}
// MarkRead handles POST /api/v1/articles/{id}/read
func (h *ArticleHandler) MarkRead(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
article, err := h.articleRepo.GetByID(articleID)
if err != nil || article == nil {
respondError(w, http.StatusNotFound, "Article not found")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(article.FeedID, userID)
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
return
}
if err := h.articleRepo.MarkAsRead(articleID); err != nil {
respondError(w, http.StatusInternalServerError, "Failed to mark as read")
return
}
// Broadcast update
if h.hub != nil {
h.hub.Broadcast("article_updated", map[string]interface{}{
"id": articleID,
"is_read": true,
})
}
respondJSON(w, http.StatusOK, map[string]bool{"is_read": true})
}
// MarkUnread handles DELETE /api/v1/articles/{id}/read
func (h *ArticleHandler) MarkUnread(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
article, err := h.articleRepo.GetByID(articleID)
if err != nil || article == nil {
respondError(w, http.StatusNotFound, "Article not found")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(article.FeedID, userID)
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
return
}
if err := h.articleRepo.MarkAsUnread(articleID); err != nil {
respondError(w, http.StatusInternalServerError, "Failed to mark as unread")
return
}
// Broadcast update
if h.hub != nil {
h.hub.Broadcast("article_updated", map[string]interface{}{
"id": articleID,
"is_read": false,
})
}
respondJSON(w, http.StatusOK, map[string]bool{"is_read": false})
}
// ToggleFavorite handles POST /api/v1/articles/{id}/favorite
func (h *ArticleHandler) ToggleFavorite(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
article, err := h.articleRepo.GetByID(articleID)
if err != nil || article == nil {
respondError(w, http.StatusNotFound, "Article not found")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(article.FeedID, userID)
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
return
}
if err := h.articleRepo.ToggleFavorite(articleID); err != nil {
respondError(w, http.StatusInternalServerError, "Failed to toggle favorite")
return
}
// Broadcast update
if h.hub != nil {
h.hub.Broadcast("article_updated", map[string]interface{}{
"id": articleID,
"is_favorite": !article.IsFavorite,
})
}
respondJSON(w, http.StatusOK, map[string]bool{"is_favorite": !article.IsFavorite})
}
// MarkAllRead handles POST /api/v1/feeds/{id}/read-all
func (h *ArticleHandler) MarkAllRead(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
feedID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid feed ID")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(feedID, userID)
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
return
}
if err := h.articleRepo.MarkAllAsRead(feedID); err != nil {
respondError(w, http.StatusInternalServerError, "Failed to mark all as read")
return
}
respondJSON(w, http.StatusOK, map[string]string{"message": "All articles marked as read"})
}
// MarkAllReadGlobal handles POST /api/v1/articles/read-all
func (h *ArticleHandler) MarkAllReadGlobal(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
if err := h.articleRepo.MarkAllAsReadGlobal(userID); err != nil {
respondError(w, http.StatusInternalServerError, "Failed to mark all as read")
return
}
respondJSON(w, http.StatusOK, map[string]string{"message": "All articles marked as read"})
h.list(w, r, domain.ArticleFilter{FeedID: &feedID})
}
// GetFavorites handles GET /api/v1/articles/favorites
func (h *ArticleHandler) GetFavorites(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
limit, _ := strconv.Atoi(r.URL.Query().Get("limit"))
if limit <= 0 || limit > 100 {
limit = 50
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
if offset < 0 {
offset = 0
}
articles, err := h.articleRepo.GetFavorites(userID, limit, offset)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to get favorites")
return
}
h.sanitizeArticles(articles)
respondJSON(w, http.StatusOK, articles)
h.list(w, r, domain.ArticleFilter{FavoritesOnly: true})
}
// Search handles GET /api/v1/articles/search
func (h *ArticleHandler) Search(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
query := r.URL.Query().Get("q")
query := strings.TrimSpace(r.URL.Query().Get("q"))
if query == "" {
respondJSON(w, http.StatusOK, []*domain.Article{})
return
}
limit, _ := strconv.Atoi(r.URL.Query().Get("limit"))
if limit <= 0 || limit > 100 {
limit = 50
}
query = utils.TruncateRunes(query, 200)
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
if offset < 0 {
if offset < 0 || offset > 10000 {
offset = 0
}
articles, err := h.articleRepo.Search(userID, query, limit, offset)
articles, err := h.articleRepo.Search(r.Context(), currentUser(r).ID, query, parseLimit(r), offset)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to search articles")
return
}
h.sanitizeArticles(articles)
respondJSON(w, http.StatusOK, articles)
}
// Summarize handles POST /api/v1/articles/{id}/summarize
func (h *ArticleHandler) Summarize(w http.ResponseWriter, r *http.Request) {
userID, err := h.getUserFromRequest(r)
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
// Get handles GET /api/v1/articles/{id} and returns the full content.
func (h *ArticleHandler) Get(w http.ResponseWriter, r *http.Request) {
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
article, err := h.articleRepo.GetByID(articleID)
if err != nil || article == nil {
article, err := h.articleRepo.GetForUser(r.Context(), articleID, currentUser(r).ID)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to get article")
return
}
if article == nil {
respondError(w, http.StatusNotFound, "Article not found")
return
}
// Verify feed ownership
_, err = h.feedService.GetFeed(article.FeedID, userID)
// Content is sanitized at ingest; sanitizing again here is cheap for a
// single article and keeps defence in depth for rows not yet backfilled.
article.Content = h.sanitizer.Sanitize(article.Content)
article.Summary = h.sanitizer.Sanitize(article.Summary)
w.Header().Set("Cache-Control", "private, no-cache")
respondJSON(w, http.StatusOK, article)
}
func (h *ArticleHandler) setRead(w http.ResponseWriter, r *http.Request, read bool) {
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusForbidden, "Access denied")
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
userID := currentUser(r).ID
found, err := h.articleRepo.SetRead(r.Context(), articleID, userID, read)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to update article")
return
}
if !found {
respondError(w, http.StatusNotFound, "Article not found")
return
}
h.notify(userID, "article_updated", map[string]any{"id": articleID, "is_read": read})
respondJSON(w, http.StatusOK, map[string]bool{"is_read": read})
}
// MarkRead handles POST /api/v1/articles/{id}/read
func (h *ArticleHandler) MarkRead(w http.ResponseWriter, r *http.Request) { h.setRead(w, r, true) }
// MarkUnread handles DELETE /api/v1/articles/{id}/read
func (h *ArticleHandler) MarkUnread(w http.ResponseWriter, r *http.Request) { h.setRead(w, r, false) }
// ToggleFavorite handles POST /api/v1/articles/{id}/favorite
func (h *ArticleHandler) ToggleFavorite(w http.ResponseWriter, r *http.Request) {
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
userID := currentUser(r).ID
fav, found, err := h.articleRepo.ToggleFavorite(r.Context(), articleID, userID)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to toggle favorite")
return
}
if !found {
respondError(w, http.StatusNotFound, "Article not found")
return
}
h.notify(userID, "article_updated", map[string]any{"id": articleID, "is_favorite": fav})
respondJSON(w, http.StatusOK, map[string]bool{"is_favorite": fav})
}
// MarkAllRead handles POST /api/v1/feeds/{id}/read-all
func (h *ArticleHandler) MarkAllRead(w http.ResponseWriter, r *http.Request) {
feedID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid feed ID")
return
}
userID := currentUser(r).ID
n, err := h.articleRepo.MarkFeedRead(r.Context(), feedID, userID)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to mark all as read")
return
}
if n > 0 {
h.notify(userID, "articles_bulk_read", map[string]any{"feed_id": feedID})
}
respondJSON(w, http.StatusOK, map[string]any{"message": "All articles marked as read", "count": n})
}
// MarkAllReadGlobal handles POST /api/v1/articles/read-all
func (h *ArticleHandler) MarkAllReadGlobal(w http.ResponseWriter, r *http.Request) {
userID := currentUser(r).ID
n, err := h.articleRepo.MarkAllRead(r.Context(), userID)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to mark all as read")
return
}
if n > 0 {
h.notify(userID, "articles_bulk_read", map[string]any{})
}
respondJSON(w, http.StatusOK, map[string]any{"message": "All articles marked as read", "count": n})
}
// Summarize handles POST /api/v1/articles/{id}/summarize
func (h *ArticleHandler) Summarize(w http.ResponseWriter, r *http.Request) {
articleID, err := uuid.Parse(chi.URLParam(r, "id"))
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid article ID")
return
}
userID := currentUser(r).ID
article, err := h.articleRepo.GetForUser(r.Context(), articleID, userID)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to get article")
return
}
if article == nil {
respondError(w, http.StatusNotFound, "Article not found")
return
}
@@ -448,44 +269,39 @@ func (h *ArticleHandler) Summarize(w http.ResponseWriter, r *http.Request) {
respondJSON(w, http.StatusOK, map[string]string{"summary": article.AISummary})
return
}
if !h.aiService.Enabled() {
respondError(w, http.StatusServiceUnavailable, "AI summaries are not configured")
return
}
// Context for AI is title + content (or summary if content empty)
content := article.Content
// Summaries take an extractor fetch plus an LLM call: allow more than the
// server's default write deadline for this route only.
_ = http.NewResponseController(w).SetWriteDeadline(time.Now().Add(90 * time.Second))
content := utils.PlainText(article.Content)
if content == "" {
content = article.Summary
content = utils.PlainText(article.Summary)
}
// Try to extract full content from URL if available
if article.URL != "" {
fullContent, err := h.extractor.Extract(r.Context(), article.URL)
if err == nil && len(fullContent) > len(content) {
content = "--- CONTENU COMPLET EXTRAIT DU SITE WEB ---\n" + fullContent
content = fullContent
}
}
aiInput := fmt.Sprintf("Titre: %s\n\nContenu: %s", article.Title, content)
// Summary generation (can be slow, but for this demo/small app we do it synchronously
// or we could use WS to notify when done. Here we follow the simple POST -> String pattern).
summary, err := h.aiService.Summarize(r.Context(), aiInput)
summary, err := h.aiService.Summarize(r.Context(), "Titre : "+article.Title+"\n\n"+content)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to generate summary: "+err.Error())
respondError(w, http.StatusBadGateway, "Failed to generate summary")
return
}
// Save to DB
if err := h.articleRepo.UpdateAISummary(articleID, summary); err != nil {
if err := h.articleRepo.UpdateAISummary(r.Context(), articleID, summary); err != nil {
respondError(w, http.StatusInternalServerError, "Failed to persist summary")
return
}
// Broadcast update via WebSocket
if h.hub != nil {
h.hub.Broadcast("article_updated", map[string]interface{}{
"id": articleID,
"ai_summary": summary,
})
}
h.notify(userID, "article_updated", map[string]any{"id": articleID, "ai_summary": summary})
respondJSON(w, http.StatusOK, map[string]string{"summary": summary})
}
+12 -39
View File
@@ -52,8 +52,12 @@ func (h *AuthHandler) Register(w http.ResponseWriter, r *http.Request) {
respondError(w, http.StatusBadRequest, "Invalid email format")
case errors.Is(err, service.ErrPasswordTooShort):
respondError(w, http.StatusBadRequest, "Password must be at least 8 characters")
case errors.Is(err, service.ErrPasswordTooLong):
respondError(w, http.StatusBadRequest, "Password is too long")
case errors.Is(err, service.ErrEmailAlreadyExists):
respondError(w, http.StatusConflict, "Email already registered")
case errors.Is(err, service.ErrRegistrationClosed):
respondError(w, http.StatusForbidden, "Registration is disabled on this instance")
default:
respondError(w, http.StatusInternalServerError, "Registration failed")
}
@@ -123,24 +127,9 @@ func (h *AuthHandler) Logout(w http.ResponseWriter, r *http.Request) {
respondJSON(w, http.StatusOK, map[string]string{"message": "Logged out successfully"})
}
// Me handles GET /api/v1/users/me
// Me handles GET /api/v1/users/me (behind RequireAuth).
func (h *AuthHandler) Me(w http.ResponseWriter, r *http.Request) {
cookie, err := r.Cookie("session_id")
if err != nil {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
user, err := h.authService.GetUserByToken(cookie.Value)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to get user")
return
}
if user == nil {
respondError(w, http.StatusUnauthorized, "Session expired")
return
}
user := currentUser(r)
respondJSON(w, http.StatusOK, service.UserInfo{
ID: user.ID,
Email: user.Email,
@@ -148,31 +137,15 @@ func (h *AuthHandler) Me(w http.ResponseWriter, r *http.Request) {
})
}
// getClientIP extracts the client IP from the request.
func getClientIP(r *http.Request) string {
// Check X-Forwarded-For header
forwarded := r.Header.Get("X-Forwarded-For")
if forwarded != "" {
// Take the first IP in the chain
parts := strings.Split(forwarded, ",")
return strings.TrimSpace(parts[0])
}
// Check X-Real-IP header
realIP := r.Header.Get("X-Real-IP")
if realIP != "" {
return realIP
}
// Fall back to RemoteAddr
return r.RemoteAddr
}
// respondJSON writes a JSON response.
func respondJSON(w http.ResponseWriter, status int, data interface{}) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
json.NewEncoder(w).Encode(data)
enc := json.NewEncoder(w)
// Safe: served as application/json with nosniff; avoids inflating HTML
// content with \u003c escapes.
enc.SetEscapeHTML(false)
enc.Encode(data)
}
// respondError writes an error response.
+24 -27
View File
@@ -29,19 +29,12 @@ func NewFeedHandler(feedService *service.FeedService, fetchService *service.Fetc
}
}
// getUserFromRequest extracts the authenticated user from the request.
// getUserFromRequest returns the user resolved by RequireAuth.
func (h *FeedHandler) getUserFromRequest(r *http.Request) (uuid.UUID, error) {
cookie, err := r.Cookie("session_id")
if err != nil {
return uuid.Nil, errors.New("not authenticated")
if u := currentUser(r); u != nil {
return u.ID, nil
}
user, err := h.authService.GetUserByToken(cookie.Value)
if err != nil || user == nil {
return uuid.Nil, errors.New("invalid session")
}
return user.ID, nil
return uuid.Nil, errors.New("not authenticated")
}
// List handles GET /api/v1/feeds
@@ -91,7 +84,7 @@ func (h *FeedHandler) Add(w http.ResponseWriter, r *http.Request) {
// Trigger immediate fetch in background
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), time.Minute)
defer cancel()
_ = h.fetchService.FetchFeed(ctx, resp.ID)
}()
@@ -107,18 +100,13 @@ func (h *FeedHandler) Refresh(w http.ResponseWriter, r *http.Request) {
return
}
// For now, we refresh all feeds for the user synchronously or in background
// Let's do background and return 202 Accepted
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
feeds, _ := h.feedService.GetUserFeeds(userID)
for _, f := range feeds {
_ = h.fetchService.FetchFeed(ctx, f.ID)
}
}()
// Runs in the background through the shared worker pool; concurrent
// clicks for the same user are coalesced and recently fetched feeds skipped.
started := h.fetchService.RefreshUser(userID)
if !started {
respondJSON(w, http.StatusAccepted, map[string]string{"message": "Refresh already running"})
return
}
respondJSON(w, http.StatusAccepted, map[string]string{"message": "Refresh started"})
}
@@ -236,8 +224,8 @@ func (h *FeedHandler) ImportOPML(w http.ResponseWriter, r *http.Request) {
return
}
// Parse multipart form (max 10MB)
if err := r.ParseMultipartForm(10 << 20); err != nil {
// Parse multipart form (body already capped by the router; keep it in memory)
if err := r.ParseMultipartForm(5 << 20); err != nil {
respondError(w, http.StatusBadRequest, "Invalid form data")
return
}
@@ -252,7 +240,7 @@ func (h *FeedHandler) ImportOPML(w http.ResponseWriter, r *http.Request) {
// Parse OPML
feeds, err := opml.Parse(file)
if err != nil {
respondError(w, http.StatusBadRequest, "Invalid OPML file: "+err.Error())
respondError(w, http.StatusBadRequest, "Invalid OPML file")
return
}
@@ -269,10 +257,19 @@ func (h *FeedHandler) ImportOPML(w http.ResponseWriter, r *http.Request) {
// Import feeds
result, err := h.feedService.ImportOPML(userID, opmlFeeds)
if err != nil {
if errors.Is(err, service.ErrTooManyFeeds) {
respondError(w, http.StatusBadRequest, err.Error())
return
}
respondError(w, http.StatusInternalServerError, "Import failed")
return
}
// Fetch the newly imported feeds right away.
if result.Imported > 0 {
h.fetchService.RefreshUser(userID)
}
respondJSON(w, http.StatusOK, result)
}
+173
View File
@@ -0,0 +1,173 @@
package handler
import (
"context"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"strings"
"github.com/michael/flowreader/internal/domain"
"github.com/michael/flowreader/internal/service"
)
type ctxKey int
const userCtxKey ctxKey = iota
// RequireAuth resolves the session cookie once per request (one SQL query)
// and stores the user in the request context. Unauthenticated requests get 401.
func RequireAuth(authService *service.AuthService) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
cookie, err := r.Cookie("session_id")
if err != nil || cookie.Value == "" {
respondError(w, http.StatusUnauthorized, "Not authenticated")
return
}
user, err := authService.GetUserByTokenCtx(r.Context(), cookie.Value)
if err != nil {
respondError(w, http.StatusInternalServerError, "Failed to resolve session")
return
}
if user == nil {
respondError(w, http.StatusUnauthorized, "Session expired")
return
}
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), userCtxKey, user)))
})
}
}
// currentUser returns the user set by RequireAuth. It is only nil for routes
// not wrapped by RequireAuth.
func currentUser(r *http.Request) *domain.User {
u, _ := r.Context().Value(userCtxKey).(*domain.User)
return u
}
// trustedProxies lists CIDRs (TRUSTED_PROXIES, comma-separated) whose
// X-Forwarded-For header is believed. Empty means: trust no proxy header.
var trustedProxies = parsePrefixes(os.Getenv("TRUSTED_PROXIES"))
func parsePrefixes(raw string) []netip.Prefix {
var out []netip.Prefix
for _, s := range strings.Split(raw, ",") {
s = strings.TrimSpace(s)
if s == "" {
continue
}
if !strings.Contains(s, "/") {
if a, err := netip.ParseAddr(s); err == nil {
out = append(out, netip.PrefixFrom(a, a.BitLen()))
}
continue
}
if p, err := netip.ParsePrefix(s); err == nil {
out = append(out, p.Masked())
}
}
return out
}
func isTrustedProxy(a netip.Addr) bool {
a = a.Unmap()
for _, p := range trustedProxies {
if p.Contains(a) {
return true
}
}
return false
}
// getClientIP returns the peer address, honouring X-Forwarded-For only when
// the direct peer is a configured trusted proxy. The right-most untrusted
// entry is used, which a client cannot spoof.
func getClientIP(r *http.Request) string {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
host = r.RemoteAddr
}
peer, err := netip.ParseAddr(host)
if err != nil || !isTrustedProxy(peer) {
return host
}
hops := strings.Split(r.Header.Get("X-Forwarded-For"), ",")
for i := len(hops) - 1; i >= 0; i-- {
hop := strings.TrimSpace(hops[i])
a, err := netip.ParseAddr(hop)
if err != nil {
break
}
if !isTrustedProxy(a) {
return a.Unmap().String()
}
}
return host
}
// SecurityHeaders sets defensive response headers, including a CSP that
// backstops the HTML sanitizer for feed content.
func SecurityHeaders(next http.Handler) http.Handler {
const csp = "default-src 'self'; " +
"script-src 'self'; " +
"style-src 'self' 'unsafe-inline'; " +
"img-src * data: blob:; " +
"media-src *; " +
"font-src 'self' data:; " +
"connect-src 'self'; " +
"frame-src 'none'; object-src 'none'; base-uri 'none'; " +
"frame-ancestors 'none'; form-action 'self'; manifest-src 'self'; worker-src 'self'"
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
h := w.Header()
h.Set("X-Content-Type-Options", "nosniff")
h.Set("X-Frame-Options", "DENY")
h.Set("Referrer-Policy", "strict-origin-when-cross-origin")
h.Set("Permissions-Policy", "geolocation=(), microphone=(), camera=()")
h.Set("Content-Security-Policy", csp)
h.Set("Cross-Origin-Opener-Policy", "same-origin")
if secureCookie(r) {
h.Set("Strict-Transport-Security", "max-age=31536000")
}
next.ServeHTTP(w, r)
})
}
// LimitBody caps request bodies to n bytes.
func LimitBody(n int64) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Body != nil {
r.Body = http.MaxBytesReader(w, r.Body, n)
}
next.ServeHTTP(w, r)
})
}
}
// SameOriginGuard rejects state-changing requests coming from another site
// (CSRF defence in depth on top of SameSite=Strict cookies).
func SameOriginGuard(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet, http.MethodHead, http.MethodOptions:
next.ServeHTTP(w, r)
return
}
if site := r.Header.Get("Sec-Fetch-Site"); site != "" {
if site != "same-origin" && site != "none" {
respondError(w, http.StatusForbidden, "Cross-site request blocked")
return
}
} else if origin := r.Header.Get("Origin"); origin != "" {
u, err := url.Parse(origin)
if err != nil || !strings.EqualFold(u.Host, r.Host) {
respondError(w, http.StatusForbidden, "Cross-site request blocked")
return
}
}
next.ServeHTTP(w, r)
})
}
+38
View File
@@ -0,0 +1,38 @@
package handler
import (
"net/http/httptest"
"testing"
)
func TestGetClientIPIgnoresSpoofedHeadersWithoutTrustedProxy(t *testing.T) {
trustedProxies = nil
r := httptest.NewRequest("POST", "/api/v1/auth/login", nil)
r.RemoteAddr = "203.0.113.7:5555"
r.Header.Set("X-Forwarded-For", "1.2.3.4")
if got := getClientIP(r); got != "203.0.113.7" {
t.Fatalf("got %q, want peer address", got)
}
}
func TestGetClientIPUsesRightmostUntrustedHop(t *testing.T) {
trustedProxies = parsePrefixes("10.0.0.0/8")
defer func() { trustedProxies = nil }()
r := httptest.NewRequest("POST", "/", nil)
r.RemoteAddr = "10.0.0.2:443"
r.Header.Set("X-Forwarded-For", "6.6.6.6, 198.51.100.9, 10.0.0.3")
if got := getClientIP(r); got != "198.51.100.9" {
t.Fatalf("got %q, want 198.51.100.9", got)
}
}
func TestSameOriginGuard(t *testing.T) {
h := SameOriginGuard(nil)
r := httptest.NewRequest("POST", "http://reader.example/api/v1/articles/read-all", nil)
r.Header.Set("Sec-Fetch-Site", "cross-site")
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
if w.Code != 403 {
t.Fatalf("cross-site POST got %d, want 403", w.Code)
}
}
+33 -7
View File
@@ -16,6 +16,9 @@ type rateLimiter struct {
capacity float64 // max tokens (burst)
}
// maxBuckets caps the number of tracked keys per limiter.
const maxBuckets = 50_000
type bucket struct {
tokens float64
last time.Time
@@ -40,6 +43,11 @@ func (rl *rateLimiter) allow(key string) bool {
now := time.Now()
b, ok := rl.buckets[key]
if !ok {
// Bound memory: under a flood of distinct keys, fail closed rather
// than growing the map without limit.
if len(rl.buckets) >= maxBuckets {
return false
}
rl.buckets[key] = &bucket{tokens: rl.capacity - 1, last: now}
return true
}
@@ -75,14 +83,21 @@ func (rl *rateLimiter) cleanupLoop() {
// Middleware returns a chi-compatible middleware enforcing the limit per IP.
func (rl *rateLimiter) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !rl.allow(getClientIP(r)) {
return rl.middlewareBy(getClientIP)(next)
}
// middlewareBy enforces the limit per key computed from the request.
func (rl *rateLimiter) middlewareBy(key func(*http.Request) string) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !rl.allow(key(r)) {
w.Header().Set("Retry-After", "60")
respondError(w, http.StatusTooManyRequests, "Too many requests. Please slow down.")
return
}
next.ServeHTTP(w, r)
})
respondError(w, http.StatusTooManyRequests, "Too many requests. Please slow down.")
return
}
next.ServeHTTP(w, r)
})
}
}
// NewAuthRateLimiter builds the limiter used for authentication routes:
@@ -90,3 +105,14 @@ func (rl *rateLimiter) Middleware(next http.Handler) http.Handler {
func NewAuthRateLimiter() func(http.Handler) http.Handler {
return newRateLimiter(10, 5).Middleware
}
// NewUserRateLimiter limits an authenticated user's calls to an expensive
// endpoint (e.g. AI summaries). Must run after RequireAuth.
func NewUserRateLimiter(perMinute, burst int) func(http.Handler) http.Handler {
return newRateLimiter(perMinute, burst).middlewareBy(func(r *http.Request) string {
if u := currentUser(r); u != nil {
return u.ID.String()
}
return getClientIP(r)
})
}
+3 -12
View File
@@ -1,7 +1,6 @@
package handler
import (
"log"
"net/http"
"github.com/michael/flowreader/internal/service"
@@ -22,20 +21,12 @@ func NewWSHandler(hub *ws.Hub, authService *service.AuthService) *WSHandler {
}
}
// Connect handles WebSocket initiation.
// Connect handles WebSocket initiation (behind RequireAuth).
func (h *WSHandler) Connect(w http.ResponseWriter, r *http.Request) {
cookie, err := r.Cookie("session_id")
if err != nil {
user := currentUser(r)
if user == nil {
http.Error(w, "Unauthorized", http.StatusUnauthorized)
return
}
user, err := h.authService.GetUserByToken(cookie.Value)
if err != nil || user == nil {
http.Error(w, "Unauthorized", http.StatusUnauthorized)
return
}
log.Printf("Setting up WS for user %s", user.ID)
h.hub.ServeWS(user.ID, w, r)
}