From 100a9bb282dd4226368b4e5fac398726c7acf653 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Thu, 30 Apr 2026 11:28:21 +0300 Subject: task da: enforce media access and owner/admin role permissions Changes: - MediaService.ListMedia now accepts userID and filters by allowed sets for non-admins via AllowedSetIDs in repository.MediaFilter. - Handlers pass userID into ListMedia; API returns 403 for forbidden. - Added verifyModifyAccess and verifySetModifyAccess helpers so only owners/admins can upload, soft-delete, restore, and regenerate thumbnails/covers; viewers are blocked. - GetMediaDetail, ToggleFavorite, AssignTag, RemoveTag, notes, and shares now consistently verifyAccess before proceeding. - Handlers handle ErrForbidden with 403 for soft-delete and restore. - Added negative tests proving viewers cannot mutate and unauthorized users cannot access/detail/tag/note/favorite/share inaccessible media. --- internal/service/media.go | 134 ++++++++++- internal/service/media_test.go | 482 ++++++++++++++++++++++++++++++++++++++- internal/service/mock.go | 6 +- internal/service/no_rows_test.go | 34 ++- internal/service/service.go | 2 +- 5 files changed, 641 insertions(+), 17 deletions(-) (limited to 'internal/service') 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 func(ctx context.Context, mediaID, userID int64) (*FileResult, error) @@ -48,9 +48,9 @@ func (m *MockMediaService) GetMediaDetail(ctx context.Context, mediaID, userID i } return nil, nil } -func (m *MockMediaService) ListMedia(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) { +func (m *MockMediaService) ListMedia(ctx context.Context, userID int64, filter repository.MediaFilter) ([]model.Media, error) { if m.ListMediaFunc != nil { - return m.ListMediaFunc(ctx, filter) + return m.ListMediaFunc(ctx, userID, filter) } return nil, nil } diff --git a/internal/service/no_rows_test.go b/internal/service/no_rows_test.go index 4d512e5..3606b63 100644 --- a/internal/service/no_rows_test.go +++ b/internal/service/no_rows_test.go @@ -22,8 +22,8 @@ func TestService_NoRows_ReturnsNil(t *testing.T) { } svc := NewMediaService(store, newMockClock(), "/tmp/media") detail, err := svc.GetMediaDetail(ctx, 99, 1) - if err != nil { - t.Fatalf("expected no error, got %v", err) + if !errors.Is(err, ErrNotFound) { + t.Fatalf("expected ErrNotFound, got %v", err) } if detail != nil { t.Fatalf("expected nil detail, got %+v", detail) @@ -32,6 +32,21 @@ func TestService_NoRows_ReturnsNil(t *testing.T) { t.Run("GetNote nil", 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 nil, nil @@ -50,6 +65,21 @@ func TestService_NoRows_ReturnsNil(t *testing.T) { t.Run("AssignTag creates missing tag", 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) { return nil, nil diff --git a/internal/service/service.go b/internal/service/service.go index 28dfdc9..59f4e7d 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -14,7 +14,7 @@ import ( type MediaService interface { ListSets(ctx context.Context, userID int64) ([]model.Set, error) GetMediaDetail(ctx context.Context, mediaID, userID int64) (*MediaDetail, error) - ListMedia(ctx context.Context, filter repository.MediaFilter) ([]model.Media, error) + ListMedia(ctx context.Context, userID int64, filter repository.MediaFilter) ([]model.Media, error) StreamMedia(ctx context.Context, mediaID, userID int64) (*FileResult, error) DownloadMedia(ctx context.Context, mediaID, userID int64) (*FileResult, error) GetThumbnail(ctx context.Context, mediaID, userID int64) (*FileResult, error) -- cgit v1.2.3