diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/repository/repository.go | 2 | ||||
| -rw-r--r-- | internal/service/mock.go | 27 | ||||
| -rw-r--r-- | internal/service/mock_test.go | 9 | ||||
| -rw-r--r-- | internal/service/progress.go | 88 | ||||
| -rw-r--r-- | internal/service/progress_test.go | 188 | ||||
| -rw-r--r-- | internal/service/service.go | 6 |
6 files changed, 315 insertions, 5 deletions
diff --git a/internal/repository/repository.go b/internal/repository/repository.go index e7ff672..2dacc90 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -55,9 +55,9 @@ type AccessHelperStore interface { // ProgressServiceStore is the subset of Store required by service.ProgressService. type ProgressServiceStore interface { + AccessHelperStore PlaybackProgressRepo PlaybackAccumulatorRepo - MediaRepo } // GCStore is the subset of Store required by service.GCWorker. diff --git a/internal/service/mock.go b/internal/service/mock.go index f585c52..2593d9e 100644 --- a/internal/service/mock.go +++ b/internal/service/mock.go @@ -404,6 +404,9 @@ func (m *MockAuthService) GetUserByID(ctx context.Context, id int64) (*model.Use // MockProgressService is a fake ProgressService for testing. type MockProgressService struct { UpdateProgressFunc func(ctx context.Context, sessionID string, userID, mediaID int64, position float64) error + MarkFinishedFunc func(ctx context.Context, userID, mediaID int64) error + MarkNotStartedFunc func(ctx context.Context, userID, mediaID int64) error + ListInProgressFunc func(ctx context.Context, userID int64) ([]model.Media, error) } // UpdateProgress calls UpdateProgressFunc or returns nil. @@ -413,3 +416,27 @@ func (m *MockProgressService) UpdateProgress(ctx context.Context, sessionID stri } return nil } + +// MarkFinished calls MarkFinishedFunc or returns nil. +func (m *MockProgressService) MarkFinished(ctx context.Context, userID, mediaID int64) error { + if m.MarkFinishedFunc != nil { + return m.MarkFinishedFunc(ctx, userID, mediaID) + } + return nil +} + +// MarkNotStarted calls MarkNotStartedFunc or returns nil. +func (m *MockProgressService) MarkNotStarted(ctx context.Context, userID, mediaID int64) error { + if m.MarkNotStartedFunc != nil { + return m.MarkNotStartedFunc(ctx, userID, mediaID) + } + return nil +} + +// ListInProgress calls ListInProgressFunc or returns nil. +func (m *MockProgressService) ListInProgress(ctx context.Context, userID int64) ([]model.Media, error) { + if m.ListInProgressFunc != nil { + return m.ListInProgressFunc(ctx, userID) + } + return nil, nil +} diff --git a/internal/service/mock_test.go b/internal/service/mock_test.go index c90bac6..4055aaa 100644 --- a/internal/service/mock_test.go +++ b/internal/service/mock_test.go @@ -156,12 +156,21 @@ func TestMockProgressService_Defaults(t *testing.T) { ctx := context.Background() m := &MockProgressService{} m.UpdateProgress(ctx, "sess", 1, 1, 10) + m.MarkFinished(ctx, 1, 1) + m.MarkNotStarted(ctx, 1, 1) + m.ListInProgress(ctx, 1) } func TestMockProgressService_WithFunc(t *testing.T) { ctx := context.Background() m := &MockProgressService{ UpdateProgressFunc: func(ctx context.Context, sessionID string, userID, mediaID int64, position float64) error { return nil }, + MarkFinishedFunc: func(ctx context.Context, userID, mediaID int64) error { return nil }, + MarkNotStartedFunc: func(ctx context.Context, userID, mediaID int64) error { return nil }, + ListInProgressFunc: func(ctx context.Context, userID int64) ([]model.Media, error) { return nil, nil }, } m.UpdateProgress(ctx, "sess", 1, 1, 10) + m.MarkFinished(ctx, 1, 1) + m.MarkNotStarted(ctx, 1, 1) + m.ListInProgress(ctx, 1) } diff --git a/internal/service/progress.go b/internal/service/progress.go index ef248cf..4e79aa5 100644 --- a/internal/service/progress.go +++ b/internal/service/progress.go @@ -12,15 +12,17 @@ import ( // progressService is the concrete implementation of ProgressService. type progressService struct { - store repository.ProgressServiceStore - clock clock.Clock + store repository.ProgressServiceStore + helper *accessHelper + clock clock.Clock } // NewProgressService creates a concrete ProgressService. func NewProgressService(store repository.ProgressServiceStore, clk clock.Clock) *progressService { return &progressService{ - store: store, - clock: clk, + store: store, + helper: NewAccessHelper(store), + clock: clk, } } @@ -82,3 +84,81 @@ func (s *progressService) UpdateProgress(ctx context.Context, sessionID string, return nil } + +func (s *progressService) MarkFinished(ctx context.Context, userID, mediaID int64) error { + if mediaID == 0 { + return errors.New("media_id required") + } + + media, err := s.helper.verifyAccess(ctx, mediaID, userID) + if err != nil { + return err + } + + if err := s.store.UpsertProgress(ctx, &model.PlaybackProgress{ + UserID: userID, + MediaID: mediaID, + PositionSeconds: media.Duration, + Finished: true, + UpdatedAt: s.clock.Now(), + }); err != nil { + return fmt.Errorf("mark finished: %w", err) + } + + return nil +} + +func (s *progressService) MarkNotStarted(ctx context.Context, userID, mediaID int64) error { + if mediaID == 0 { + return errors.New("media_id required") + } + + if _, err := s.helper.verifyAccess(ctx, mediaID, userID); err != nil { + return err + } + + if err := s.store.DeleteProgress(ctx, userID, mediaID); err != nil { + return fmt.Errorf("delete progress: %w", err) + } + if err := s.store.DeleteAccumulatorByMedia(ctx, mediaID); err != nil { + return fmt.Errorf("delete accumulator: %w", err) + } + + return nil +} + +func (s *progressService) ListInProgress(ctx context.Context, userID int64) ([]model.Media, error) { + allowed, err := s.helper.allowedSetIDs(ctx, userID) + if err != nil { + return nil, err + } + if allowed != nil && len(allowed) == 0 { + return []model.Media{}, nil + } + + return s.store.ListInProgressMedia(ctx, userID, repository.MediaFilter{ + AllowedSetIDs: allowed, + }) +} + +func (h *accessHelper) allowedSetIDs(ctx context.Context, userID int64) ([]int64, error) { + user, err := h.store.GetUserByID(ctx, userID) + if err != nil { + return nil, fmt.Errorf("get user: %w", err) + } + if user != nil && user.IsAdmin { + return nil, nil + } + + perms, err := h.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) + } + + return allowed, nil +} diff --git a/internal/service/progress_test.go b/internal/service/progress_test.go index 85eb1b3..a2b36db 100644 --- a/internal/service/progress_test.go +++ b/internal/service/progress_test.go @@ -197,3 +197,191 @@ func TestProgressService_UpdateProgress(t *testing.T) { }) } } + +func TestProgressService_MarkFinished(t *testing.T) { + ctx := context.Background() + var saved *model.PlaybackProgress + + store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: id, SetID: 7, Duration: 123.5}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + PlaybackProgressRepo: repository.MockPlaybackProgressRepo{ + UpsertProgressFunc: func(ctx context.Context, progress *model.PlaybackProgress) error { + saved = progress + return nil + }, + }, + } + + svc := NewProgressService(store, newMockClock()) + if err := svc.MarkFinished(ctx, 1, 10); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if saved == nil { + t.Fatal("expected saved progress") + } + if !saved.Finished { + t.Fatal("expected progress marked finished") + } + if saved.PositionSeconds != 123.5 { + t.Fatalf("expected position_seconds=duration, got %v", saved.PositionSeconds) + } +} + +func TestProgressService_MarkFinished_Validation(t *testing.T) { + ctx := context.Background() + svc := NewProgressService(&repository.MockStore{}, newMockClock()) + + if err := svc.MarkFinished(ctx, 1, 0); err == nil { + t.Fatal("expected error for mediaID=0") + } +} + +func TestProgressService_MarkNotStarted(t *testing.T) { + ctx := context.Background() + var deletedProgress bool + var deletedAccumulator bool + + store := &repository.MockStore{ + MediaRepo: repository.MockMediaRepo{ + GetMediaByIDFunc: func(ctx context.Context, id int64) (*model.Media, error) { + return &model.Media{ID: id, SetID: 7}, nil + }, + }, + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return &model.User{ID: id, IsAdmin: true}, nil + }, + }, + PlaybackProgressRepo: repository.MockPlaybackProgressRepo{ + DeleteProgressFunc: func(ctx context.Context, userID, mediaID int64) error { + deletedProgress = userID == 1 && mediaID == 10 + return nil + }, + }, + PlaybackAccumulatorRepo: repository.MockPlaybackAccumulatorRepo{ + DeleteAccumulatorByMediaFunc: func(ctx context.Context, mediaID int64) error { + deletedAccumulator = mediaID == 10 + return nil + }, + }, + } + + svc := NewProgressService(store, newMockClock()) + if err := svc.MarkNotStarted(ctx, 1, 10); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !deletedProgress { + t.Fatal("expected DeleteProgress called") + } + if !deletedAccumulator { + t.Fatal("expected DeleteAccumulatorByMedia called") + } +} + +func TestProgressService_MarkNotStarted_Validation(t *testing.T) { + ctx := context.Background() + svc := NewProgressService(&repository.MockStore{}, newMockClock()) + + if err := svc.MarkNotStarted(ctx, 1, 0); err == nil { + t.Fatal("expected error for mediaID=0") + } +} + +func TestProgressService_ListInProgress(t *testing.T) { + ctx := context.Background() + + tests := []struct { + name string + user *model.User + perms []model.SetPermission + want []model.Media + wantAllow []int64 + wantCalls int + }{ + { + name: "admin lists without allowed set filter", + user: &model.User{ID: 1, IsAdmin: true}, + want: []model.Media{{ID: 10, SetID: 7}}, + wantAllow: nil, + wantCalls: 1, + }, + { + name: "viewer lists only permitted sets", + user: &model.User{ID: 2, IsAdmin: false}, + perms: []model.SetPermission{{SetID: 7, UserID: 2}, {SetID: 8, UserID: 2}}, + want: []model.Media{{ID: 10, SetID: 7}}, + wantAllow: []int64{7, 8}, + wantCalls: 1, + }, + { + name: "viewer with no permissions does not query media", + user: &model.User{ID: 3, IsAdmin: false}, + want: []model.Media{}, + wantAllow: nil, + wantCalls: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var calls int + var gotAllowed []int64 + + store := &repository.MockStore{ + UserRepo: repository.MockUserRepo{ + GetUserByIDFunc: func(ctx context.Context, id int64) (*model.User, error) { + return tt.user, nil + }, + }, + SetPermissionRepo: repository.MockSetPermissionRepo{ + ListPermissionsByUserFunc: func(ctx context.Context, userID int64) ([]model.SetPermission, error) { + return tt.perms, nil + }, + }, + PlaybackProgressRepo: repository.MockPlaybackProgressRepo{ + ListInProgressMediaFunc: func(ctx context.Context, userID int64, filter repository.MediaFilter) ([]model.Media, error) { + calls++ + gotAllowed = filter.AllowedSetIDs + return tt.want, nil + }, + }, + } + + svc := NewProgressService(store, newMockClock()) + got, err := svc.ListInProgress(ctx, tt.user.ID) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if calls != tt.wantCalls { + t.Fatalf("expected %d ListInProgressMedia calls, got %d", tt.wantCalls, calls) + } + if len(got) != len(tt.want) { + t.Fatalf("expected %d media, got %d", len(tt.want), len(got)) + } + if tt.wantCalls > 0 && !equalInt64Slices(gotAllowed, tt.wantAllow) { + t.Fatalf("expected AllowedSetIDs=%v, got %v", tt.wantAllow, gotAllowed) + } + }) + } +} + +func equalInt64Slices(a, b []int64) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} diff --git a/internal/service/service.go b/internal/service/service.go index ba826be..79298d7 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -223,6 +223,12 @@ type AuthResult struct { type ProgressService interface { // UpdateProgress stores a playback position and updates play-count accounting. UpdateProgress(ctx context.Context, sessionID string, userID, mediaID int64, position float64) error + // MarkFinished stores completed playback progress for a media item. + MarkFinished(ctx context.Context, userID, mediaID int64) error + // MarkNotStarted clears saved playback progress and playback counters for a media item. + MarkNotStarted(ctx context.Context, userID, mediaID int64) error + // ListInProgress returns unfinished media with saved playback positions visible to the user. + ListInProgress(ctx context.Context, userID int64) ([]model.Media, error) } // MediaStreamer prepares authorized file results for HTTP streaming. |
