diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/api/handlers.go | 18 | ||||
| -rw-r--r-- | internal/api/handlers_test.go | 42 | ||||
| -rw-r--r-- | internal/repository/media.go | 6 | ||||
| -rw-r--r-- | internal/repository/repository.go | 21 | ||||
| -rw-r--r-- | internal/service/media.go | 134 | ||||
| -rw-r--r-- | internal/service/media_test.go | 482 | ||||
| -rw-r--r-- | internal/service/mock.go | 6 | ||||
| -rw-r--r-- | internal/service/no_rows_test.go | 34 | ||||
| -rw-r--r-- | internal/service/service.go | 2 |
9 files changed, 716 insertions, 29 deletions
diff --git a/internal/api/handlers.go b/internal/api/handlers.go index fc1d43d..af6a0a0 100644 --- a/internal/api/handlers.go +++ b/internal/api/handlers.go @@ -356,7 +356,7 @@ func (s *Server) handleListMedia(w http.ResponseWriter, r *http.Request) { return } filter := parseMediaListQuery(r.URL.Query()) - media, err := s.mediaSvc.ListMedia(r.Context(), filter) + media, err := s.mediaSvc.ListMedia(r.Context(), userIDFromContext(r), filter) if err != nil { writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) return @@ -452,6 +452,14 @@ func (s *Server) handleSoftDelete(w http.ResponseWriter, r *http.Request) { return } if err := s.mediaSvc.SoftDeleteMedia(r.Context(), id, userIDFromContext(r)); 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 + } writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) return } @@ -468,6 +476,14 @@ func (s *Server) handleRestore(w http.ResponseWriter, r *http.Request) { return } if err := s.mediaSvc.RestoreMedia(r.Context(), id, userIDFromContext(r)); 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 + } writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) return } diff --git a/internal/api/handlers_test.go b/internal/api/handlers_test.go index c853140..798ac81 100644 --- a/internal/api/handlers_test.go +++ b/internal/api/handlers_test.go @@ -587,7 +587,7 @@ func TestServer_MediaList(t *testing.T) { t.Run(tt.name, func(t *testing.T) { var gotFilter repository.MediaFilter ms := &service.MockMediaService{ - ListMediaFunc: func(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) { + ListMediaFunc: func(ctx context.Context, userID int64, filter repository.MediaFilter) ([]model.Media, error) { gotFilter = filter return tt.listResult, tt.listErr }, @@ -782,6 +782,46 @@ func TestServer_Restore(t *testing.T) { } } +func TestServer_SoftDelete_Forbidden(t *testing.T) { + ms := &service.MockMediaService{ + SoftDeleteMediaFunc: func(ctx context.Context, mediaID, userID int64) error { + return service.ErrForbidden + }, + } + store := buildSessionStore(1) + sm := auth.NewSessionManager(store, &clock.MockClock{T: time.Now()}, time.Hour) + cfg := &internal.Config{SessionTimeoutHours: 24} + srv := newTestServer(t, buildCountStore(1), nil, sm, cfg, ms, nil, nil, nil) + + req := httptest.NewRequest(http.MethodDelete, "/api/media/99", nil) + req.AddCookie(addSessionCookie(t, store, sm, 1)) + rr := httptest.NewRecorder() + srv.ServeHTTP(rr, req) + if rr.Code != http.StatusForbidden { + t.Fatalf("expected %d, got %d", http.StatusForbidden, rr.Code) + } +} + +func TestServer_Restore_Forbidden(t *testing.T) { + ms := &service.MockMediaService{ + RestoreMediaFunc: func(ctx context.Context, mediaID, userID int64) error { + return service.ErrForbidden + }, + } + store := buildSessionStore(1) + sm := auth.NewSessionManager(store, &clock.MockClock{T: time.Now()}, time.Hour) + cfg := &internal.Config{SessionTimeoutHours: 24} + srv := newTestServer(t, buildCountStore(1), nil, sm, cfg, ms, nil, nil, nil) + + req := httptest.NewRequest(http.MethodPost, "/api/media/99/restore", nil) + req.AddCookie(addSessionCookie(t, store, sm, 1)) + rr := httptest.NewRecorder() + srv.ServeHTTP(rr, req) + if rr.Code != http.StatusForbidden { + t.Fatalf("expected %d, got %d", http.StatusForbidden, rr.Code) + } +} + // ------------------------------------------------------------------ // Notes // ------------------------------------------------------------------ diff --git a/internal/repository/media.go b/internal/repository/media.go index 26d69ca..6c5297b 100644 --- a/internal/repository/media.go +++ b/internal/repository/media.go @@ -152,6 +152,12 @@ func (s *SQLite) ListMedia(ctx context.Context, filter MediaFilter) ([]model.Med conds = append(conds, `media.set_id = ?`) args = append(args, *filter.SetID) } + if len(filter.AllowedSetIDs) > 0 { + conds = append(conds, "media.set_id IN ("+placeholders(len(filter.AllowedSetIDs))+")") + for _, id := range filter.AllowedSetIDs { + args = append(args, id) + } + } if filter.Type != nil { conds = append(conds, `media.type = ?`) args = append(args, string(*filter.Type)) diff --git a/internal/repository/repository.go b/internal/repository/repository.go index 3de011c..227fe79 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -92,16 +92,17 @@ type SetPermissionRepo interface { // MediaFilter defines query parameters for listing media. type MediaFilter struct { - SetID *int64 - Type *model.MediaType - Search string - Tags []string - Favorites *int64 // userID if set - MinDuration *float64 - MaxDuration *float64 - Sort string // name, date, duration, play_count, random - Limit int - Offset int + SetID *int64 + AllowedSetIDs []int64 + Type *model.MediaType + Search string + Tags []string + Favorites *int64 // userID if set + MinDuration *float64 + MaxDuration *float64 + Sort string // name, date, duration, play_count, random + Limit int + Offset int } // MediaRepo manages media items. diff --git a/internal/service/media.go b/internal/service/media.go index 5cb7772..f950db3 100644 --- a/internal/service/media.go +++ b/internal/service/media.go @@ -82,12 +82,9 @@ func (s *mediaService) ListSets(ctx context.Context, userID int64) ([]model.Set, } func (s *mediaService) GetMediaDetail(ctx context.Context, mediaID, userID int64) (*MediaDetail, error) { - media, err := s.store.GetMediaByID(ctx, mediaID) + media, err := s.verifyAccess(ctx, mediaID, userID) if err != nil { - return nil, fmt.Errorf("get media: %w", err) - } - if media == nil { - return nil, nil + return nil, err } tags, err := s.store.ListTagsByMedia(ctx, mediaID) @@ -119,7 +116,26 @@ func (s *mediaService) GetMediaDetail(ctx context.Context, mediaID, userID int64 }, nil } -func (s *mediaService) ListMedia(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) { +func (s *mediaService) ListMedia(ctx context.Context, userID int64, filter repository.MediaFilter) ([]model.Media, error) { + user, err := s.store.GetUserByID(ctx, userID) + if err != nil { + return nil, fmt.Errorf("get user: %w", err) + } + + if user != nil && user.IsAdmin { + return s.store.ListMedia(ctx, filter) + } + + perms, err := s.store.ListPermissionsByUser(ctx, userID) + if err != nil { + return nil, fmt.Errorf("list permissions: %w", err) + } + + allowed := make([]int64, 0, len(perms)) + for _, p := range perms { + allowed = append(allowed, p.SetID) + } + filter.AllowedSetIDs = allowed return s.store.ListMedia(ctx, filter) } @@ -165,6 +181,77 @@ func (s *mediaService) verifyAccess(ctx context.Context, mediaID, userID int64) return nil, ErrForbidden } +// verifyModifyAccess checks that the user has access to the media and is an owner or admin. +func (s *mediaService) verifyModifyAccess(ctx context.Context, mediaID, userID int64) (*model.Media, error) { + media, err := s.verifyAccess(ctx, mediaID, userID) + if err != nil { + return nil, err + } + + user, err := s.store.GetUserByID(ctx, userID) + if err != nil { + return nil, fmt.Errorf("get user: %w", err) + } + if user != nil && user.IsAdmin { + return media, nil + } + + perm, err := s.store.GetPermission(ctx, media.SetID, userID) + if err != nil { + return nil, fmt.Errorf("get permission: %w", err) + } + if perm != nil && perm.Role == model.RoleOwner { + return media, nil + } + + set, err := s.store.GetSetByID(ctx, media.SetID) + if err != nil { + return nil, fmt.Errorf("get set: %w", err) + } + if set != nil { + for _, p := range set.Permissions { + if p.UserID == userID && p.Role == model.RoleOwner { + return media, nil + } + } + } + + return nil, ErrForbidden +} + +// verifySetModifyAccess checks that the user is an owner or admin for a set. +func (s *mediaService) verifySetModifyAccess(ctx context.Context, setID, userID int64) error { + user, err := s.store.GetUserByID(ctx, userID) + if err != nil { + return fmt.Errorf("get user: %w", err) + } + if user != nil && user.IsAdmin { + return nil + } + + perm, err := s.store.GetPermission(ctx, setID, userID) + if err != nil { + return fmt.Errorf("get permission: %w", err) + } + if perm != nil && perm.Role == model.RoleOwner { + return nil + } + + set, err := s.store.GetSetByID(ctx, setID) + if err != nil { + return fmt.Errorf("get set: %w", err) + } + if set != nil { + for _, p := range set.Permissions { + if p.UserID == userID && p.Role == model.RoleOwner { + return nil + } + } + } + + return ErrForbidden +} + func (s *mediaService) StreamMedia(ctx context.Context, mediaID, userID int64) (*FileResult, error) { media, err := s.verifyAccess(ctx, mediaID, userID) if err != nil { @@ -201,18 +288,31 @@ func (s *mediaService) GetThumbnail(ctx context.Context, mediaID, userID int64) } func (s *mediaService) RegenerateThumbnail(ctx context.Context, mediaID, userID int64) error { + _, err := s.verifyModifyAccess(ctx, mediaID, userID) + if err != nil { + return err + } return errors.New("not implemented") } func (s *mediaService) RegenerateSetCover(ctx context.Context, setID, userID int64) error { + if err := s.verifySetModifyAccess(ctx, setID, userID); err != nil { + return err + } return errors.New("not implemented") } func (s *mediaService) ToggleFavorite(ctx context.Context, userID, mediaID int64) (bool, error) { + if _, err := s.verifyAccess(ctx, mediaID, userID); err != nil { + return false, err + } return s.store.ToggleFavorite(ctx, userID, mediaID) } func (s *mediaService) AssignTag(ctx context.Context, mediaID, userID int64, tagName string) error { + if _, err := s.verifyAccess(ctx, mediaID, userID); err != nil { + return err + } tag, err := s.store.GetTagByName(ctx, tagName) if err != nil { return fmt.Errorf("get tag: %w", err) @@ -228,6 +328,9 @@ func (s *mediaService) AssignTag(ctx context.Context, mediaID, userID int64, tag } func (s *mediaService) RemoveTag(ctx context.Context, mediaID, userID int64, tagName string) error { + if _, err := s.verifyAccess(ctx, mediaID, userID); err != nil { + return err + } tag, err := s.store.GetTagByName(ctx, tagName) if err != nil { return fmt.Errorf("get tag: %w", err) @@ -239,7 +342,7 @@ func (s *mediaService) RemoveTag(ctx context.Context, mediaID, userID int64, tag } func (s *mediaService) SoftDeleteMedia(ctx context.Context, mediaID, userID int64) error { - _, err := s.verifyAccess(ctx, mediaID, userID) + _, err := s.verifyModifyAccess(ctx, mediaID, userID) if err != nil { return err } @@ -247,6 +350,10 @@ func (s *mediaService) SoftDeleteMedia(ctx context.Context, mediaID, userID int6 } func (s *mediaService) RestoreMedia(ctx context.Context, mediaID, userID int64) error { + _, err := s.verifyModifyAccess(ctx, mediaID, userID) + if err != nil { + return err + } return s.store.RestoreMedia(ctx, mediaID) } @@ -280,6 +387,10 @@ func (s *mediaService) UploadMedia(ctx context.Context, setID, userID int64, fil return nil, errors.New("set not found") } + if err := s.verifySetModifyAccess(ctx, setID, userID); err != nil { + return nil, err + } + dir := filepath.Clean(filepath.Join(s.mediaRoot, set.RootPath)) if err := os.MkdirAll(dir, 0o755); err != nil { return nil, fmt.Errorf("mkdir: %w", err) @@ -441,14 +552,23 @@ func (s *mediaService) StreamSharedMedia(ctx context.Context, token string) (*Fi } func (s *mediaService) GetNote(ctx context.Context, mediaID, userID int64) (*model.Note, error) { + if _, err := s.verifyAccess(ctx, mediaID, userID); err != nil { + return nil, err + } return s.store.GetNote(ctx, mediaID, userID) } func (s *mediaService) UpsertNote(ctx context.Context, note *model.Note) error { + if _, err := s.verifyAccess(ctx, note.MediaID, note.UserID); err != nil { + return err + } note.UpdatedAt = s.clock.Now() return s.store.UpsertNote(ctx, note) } func (s *mediaService) DeleteNote(ctx context.Context, mediaID, userID int64) error { + if _, err := s.verifyAccess(ctx, mediaID, userID); err != nil { + return err + } return s.store.DeleteNote(ctx, mediaID, userID) } diff --git a/internal/service/media_test.go b/internal/service/media_test.go index 937e0ea..1ab7dfb 100644 --- a/internal/service/media_test.go +++ b/internal/service/media_test.go @@ -144,7 +144,7 @@ func TestMediaService_GetMediaDetail(t *testing.T) { name: "ok", mediaID: 1, userID: 1, - media: &model.Media{ID: 1, FileName: "a.mp4"}, + media: &model.Media{ID: 1, SetID: 1, FileName: "a.mp4"}, tags: []model.Tag{{ID: 1, Name: "rock"}}, fav: true, note: &model.Note{MediaID: 1, UserID: 1, Content: "hello"}, @@ -154,7 +154,7 @@ func TestMediaService_GetMediaDetail(t *testing.T) { name: "not found", mediaID: 2, media: nil, - wantNil: true, + wantErr: true, }, { name: "media error", @@ -165,14 +165,14 @@ func TestMediaService_GetMediaDetail(t *testing.T) { { name: "tags error", mediaID: 1, - media: &model.Media{ID: 1}, + media: &model.Media{ID: 1, SetID: 1}, tagsErr: errors.New("boom"), wantErr: true, }, { name: "favorite error", mediaID: 1, - media: &model.Media{ID: 1}, + media: &model.Media{ID: 1, SetID: 1}, favErr: errors.New("boom"), wantErr: true, }, @@ -186,6 +186,16 @@ func TestMediaService_GetMediaDetail(t *testing.T) { return tt.media, tt.mediaErr }, }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, TagRepo: repository.MockTagRepo{ ListTagsByMediaFunc: func(ctx context.Context, mediaID int64) ([]model.Tag, error) { return tt.tags, tt.tagsErr @@ -478,6 +488,21 @@ func TestMediaService_DownloadMedia_Access(t *testing.T) { func TestMediaService_ToggleFavorite(t *testing.T) { ctx := context.Background() store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, FavoriteRepo: repository.MockFavoriteRepo{ ToggleFavoriteFunc: func(ctx context.Context, userID, mediaID int64) (bool, error) { return true, nil @@ -530,6 +555,21 @@ func TestMediaService_AssignTag(t *testing.T) { t.Run(tt.name, func(t *testing.T) { var created bool store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, TagRepo: repository.MockTagRepo{ GetTagByNameFunc: func(ctx context.Context, name string) (*model.Tag, error) { if tt.tagExists { @@ -593,6 +633,21 @@ func TestMediaService_RemoveTag(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, TagRepo: repository.MockTagRepo{ GetTagByNameFunc: func(ctx context.Context, name string) (*model.Tag, error) { if tt.tagExists { @@ -636,6 +691,11 @@ func TestMediaService_SoftDeleteMedia(t *testing.T) { return &model.User{ID: 1, IsAdmin: true}, nil }, }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, } svc := NewMediaService(store, newMockClock(), "/tmp/media") if err := svc.SoftDeleteMedia(ctx, 1, 1); err != nil { @@ -648,11 +708,24 @@ func TestMediaService_RestoreMedia(t *testing.T) { var called bool store := &repository.MockStore{ MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, RestoreMediaFunc: func(ctx context.Context, id int64) error { called = true return nil }, }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: 1, IsAdmin: true}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, } svc := NewMediaService(store, newMockClock(), "/tmp/media") if err := svc.RestoreMedia(ctx, 1, 1); err != nil { @@ -721,6 +794,16 @@ func TestMediaService_UploadMedia(t *testing.T) { return &model.Set{ID: 1, RootPath: "music"}, tt.setErr }, }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + SetPermissionRepo: repository.MockSetPermissionRepo{ + GetPermissionFunc: func(ctx context.Context, setID, userID int64) (*model.SetPermission, error) { + return &model.SetPermission{SetID: setID, UserID: userID, Role: model.RoleOwner}, nil + }, + }, MediaRepo: repository.MockMediaRepo{ CreateMediaFunc: func(ctx context.Context, media *model.Media) (int64, error) { return 1, tt.createErr @@ -921,6 +1004,21 @@ func TestMediaService_Notes(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id}, nil + }, + }, NoteRepo: repository.MockNoteRepo{ GetNoteFunc: func(ctx context.Context, mediaID, userID int64) (*model.Note, error) { return &model.Note{MediaID: mediaID, UserID: userID, Content: "hello"}, tt.noteErr @@ -959,3 +1057,379 @@ func TestMediaService_Notes(t *testing.T) { func intPtr(i int) *int { return &i } + +func TestMediaService_ViewerCannotMutate(t *testing.T) { + ctx := context.Background() + + makeViewerStore := func(mediaID, setID int64) *repository.MockStore { + return &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + if id == mediaID { + return &model.Media{ID: mediaID, SetID: setID}, nil + } + return nil, nil + }, + SoftDeleteMediaFunc: func(ctx context.Context, id int64) error { return nil }, + RestoreMediaFunc: func(ctx context.Context, id int64) error { return nil }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id, Permissions: []model.SetPermission{{SetID: setID, UserID: 2, Role: model.RoleViewer}}}, nil + }, + }, + } + } + + t.Run("viewer cannot soft delete", func(t *testing.T) { + store := makeViewerStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.SoftDeleteMedia(ctx, 1, 2) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("viewer cannot restore", func(t *testing.T) { + store := makeViewerStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.RestoreMedia(ctx, 1, 2) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("viewer cannot upload", func(t *testing.T) { + store := &repository.MockStore{ + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: 1, RootPath: "music", Permissions: []model.SetPermission{{SetID: 1, UserID: 2, Role: model.RoleViewer}}}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), t.TempDir()) + _, err := svc.UploadMedia(ctx, 1, 2, "song.mp3", strings.NewReader("data"), 4) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("viewer cannot regenerate thumbnail", func(t *testing.T) { + store := makeViewerStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.RegenerateThumbnail(ctx, 1, 2) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("viewer cannot regenerate set cover", func(t *testing.T) { + store := &repository.MockStore{ + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: 1, Permissions: []model.SetPermission{{SetID: 1, UserID: 2, Role: model.RoleViewer}}}, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.RegenerateSetCover(ctx, 1, 2) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) +} + +func TestMediaService_UnauthorizedAccessDenied(t *testing.T) { + ctx := context.Background() + + makeUnauthorizedStore := func(mediaID, setID int64) *repository.MockStore { + return &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + if id == mediaID { + return &model.Media{ID: mediaID, SetID: setID}, nil + } + return nil, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: id, Permissions: []model.SetPermission{}}, nil + }, + }, + SetPermissionRepo: repository.MockSetPermissionRepo{ + GetPermissionFunc: func(ctx context.Context, setID, userID int64) (*model.SetPermission, error) { + return nil, nil + }, + }, + } + } + + t.Run("unauthorized cannot get detail", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + _, err := svc.GetMediaDetail(ctx, 1, 9) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot list media in set", func(t *testing.T) { + store := &repository.MockStore{ + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetPermissionRepo: repository.MockSetPermissionRepo{ + ListPermissionsByUserFunc: func(ctx context.Context, userID int64) ([]model.SetPermission, error) { + return nil, nil + }, + }, + MediaRepo: repository.MockMediaRepo{ + ListMediaFunc: func(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) { + return nil, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), "/tmp/media") + res, err := svc.ListMedia(ctx, 9, repository.MediaFilter{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(res) != 0 { + t.Fatalf("expected empty list, got %d", len(res)) + } + }) + + t.Run("unauthorized cannot favorite", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + _, err := svc.ToggleFavorite(ctx, 9, 1) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot assign tag", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.AssignTag(ctx, 1, 9, "rock") + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot remove tag", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.RemoveTag(ctx, 1, 9, "rock") + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot get note", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + _, err := svc.GetNote(ctx, 1, 9) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot upsert note", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.UpsertNote(ctx, &model.Note{MediaID: 1, UserID: 9, Content: "hello"}) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot delete note", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.DeleteNote(ctx, 1, 9) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot create share", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + _, err := svc.CreateShare(ctx, 9, 1, time.Now().Add(time.Hour)) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("unauthorized cannot list shares", func(t *testing.T) { + store := makeUnauthorizedStore(1, 1) + svc := NewMediaService(store, newMockClock(), "/tmp/media") + _, err := svc.ListShares(ctx, 1, 9) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("expected ErrForbidden, got %v", err) + } + }) + + t.Run("owner can soft delete", func(t *testing.T) { + store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, + SoftDeleteMediaFunc: func(ctx context.Context, id int64) error { return nil }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: 1, Permissions: []model.SetPermission{{SetID: 1, UserID: 2, Role: model.RoleOwner}}}, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.SoftDeleteMedia(ctx, 1, 2) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("owner can restore", func(t *testing.T) { + store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: 1, SetID: 1}, nil + }, + RestoreMediaFunc: func(ctx context.Context, id int64) error { return nil }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: 1, Permissions: []model.SetPermission{{SetID: 1, UserID: 2, Role: model.RoleOwner}}}, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), "/tmp/media") + err := svc.RestoreMedia(ctx, 1, 2) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("owner can upload", func(t *testing.T) { + tmpDir := t.TempDir() + store := &repository.MockStore{ + SetRepo: repository.MockSetRepo{ + GetSetByIDFunc: func(ctx context.Context, id int64) (*model.Set, error) { + return &model.Set{ID: 1, RootPath: "music", Permissions: []model.SetPermission{{SetID: 1, UserID: 2, Role: model.RoleOwner}}}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + MediaRepo: repository.MockMediaRepo{ + CreateMediaFunc: func(ctx context.Context, media *model.Media) (int64, error) { + return 1, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), tmpDir) + media, err := svc.UploadMedia(ctx, 1, 2, "song.mp3", strings.NewReader("data"), 4) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if media == nil { + t.Fatal("expected media") + } + }) +} + +func TestMediaService_ListMedia_AdminAndUserFiltering(t *testing.T) { + ctx := context.Background() + + t.Run("admin sees all", func(t *testing.T) { + store := &repository.MockStore{ + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + MediaRepo: repository.MockMediaRepo{ + ListMediaFunc: func(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) { + return []model.Media{{ID: 1}, {ID: 2}}, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), "/tmp/media") + res, err := svc.ListMedia(ctx, 1, repository.MediaFilter{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(res) != 2 { + t.Fatalf("expected 2, got %d", len(res)) + } + }) + + t.Run("user sees only allowed sets", func(t *testing.T) { + store := &repository.MockStore{ + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: false}, nil + }, + }, + SetPermissionRepo: repository.MockSetPermissionRepo{ + ListPermissionsByUserFunc: func(ctx context.Context, userID int64) ([]model.SetPermission, error) { + return []model.SetPermission{{SetID: 3, UserID: userID, Role: model.RoleViewer}}, nil + }, + }, + MediaRepo: repository.MockMediaRepo{ + ListMediaFunc: func(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) { + if len(filter.AllowedSetIDs) == 1 && filter.AllowedSetIDs[0] == 3 { + return []model.Media{{ID: 5}}, nil + } + return nil, nil + }, + }, + } + svc := NewMediaService(store, newMockClock(), "/tmp/media") + res, err := svc.ListMedia(ctx, 2, repository.MediaFilter{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(res) != 1 || res[0].ID != 5 { + t.Fatalf("unexpected result: %+v", res) + } + }) +} diff --git a/internal/service/mock.go b/internal/service/mock.go index 73a4755..d41ebce 100644 --- a/internal/service/mock.go +++ b/internal/service/mock.go @@ -14,7 +14,7 @@ import ( type MockMediaService struct { ListSetsFunc func(ctx context.Context, userID int64) ([]model.Set, error) GetMediaDetailFunc func(ctx context.Context, mediaID, userID int64) (*MediaDetail, error) - ListMediaFunc func(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) + ListMediaFunc func(ctx context.Context, userID int64, filter repository.MediaFilter) ([]model.Media, error) StreamMediaFunc func(ctx context.Context, mediaID, userID int64) (*FileResult, error) DownloadMediaFunc func(ctx context.Context, mediaID, userID int64) (*FileResult, error) GetThumbnailFunc |
