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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
Reference in new issue
Block a user