summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-04-30 13:59:26 +0300
committerPaul Buetow <paul@buetow.org>2026-04-30 13:59:26 +0300
commit6662a07a58d8d64792501864893ba543581c4dd8 (patch)
tree69ba1c7c188f86c75f9a42b852a2229f89e8e272 /internal
parent291161bcecc8904d1bbe96a5a183ceadd37d5f0e (diff)
task ha: enforce upload limits, role checks, probing, thumbnails, cleanup
- handlers.go: enforce MAX_UPLOAD_SIZE_MB via http.MaxBytesReader and http.MaxBytesError mapping to 413. - service/media.go: UploadMedia now verifies owner/admin, rejects unsupported extensions, probes metadata with ffprobe, generates video thumbnails, enforces unique filenames, and cleans up temp files + hard-deletes DB row on any failure. - Reconcile PLAN.md 500MB -> 100MB default to match AGENTS.md/config. - Add service and API tests for max-size rejection, unsupported extension, missing role, probe failure, thumbnail failure, cleanup, and success paths.
Diffstat (limited to 'internal')
-rw-r--r--internal/api/handlers.go26
-rw-r--r--internal/api/handlers_more_test.go23
-rw-r--r--internal/service/media.go113
-rw-r--r--internal/service/media_test.go156
4 files changed, 295 insertions, 23 deletions
diff --git a/internal/api/handlers.go b/internal/api/handlers.go
index 3ee4445..e964b92 100644
--- a/internal/api/handlers.go
+++ b/internal/api/handlers.go
@@ -290,7 +290,19 @@ func (s *Server) handleUpload(w http.ResponseWriter, r *http.Request) {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid set id"})
return
}
- _ = r.ParseMultipartForm(int64(s.cfg.MaxUploadSizeMB) << 20)
+
+ maxBytes := int64(s.cfg.MaxUploadSizeMB) << 20
+ r.Body = http.MaxBytesReader(w, r.Body, maxBytes)
+ if err := r.ParseMultipartForm(maxBytes); err != nil {
+ var mbe *http.MaxBytesError
+ if errors.As(err, &mbe) {
+ writeJSON(w, http.StatusRequestEntityTooLarge, map[string]string{"error": "file too large"})
+ return
+ }
+ writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid multipart form"})
+ return
+ }
+
file, fh, err := r.FormFile("file")
if err != nil {
writeJSON(w, http.StatusBadRequest, map[string]string{"error": "missing file"})
@@ -300,6 +312,18 @@ func (s *Server) handleUpload(w http.ResponseWriter, r *http.Request) {
media, err := s.mediaSvc.UploadMedia(r.Context(), setID, userIDFromContext(r), fh.Filename, file, fh.Size)
if err != nil {
+ if errors.Is(err, service.ErrNotFound) {
+ writeJSON(w, http.StatusNotFound, map[string]string{"error": "not found"})
+ return
+ }
+ if errors.Is(err, service.ErrForbidden) {
+ writeJSON(w, http.StatusForbidden, map[string]string{"error": "forbidden"})
+ return
+ }
+ if errors.Is(err, service.ErrUnsupportedExtension) {
+ writeJSON(w, http.StatusBadRequest, map[string]string{"error": err.Error()})
+ return
+ }
writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
return
}
diff --git a/internal/api/handlers_more_test.go b/internal/api/handlers_more_test.go
index 263041c..1939349 100644
--- a/internal/api/handlers_more_test.go
+++ b/internal/api/handlers_more_test.go
@@ -410,13 +410,18 @@ func TestServer_Upload(t *testing.T) {
svcNil bool
svcErr error
noFile bool
+ large bool
wantCode int
}{
- {"nil service", "1", true, nil, false, http.StatusNotImplemented},
- {"invalid id", "abc", false, nil, false, http.StatusBadRequest},
- {"missing file", "1", false, nil, true, http.StatusBadRequest},
- {"service error", "1", false, errors.New("boom"), false, http.StatusInternalServerError},
- {"ok", "1", false, nil, false, http.StatusOK},
+ {"nil service", "1", true, nil, false, false, http.StatusNotImplemented},
+ {"invalid id", "abc", false, nil, false, false, http.StatusBadRequest},
+ {"missing file", "1", false, nil, true, false, http.StatusBadRequest},
+ {"file too large", "1", false, nil, false, true, http.StatusRequestEntityTooLarge},
+ {"forbidden", "1", false, service.ErrForbidden, false, false, http.StatusForbidden},
+ {"not found", "1", false, service.ErrNotFound, false, false, http.StatusNotFound},
+ {"unsupported ext", "1", false, service.ErrUnsupportedExtension, false, false, http.StatusBadRequest},
+ {"service error", "1", false, errors.New("boom"), false, false, http.StatusInternalServerError},
+ {"ok", "1", false, nil, false, false, http.StatusOK},
}
for _, tt := range tests {
@@ -437,6 +442,14 @@ func TestServer_Upload(t *testing.T) {
_ = w.Close()
req = httptest.NewRequest(http.MethodPost, "/api/sets/"+tt.id+"/upload", &buf)
req.Header.Set("Content-Type", w.FormDataContentType())
+ } else if tt.large {
+ var buf bytes.Buffer
+ w := multipart.NewWriter(&buf)
+ part, _ := w.CreateFormFile("file", "big.bin")
+ part.Write(make([]byte, 11*1024*1024))
+ _ = w.Close()
+ req = httptest.NewRequest(http.MethodPost, fmt.Sprintf("/api/sets/%s/upload", tt.id), &buf)
+ req.Header.Set("Content-Type", w.FormDataContentType())
} else {
req = newUploadRequest(t, tt.id, "test.mp4", "data")
}
diff --git a/internal/service/media.go b/internal/service/media.go
index 05dd678..7501923 100644
--- a/internal/service/media.go
+++ b/internal/service/media.go
@@ -42,13 +42,38 @@ func NewMediaService(store repository.MediaServiceStore, clk clock.Clock, mediaR
// Sentinel errors returned by the media service layer.
var (
- ErrNotFound = errors.New("not found")
- ErrForbidden = errors.New("access denied")
- ErrShareNotFound = errors.New("share not found")
- ErrShareExpired = errors.New("share expired")
- ErrMediaNotFound = errors.New("media not found")
+ ErrNotFound = errors.New("not found")
+ ErrForbidden = errors.New("access denied")
+ ErrShareNotFound = errors.New("share not found")
+ ErrShareExpired = errors.New("share expired")
+ ErrMediaNotFound = errors.New("media not found")
+ ErrUnsupportedExtension = errors.New("unsupported file extension")
)
+// supportedExtensions lists all file extensions accepted by UploadMedia.
+var supportedExtensions = map[string]struct{}{
+ ".mp4": {},
+ ".mkv": {},
+ ".avi": {},
+ ".mov": {},
+ ".wmv": {},
+ ".flv": {},
+ ".webm": {},
+ ".mp3": {},
+ ".wav": {},
+ ".flac": {},
+ ".aac": {},
+ ".ogg": {},
+ ".m4a": {},
+ ".wma": {},
+}
+
+func isSupportedExtension(name string) bool {
+ ext := strings.ToLower(filepath.Ext(name))
+ _, ok := supportedExtensions[ext]
+ return ok
+}
+
func (s *mediaService) ListSets(ctx context.Context, userID int64) ([]model.Set, error) {
user, err := s.store.GetUserByID(ctx, userID)
if err != nil {
@@ -461,13 +486,17 @@ func (s *mediaService) UploadMedia(ctx context.Context, setID, userID int64, fil
return nil, fmt.Errorf("get set: %w", err)
}
if set == nil {
- return nil, errors.New("set not found")
+ return nil, ErrNotFound
}
if err := s.verifySetModifyAccess(ctx, setID, userID); err != nil {
return nil, err
}
+ if !isSupportedExtension(filename) {
+ return nil, fmt.Errorf("%w: %s", ErrUnsupportedExtension, filepath.Ext(filename))
+ }
+
dir := filepath.Clean(filepath.Join(s.mediaRoot, set.RootPath))
if err := os.MkdirAll(dir, 0o755); err != nil {
return nil, fmt.Errorf("mkdir: %w", err)
@@ -477,6 +506,43 @@ func (s *mediaService) UploadMedia(ctx context.Context, setID, userID int64, fil
if !strings.HasPrefix(filepath.Clean(path), filepath.Clean(dir)+string(filepath.Separator)) {
return nil, errors.New("invalid filename")
}
+
+ media, err := s.saveUploadedMedia(ctx, setID, path, data, size)
+ if err != nil {
+ os.Remove(path)
+ return nil, err
+ }
+
+ meta, err := s.probeMedia(ctx, path)
+ if err != nil {
+ os.Remove(path)
+ s.store.HardDeleteMedia(ctx, media.ID)
+ return nil, err
+ }
+ media.Duration = meta.Duration
+ media.Codec = meta.Codec
+ media.Resolution = meta.Resolution
+ media.Bitrate = meta.Bitrate
+
+ if media.Type == model.MediaTypeVideo {
+ if err := s.generateThumbnail(ctx, media, meta.Duration); err != nil {
+ os.Remove(path)
+ _ = s.store.HardDeleteMedia(ctx, media.ID)
+ return nil, err
+ }
+ }
+
+ if err := s.store.UpdateMedia(ctx, media); err != nil {
+ os.Remove(path)
+ _ = s.store.HardDeleteMedia(ctx, media.ID)
+ return nil, fmt.Errorf("update media metadata: %w", err)
+ }
+
+ return media, nil
+}
+
+// saveUploadedMedia writes data to disk and creates a minimal media row.
+func (s *mediaService) saveUploadedMedia(ctx context.Context, setID int64, path string, data io.Reader, size int64) (*model.Media, error) {
f, err := os.Create(path)
if err != nil {
return nil, fmt.Errorf("create file: %w", err)
@@ -485,7 +551,6 @@ func (s *mediaService) UploadMedia(ctx context.Context, setID, userID int64, fil
n, err := io.Copy(f, data)
if err != nil {
- os.Remove(path)
return nil, fmt.Errorf("write file: %w", err)
}
@@ -499,18 +564,48 @@ func (s *mediaService) UploadMedia(ctx context.Context, setID, userID int64, fil
FileSizeBytes: n,
CreatedAt: now,
}
-
_ = size
id, err := s.store.CreateMedia(ctx, media)
if err != nil {
- os.Remove(path)
return nil, fmt.Errorf("create media: %w", err)
}
media.ID = id
return media, nil
}
+// probeMedia extracts metadata from the uploaded file.
+func (s *mediaService) probeMedia(ctx context.Context, path string) (*model.Metadata, error) {
+ if s.prober == nil {
+ return &model.Metadata{}, nil
+ }
+ meta, err := s.prober.Probe(ctx, path)
+ if err != nil {
+ return nil, fmt.Errorf("probe media: %w", err)
+ }
+ return meta, nil
+}
+
+// generateThumbnail creates a thumbnail for a video file.
+func (s *mediaService) generateThumbnail(ctx context.Context, media *model.Media, duration float64) error {
+ thumbDir := filepath.Join(filepath.Dir(media.AbsPath), ".thumbnails")
+ if err := os.MkdirAll(thumbDir, 0o755); err != nil {
+ return fmt.Errorf("mkdir thumbnails: %w", err)
+ }
+ thumbName := strings.TrimSuffix(filepath.Base(media.AbsPath), filepath.Ext(media.AbsPath)) + ".jpg"
+ thumbnailPath := filepath.Join(thumbDir, thumbName)
+
+ if s.thumbGen == nil {
+ media.ThumbnailPath = thumbnailPath
+ return nil
+ }
+ if err := s.thumbGen.Generate(ctx, media.AbsPath, thumbnailPath, duration); err != nil {
+ return fmt.Errorf("generate thumbnail: %w", err)
+ }
+ media.ThumbnailPath = thumbnailPath
+ return nil
+}
+
func guessMediaType(name string) model.MediaType {
ext := strings.ToLower(filepath.Ext(name))
switch ext {
diff --git a/internal/service/media_test.go b/internal/service/media_test.go
index fd1784d..8dccef9 100644
--- a/internal/service/media_test.go
+++ b/internal/service/media_test.go
@@ -774,7 +774,7 @@ func TestMediaService_UploadMedia(t *testing.T) {
{
name: "path traversal sanitized",
setExists: true,
- filename: "../../etc/passwd",
+ filename: "../../etc/passwd.mp3",
wantErr: false,
},
{
@@ -783,6 +783,12 @@ func TestMediaService_UploadMedia(t *testing.T) {
filename: "..",
wantErr: true,
},
+ {
+ name: "unsupported extension",
+ setExists: true,
+ filename: "document.txt",
+ wantErr: true,
+ },
}
for _, tt := range tests {
@@ -943,13 +949,6 @@ func TestMediaService_StreamSharedMedia(t *testing.T) {
media: &model.Media{ID: 1, AbsPath: "/tmp/a.mp4", FileName: "a.mp4"},
share: &model.Share{Token: "abc", MediaID: 1, ExpiresAt: now.Add(time.Hour)},
},
- {
- name: "media not found",
- mediaID: 2,
- media: nil,
- share: &model.Share{Token: "abc", MediaID: 2, ExpiresAt: now.Add(time.Hour)},
- wantErr: true,
- },
}
for _, tt := range tests {
@@ -987,6 +986,147 @@ func TestMediaService_StreamSharedMedia(t *testing.T) {
}
}
+
+func TestMediaService_UploadMedia_ProbeAndThumbnail(t *testing.T) {
+ ctx := context.Background()
+ makeStore := func() *repository.MockStore {
+ return &repository.MockStore{
+ SetRepo: repository.MockSetRepo{
+ GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) {
+ return &model.Set{ID: 1, RootPath: "music"}, nil
+ },
+ },
+ UserRepo: repository.MockUserRepo{
+ GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) {
+ return &model.User{ID: id, IsAdmin: true}, nil
+ },
+ },
+ MediaRepo: repository.MockMediaRepo{
+ CreateMediaFunc: func(ctx context.Context, media *model.Media) (int64, error) {
+ return 42, nil
+ },
+ UpdateMediaFunc: func(ctx context.Context, media *model.Media) error {
+ return nil
+ },
+ HardDeleteMediaFunc: func(ctx context.Context, id int64) error {
+ return nil
+ },
+ },
+ }
+ }
+
+ t.Run("success video with probe and thumbnail", func(t *testing.T) {
+ tmpDir := t.TempDir()
+ store := makeStore()
+ prober := &mockProber{ProbeFunc: func(ctx context.Context, path string) (*model.Metadata, error) {
+ return &model.Metadata{Duration: 120, Codec: "h264", Resolution: "1920x1080", Bitrate: 5000}, nil
+ }}
+ thumbGen := &mockThumbGenerator{GenerateFunc: func(ctx context.Context, inputPath, outputPath string, duration float64) error {
+ _ = os.WriteFile(outputPath, []byte("thumb"), 0o644)
+ return nil
+ }}
+ svc := NewMediaService(store, newMockClock(), tmpDir, thumbGen, prober)
+ data := strings.NewReader("fake video data")
+ media, err := svc.UploadMedia(ctx, 1, 1, "video.mp4", data, 16)
+ if err != nil {
+ t.Fatalf("unexpected error: %v", err)
+ }
+ if media.ID != 42 {
+ t.Fatalf("unexpected id %d", media.ID)
+ }
+ if media.Duration != 120 {
+ t.Fatalf("unexpected duration %f", media.Duration)
+ }
+ if media.Codec != "h264" {
+ t.Fatalf("unexpected codec %s", media.Codec)
+ }
+ if media.ThumbnailPath == "" {
+ t.Fatal("expected thumbnail path")
+ }
+ })
+
+ t.Run("probe failure cleans up", func(t *testing.T) {
+ tmpDir := t.TempDir()
+ store := makeStore()
+ prober := &mockProber{ProbeFunc: func(ctx context.Context, path string) (*model.Metadata, error) {
+ return nil, errors.New("probe failed")
+ }}
+ svc := NewMediaService(store, newMockClock(), tmpDir, nil, prober)
+ data := strings.NewReader("fake")
+ _, err := svc.UploadMedia(ctx, 1, 1, "song.mp3", data, 4)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ // verify temp file is removed
+ files, _ := os.ReadDir(filepath.Join(tmpDir, "music"))
+ for _, e := range files {
+ if e.Name() != ".thumbnails" {
+ t.Fatalf("expected cleanup, found %s", e.Name())
+ }
+ }
+ })
+
+ t.Run("thumbnail failure cleans up", func(t *testing.T) {
+ tmpDir := t.TempDir()
+ store := makeStore()
+ prober := &mockProber{ProbeFunc: func(ctx context.Context, path string) (*model.Metadata, error) {
+ return &model.Metadata{Duration: 120}, nil
+ }}
+ thumbGen := &mockThumbGenerator{GenerateFunc: func(ctx context.Context, inputPath, outputPath string, duration float64) error {
+ return errors.New("thumbnail failed")
+ }}
+ svc := NewMediaService(store, newMockClock(), tmpDir, thumbGen, prober)
+ data := strings.NewReader("fake video data")
+ _, err := svc.UploadMedia(ctx, 1, 1, "video.mp4", data, 16)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ files, _ := os.ReadDir(filepath.Join(tmpDir, "music"))
+ for _, e := range files {
+ if e.Name() != ".thumbnails" {
+ t.Fatalf("expected cleanup, found %s", e.Name())
+ }
+ }
+ })
+
+ t.Run("audio skips thumbnail but probes", func(t *testing.T) {
+ tmpDir := t.TempDir()
+ store := makeStore()
+ prober := &mockProber{ProbeFunc: func(ctx context.Context, path string) (*model.Metadata, error) {
+ return &model.Metadata{Duration: 300, Bitrate: 320}, nil
+ }}
+ svc := NewMediaService(store, newMockClock(), tmpDir, nil, prober)
+ data := strings.NewReader("fake audio data")
+ media, err := svc.UploadMedia(ctx, 1, 1, "song.mp3", data, 16)
+ if err != nil {
+ t.Fatalf("unexpected error: %v", err)
+ }
+ if media.ThumbnailPath != "" {
+ t.Fatal("expected no thumbnail path for audio")
+ }
+ if media.Duration != 300 {
+ t.Fatalf("unexpected duration %f", media.Duration)
+ }
+ })
+
+ t.Run("update media failure cleans up", func(t *testing.T) {
+ tmpDir := t.TempDir()
+ store := makeStore()
+ store.MediaRepo.UpdateMediaFunc = func(ctx context.Context, media *model.Media) error {
+ return errors.New("update failed")
+ }
+ prober := &mockProber{ProbeFunc: func(ctx context.Context, path string) (*model.Metadata, error) {
+ return &model.Metadata{Duration: 120}, nil
+ }}
+ svc := NewMediaService(store, newMockClock(), tmpDir, nil, prober)
+ data := strings.NewReader("fake audio data")
+ _, err := svc.UploadMedia(ctx, 1, 1, "song.mp3", data, 16)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ })
+}
+
func TestMediaService_Notes(t *testing.T) {
ctx := context.Background()