Article
%s

diff --git a/internal/cache/media_cache.go b/internal/cache/media_cache.go
index 584dbea7e..5ba32e023 100644
--- a/internal/cache/media_cache.go
+++ b/internal/cache/media_cache.go
@@ -5,6 +5,7 @@ import (
"context"
"crypto/sha256"
"encoding/hex"
+ "errors"
"fmt"
"io"
"net/http"
@@ -14,6 +15,8 @@ import (
"sort"
"strings"
"time"
+
+ "MrRSS/internal/utils/httputil"
)
// MediaCache handles caching of images and videos to work around anti-hotlinking
@@ -21,6 +24,16 @@ type MediaCache struct {
cacheDir string
}
+// MediaDownloadError distinguishes a failed upstream fetch from a cache storage
+// error, so callers do not repeat an already exhausted Referer fallback.
+type MediaDownloadError struct {
+ Err error
+ RefererFallbackAttempted bool
+}
+
+func (e *MediaDownloadError) Error() string { return e.Err.Error() }
+func (e *MediaDownloadError) Unwrap() error { return e.Err }
+
// NewMediaCache creates a new media cache instance
func NewMediaCache(cacheDir string) (*MediaCache, error) {
// Create cache directory if it doesn't exist
@@ -90,6 +103,7 @@ func (mc *MediaCache) Get(ctx context.Context, client *http.Client, url, referer
if err != nil {
return nil, "", fmt.Errorf("failed to download media: %w", err)
}
+ cachedPath = mc.GetCachedPath(url)
// Determine better file extension from Content-Type if available
if contentType != "" {
@@ -100,8 +114,24 @@ func (mc *MediaCache) Get(ctx context.Context, client *http.Client, url, referer
}
}
- // Save to cache
- if err := os.WriteFile(cachedPath, data, 0644); err != nil {
+ // Publish only a complete file, just as DownloadToFile does.
+ tmpFile, err := os.CreateTemp(mc.cacheDir, "download-*")
+ if err != nil {
+ return nil, "", fmt.Errorf("failed to create temporary cache file: %w", err)
+ }
+ defer os.Remove(tmpFile.Name())
+ _, writeErr := tmpFile.Write(data)
+ closeErr := tmpFile.Close()
+ if writeErr != nil {
+ return nil, "", fmt.Errorf("failed to cache media: %w", writeErr)
+ }
+ if closeErr != nil {
+ return nil, "", fmt.Errorf("failed to finish cache file: %w", closeErr)
+ }
+ if err := ctx.Err(); err != nil {
+ return nil, "", err
+ }
+ if err := os.Rename(tmpFile.Name(), cachedPath); err != nil {
return nil, "", fmt.Errorf("failed to cache media: %w", err)
}
@@ -140,14 +170,14 @@ func (mc *MediaCache) DownloadToFile(ctx context.Context, client *http.Client, u
return "", "", err
}
- resp, err := client.Do(req)
+ resp, retried, err := httputil.DoWithRefererFallback(client, req)
if err != nil {
- return "", "", fmt.Errorf("failed to fetch media: %w", err)
+ return "", "", &MediaDownloadError{fmt.Errorf("failed to fetch media: %w", err), retried}
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
- return "", "", fmt.Errorf("unexpected status code: %d", resp.StatusCode)
+ return "", "", &MediaDownloadError{fmt.Errorf("unexpected status code: %d", resp.StatusCode), retried}
}
tmpFile, err := os.CreateTemp(mc.cacheDir, "download-*")
@@ -155,18 +185,25 @@ func (mc *MediaCache) DownloadToFile(ctx context.Context, client *http.Client, u
return "", "", fmt.Errorf("failed to create temporary cache file: %w", err)
}
tmpPath := tmpFile.Name()
+ defer os.Remove(tmpPath)
written, copyErr := io.Copy(tmpFile, resp.Body)
closeErr := tmpFile.Close()
- if copyErr != nil || closeErr != nil || written == 0 {
- _ = os.Remove(tmpPath)
- if copyErr != nil {
- return "", "", fmt.Errorf("failed to stream media: %w", copyErr)
+ if copyErr != nil {
+ var pathErr *os.PathError
+ if errors.As(copyErr, &pathErr) {
+ return "", "", fmt.Errorf("failed to write cache file: %w", copyErr)
}
- if closeErr != nil {
- return "", "", fmt.Errorf("failed to finish cache file: %w", closeErr)
- }
- return "", "", fmt.Errorf("empty media response")
+ return "", "", &MediaDownloadError{fmt.Errorf("failed to stream media: %w", copyErr), retried}
+ }
+ if closeErr != nil {
+ return "", "", fmt.Errorf("failed to finish cache file: %w", closeErr)
+ }
+ if err := ctx.Err(); err != nil {
+ return "", "", &MediaDownloadError{err, retried}
+ }
+ if written == 0 || (resp.ContentLength >= 0 && written != resp.ContentLength) {
+ return "", "", &MediaDownloadError{fmt.Errorf("incomplete media response: received %d bytes, expected %d", written, resp.ContentLength), retried}
}
contentType := resp.Header.Get("Content-Type")
@@ -180,7 +217,6 @@ func (mc *MediaCache) DownloadToFile(ctx context.Context, client *http.Client, u
}
if err := os.Rename(tmpPath, finalPath); err != nil {
- _ = os.Remove(tmpPath)
return "", "", fmt.Errorf("failed to store media in cache: %w", err)
}
@@ -229,19 +265,25 @@ func (mc *MediaCache) download(ctx context.Context, client *http.Client, url, re
return nil, "", err
}
- resp, err := client.Do(req)
+ resp, retried, err := httputil.DoWithRefererFallback(client, req)
if err != nil {
- return nil, "", fmt.Errorf("failed to fetch media: %w", err)
+ return nil, "", &MediaDownloadError{fmt.Errorf("failed to fetch media: %w", err), retried}
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
- return nil, "", fmt.Errorf("unexpected status code: %d", resp.StatusCode)
+ return nil, "", &MediaDownloadError{fmt.Errorf("unexpected status code: %d", resp.StatusCode), retried}
}
data, err := io.ReadAll(resp.Body)
if err != nil {
- return nil, "", fmt.Errorf("failed to read response body: %w", err)
+ return nil, "", &MediaDownloadError{fmt.Errorf("failed to read response body: %w", err), retried}
+ }
+ if err := ctx.Err(); err != nil {
+ return nil, "", &MediaDownloadError{err, retried}
+ }
+ if len(data) == 0 || (resp.ContentLength >= 0 && int64(len(data)) != resp.ContentLength) {
+ return nil, "", &MediaDownloadError{fmt.Errorf("incomplete media response: received %d bytes, expected %d", len(data), resp.ContentLength), retried}
}
contentType := resp.Header.Get("Content-Type")
diff --git a/internal/cache/media_cache_test.go b/internal/cache/media_cache_test.go
index ea16ca7a0..d11dbd74d 100644
--- a/internal/cache/media_cache_test.go
+++ b/internal/cache/media_cache_test.go
@@ -3,15 +3,189 @@ package cache
import (
"context"
"errors"
+ "fmt"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
+ "strings"
+ "sync/atomic"
"testing"
"time"
+
+ "MrRSS/internal/utils/httputil"
)
+func TestMediaCacheRefererFallback(t *testing.T) {
+ for _, inMemory := range []bool{false, true} {
+ for _, finalStatus := range []int{http.StatusOK, http.StatusForbidden, http.StatusNotFound} {
+ name := fmt.Sprintf("memory=%v/status=%d", inMemory, finalStatus)
+ t.Run(name, func(t *testing.T) {
+ dir := t.TempDir()
+ mc, err := NewMediaCache(dir)
+ if err != nil {
+ t.Fatal(err)
+ }
+ var calls atomic.Int32
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ calls.Add(1)
+ if r.Header.Get("Referer") != "" && finalStatus != http.StatusNotFound {
+ w.WriteHeader(http.StatusForbidden)
+ return
+ }
+ w.Header().Set("Content-Type", "image/png")
+ w.WriteHeader(finalStatus)
+ _, _ = io.WriteString(w, "image")
+ }))
+ defer server.Close()
+ mediaURL := server.URL + "/image.png"
+ if inMemory {
+ _, _, err = mc.Get(context.Background(), server.Client(), mediaURL, "https://article.example")
+ } else {
+ _, _, err = mc.DownloadToFile(context.Background(), server.Client(), mediaURL, "https://article.example")
+ }
+ wantCalls := int32(2)
+ if finalStatus == http.StatusNotFound {
+ wantCalls = 1
+ }
+ if calls.Load() != wantCalls {
+ t.Fatalf("requests=%d want=%d", calls.Load(), wantCalls)
+ }
+ if finalStatus == http.StatusOK {
+ if err != nil {
+ t.Fatal(err)
+ }
+ data, contentType, err := mc.Get(context.Background(), server.Client(), mediaURL, "https://article.example")
+ if err != nil || string(data) != "image" || contentType != "image/png" || calls.Load() != wantCalls {
+ t.Fatalf("cache hit: data=%q type=%q err=%v requests=%d", data, contentType, err, calls.Load())
+ }
+ } else {
+ var downloadErr *MediaDownloadError
+ if !errors.As(err, &downloadErr) || downloadErr.RefererFallbackAttempted != (wantCalls == 2) {
+ t.Fatalf("download error lost retry state: %v", err)
+ }
+ entries, err := os.ReadDir(dir)
+ if err != nil || len(entries) != 0 {
+ t.Fatalf("failed download left files: %v, %v", entries, err)
+ }
+ }
+ })
+ }
+ }
+}
+
+type failingMediaBody struct {
+ data string
+ err error
+ cancel context.CancelFunc
+ closed bool
+}
+
+func (b *failingMediaBody) Read(p []byte) (int, error) {
+ if b.data != "" {
+ n := copy(p, b.data)
+ b.data = b.data[n:]
+ return n, nil
+ }
+ if b.cancel != nil {
+ b.cancel()
+ }
+ return 0, b.err
+}
+
+func (b *failingMediaBody) Close() error {
+ b.closed = true
+ return nil
+}
+
+func TestMediaCacheIncompleteOrCancelledBodyNeverCached(t *testing.T) {
+ for _, inMemory := range []bool{false, true} {
+ for _, mode := range []string{"short", "read failure", "cancelled", "empty"} {
+ t.Run(fmt.Sprintf("memory=%v/%s", inMemory, mode), func(t *testing.T) {
+ dir := t.TempDir()
+ mc, err := NewMediaCache(dir)
+ if err != nil {
+ t.Fatal(err)
+ }
+ ctx, cancel := context.WithCancel(context.Background())
+ defer cancel()
+ body := &failingMediaBody{data: "partial", err: io.EOF}
+ length := int64(10)
+ if mode == "read failure" {
+ body.err = io.ErrUnexpectedEOF
+ } else if mode == "cancelled" {
+ body.cancel = cancel
+ length = -1
+ } else if mode == "empty" {
+ body.data = ""
+ length = 0
+ }
+ calls := 0
+ client := &http.Client{Transport: httputil.RoundTripFunc(func(req *http.Request) (*http.Response, error) {
+ calls++
+ if calls == 1 {
+ return &http.Response{StatusCode: 403, Header: make(http.Header), Body: io.NopCloser(strings.NewReader("forbidden")), Request: req}, nil
+ }
+ return &http.Response{StatusCode: 200, Header: http.Header{"Content-Type": {"image/png"}}, Body: body, ContentLength: length, Request: req}, nil
+ })}
+ if inMemory {
+ _, _, err = mc.Get(ctx, client, "https://media.example/image.png", "https://article.example")
+ } else {
+ _, _, err = mc.DownloadToFile(ctx, client, "https://media.example/image.png", "https://article.example")
+ }
+ var downloadErr *MediaDownloadError
+ if !errors.As(err, &downloadErr) || !downloadErr.RefererFallbackAttempted || calls != 2 {
+ t.Fatalf("err=%v calls=%d", err, calls)
+ }
+ if mode == "cancelled" && !errors.Is(err, context.Canceled) {
+ t.Fatalf("cancellation lost: %v", err)
+ }
+ if !body.closed {
+ t.Fatal("failed response body was not closed")
+ }
+ entries, err := os.ReadDir(dir)
+ if err != nil || len(entries) != 0 {
+ t.Fatalf("incomplete download left files: %v, %v", entries, err)
+ }
+ })
+ }
+ }
+}
+
+func TestMediaCacheStorageErrorAllowsDirectFallback(t *testing.T) {
+ for _, inMemory := range []bool{false, true} {
+ t.Run(fmt.Sprintf("memory=%v", inMemory), func(t *testing.T) {
+ // A file in place of the cache directory makes cache storage fail
+ // deterministically, including when tests run as a privileged user.
+ cachePath := filepath.Join(t.TempDir(), "not-a-directory")
+ if err := os.WriteFile(cachePath, []byte("occupied"), 0600); err != nil {
+ t.Fatal(err)
+ }
+ mc := &MediaCache{cacheDir: cachePath}
+ calls := 0
+ client := &http.Client{Transport: httputil.RoundTripFunc(func(req *http.Request) (*http.Response, error) {
+ calls++
+ code := 200
+ if calls == 1 {
+ code = 403
+ }
+ return &http.Response{StatusCode: code, Header: http.Header{"Content-Type": {"image/png"}}, Body: io.NopCloser(strings.NewReader("image")), ContentLength: 5, Request: req}, nil
+ })}
+ var err error
+ if inMemory {
+ _, _, err = mc.Get(context.Background(), client, "https://media.example/image.png", "https://article.example")
+ } else {
+ _, _, err = mc.DownloadToFile(context.Background(), client, "https://media.example/image.png", "https://article.example")
+ }
+ var downloadErr *MediaDownloadError
+ if err == nil || errors.As(err, &downloadErr) || calls != 2 {
+ t.Fatalf("storage error incorrectly prevents direct fallback: err=%v calls=%d", err, calls)
+ }
+ })
+ }
+}
+
func TestMediaCacheCancelledDownloadDoesNotWriteFile(t *testing.T) {
mc, err := NewMediaCache(t.TempDir())
if err != nil {
diff --git a/internal/handlers/core/full_text.go b/internal/handlers/core/full_text.go
index ad7463a33..2395aa768 100644
--- a/internal/handlers/core/full_text.go
+++ b/internal/handlers/core/full_text.go
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"fmt"
+ "html"
"io"
"net/http"
"net/url"
@@ -120,7 +121,9 @@ func (h *Handler) FetchFullArticleContentContext(ctx context.Context, articleURL
}
extracted, extractErr := readability.FromReader(strings.NewReader(page), base)
var output bytes.Buffer
+ leadImageURL := ""
if extractErr == nil {
+ leadImageURL = extracted.ImageURL()
extractErr = extracted.RenderHTML(&output)
}
content := output.String()
@@ -128,6 +131,12 @@ func (h *Handler) FetchFullArticleContentContext(ctx context.Context, articleURL
// Explicit semantic article containers are a useful fallback for short pages.
content, _ = doc.Find("article,main,[role=main]").First().Html()
}
+ if strings.TrimSpace(leadImageURL) != "" && !strings.Contains(strings.ToLower(content), "
%s

%s
%s
%s

