summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/repository/repository.go2
-rw-r--r--internal/service/mock.go27
-rw-r--r--internal/service/mock_test.go9
-rw-r--r--internal/service/progress.go88
-rw-r--r--internal/service/progress_test.go188
-rw-r--r--internal/service/service.go6
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.