diff options
| author | Paul Buetow <paul@buetow.org> | 2026-04-30 13:59:26 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-04-30 13:59:26 +0300 |
| commit | 6662a07a58d8d64792501864893ba543581c4dd8 (patch) | |
| tree | 69ba1c7c188f86c75f9a42b852a2229f89e8e272 /internal | |
| parent | 291161bcecc8904d1bbe96a5a183ceadd37d5f0e (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.go | 26 | ||||
| -rw-r--r-- | internal/api/handlers_more_test.go | 23 | ||||
| -rw-r--r-- | internal/service/media.go | 113 | ||||
| -rw-r--r-- | internal/service/media_test.go | 156 |
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() |
