Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 60 additions & 18 deletions internal/cache/media_cache.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"net/http"
Expand All @@ -14,13 +15,25 @@ import (
"sort"
"strings"
"time"

"MrRSS/internal/utils/httputil"
)

// MediaCache handles caching of images and videos to work around anti-hotlinking
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
Expand Down Expand Up @@ -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 != "" {
Expand All @@ -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)
}

Expand Down Expand Up @@ -140,33 +170,40 @@ 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-*")
if err != nil {
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")
Expand All @@ -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)
}

Expand Down Expand Up @@ -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")
Expand Down
174 changes: 174 additions & 0 deletions internal/cache/media_cache_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
Loading
Loading