diff options
| author | Paul Buetow <paul@buetow.org> | 2026-05-17 21:44:42 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-05-17 21:44:42 +0300 |
| commit | 60df3397fbd14a6ad200e1fdd2530f708403ee99 (patch) | |
| tree | 2d8615adad60a0d37d6d670a4dabebeea144c869 /player-server/internal | |
| parent | 0462f639bc1e1973d30b9ffcd358ac62fa10fb46 (diff) | |
Add bulk progress sync endpoint for offline mobile clients
Diffstat (limited to 'player-server/internal')
| -rw-r--r-- | player-server/internal/api/handlers_media.go | 85 | ||||
| -rw-r--r-- | player-server/internal/api/handlers_more_test.go | 1 | ||||
| -rw-r--r-- | player-server/internal/api/handlers_progress.go | 130 | ||||
| -rw-r--r-- | player-server/internal/api/handlers_test.go | 43 | ||||
| -rw-r--r-- | player-server/internal/api/server.go | 1 | ||||
| -rw-r--r-- | player-server/internal/repository/media.go | 6 | ||||
| -rw-r--r-- | player-server/internal/repository/playback_accumulator.go | 12 | ||||
| -rw-r--r-- | player-server/internal/repository/playback_progress.go | 12 | ||||
| -rw-r--r-- | player-server/internal/repository/progress_transaction.go | 59 | ||||
| -rw-r--r-- | player-server/internal/repository/repository.go | 14 | ||||
| -rw-r--r-- | player-server/internal/repository/sqlite.go | 9 | ||||
| -rw-r--r-- | player-server/internal/service/mock.go | 18 | ||||
| -rw-r--r-- | player-server/internal/service/progress.go | 75 | ||||
| -rw-r--r-- | player-server/internal/service/progress_test.go | 111 | ||||
| -rw-r--r-- | player-server/internal/service/service.go | 9 |
15 files changed, 484 insertions, 101 deletions
diff --git a/player-server/internal/api/handlers_media.go b/player-server/internal/api/handlers_media.go index 5d7a970..fb49189 100644 --- a/player-server/internal/api/handlers_media.go +++ b/player-server/internal/api/handlers_media.go @@ -451,88 +451,3 @@ func (s *Server) handleDeleteNote(w http.ResponseWriter, r *http.Request) { } writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } - -// ------------------------------------------------------------------ -// Progress -// ------------------------------------------------------------------ - -func (s *Server) handleProgress(w http.ResponseWriter, r *http.Request) { - if !requireService(w, s.progressSvc) { - return - } - var req struct { - MediaID int64 `json:"media_id"` - Position float64 `json:"position_seconds"` - } - if err := readJSON(r, &req); err != nil { - badRequest(w, "invalid body") - return - } - if req.MediaID == 0 { - badRequest(w, "media_id required") - return - } - sessionID := sessionIDFromContext(r) - if sessionID == "" { - badRequest(w, "session required") - return - } - err := s.progressSvc.UpdateProgress( - r.Context(), - sessionID, - userIDFromContext(r), - req.MediaID, - req.Position, - ) - if err != nil { - handleError(w, err) - return - } - writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) -} - -func (s *Server) handleProgressStatus(w http.ResponseWriter, r *http.Request) { - if !requireService(w, s.progressSvc) { - return - } - var req struct { - MediaID int64 `json:"media_id"` - Status string `json:"status"` - } - if err := readJSON(r, &req); err != nil { - badRequest(w, "invalid body") - return - } - if req.MediaID == 0 { - badRequest(w, "media_id required") - return - } - - var err error - switch req.Status { - case "finished": - err = s.progressSvc.MarkFinished(r.Context(), userIDFromContext(r), req.MediaID) - case "not_started": - err = s.progressSvc.MarkNotStarted(r.Context(), userIDFromContext(r), req.MediaID) - default: - badRequest(w, "invalid status") - return - } - if err != nil { - handleError(w, err) - return - } - writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) -} - -func (s *Server) handleInProgress(w http.ResponseWriter, r *http.Request) { - if !requireService(w, s.progressSvc) { - return - } - media, err := s.progressSvc.ListInProgress(r.Context(), userIDFromContext(r)) - if err != nil { - handleError(w, err) - return - } - writeJSON(w, http.StatusOK, media) -} diff --git a/player-server/internal/api/handlers_more_test.go b/player-server/internal/api/handlers_more_test.go index 451e911..5ae71bb 100644 --- a/player-server/internal/api/handlers_more_test.go +++ b/player-server/internal/api/handlers_more_test.go @@ -2177,6 +2177,7 @@ func TestServer_NilProgressSvc(t *testing.T) { body string }{ {http.MethodPost, "/api/progress", `{"media_id":1,"position_seconds":5}`}, + {http.MethodPost, "/api/progress/batch", `{"updates":[{"media_id":1,"position_seconds":5}]}`}, {http.MethodPost, "/api/progress/status", `{"media_id":1,"status":"finished"}`}, {http.MethodGet, "/api/in-progress", ``}, } diff --git a/player-server/internal/api/handlers_progress.go b/player-server/internal/api/handlers_progress.go new file mode 100644 index 0000000..0d82bb7 --- /dev/null +++ b/player-server/internal/api/handlers_progress.go @@ -0,0 +1,130 @@ +package api + +import ( + "net/http" + "time" + + "codeberg.org/snonux/player/internal/service" +) + +func (s *Server) handleProgress(w http.ResponseWriter, r *http.Request) { + if !requireService(w, s.progressSvc) { + return + } + var req struct { + MediaID int64 `json:"media_id"` + Position float64 `json:"position_seconds"` + } + if err := readJSON(r, &req); err != nil { + badRequest(w, "invalid body") + return + } + if req.MediaID == 0 { + badRequest(w, "media_id required") + return + } + sessionID := sessionIDFromContext(r) + if sessionID == "" { + badRequest(w, "session required") + return + } + err := s.progressSvc.UpdateProgress( + r.Context(), + sessionID, + userIDFromContext(r), + req.MediaID, + req.Position, + ) + if err != nil { + handleError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) +} + +func (s *Server) handleBatchProgress(w http.ResponseWriter, r *http.Request) { + if !requireService(w, s.progressSvc) { + return + } + var req struct { + Updates []struct { + MediaID int64 `json:"media_id"` + PositionSeconds float64 `json:"position_seconds"` + ObservedAt time.Time `json:"observed_at"` + } `json:"updates"` + } + if err := readJSON(r, &req); err != nil { + badRequest(w, "invalid body") + return + } + + updates := make([]service.ProgressUpdate, len(req.Updates)) + for i, update := range req.Updates { + if update.MediaID == 0 { + badRequest(w, "media_id required") + return + } + updates[i] = service.ProgressUpdate{ + MediaID: update.MediaID, + PositionSeconds: update.PositionSeconds, + ObservedAt: update.ObservedAt, + } + } + + sessionID := sessionIDFromContext(r) + if sessionID == "" { + badRequest(w, "session required") + return + } + if err := s.progressSvc.BatchUpdateProgress(r.Context(), sessionID, userIDFromContext(r), updates); err != nil { + handleError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) +} + +func (s *Server) handleProgressStatus(w http.ResponseWriter, r *http.Request) { + if !requireService(w, s.progressSvc) { + return + } + var req struct { + MediaID int64 `json:"media_id"` + Status string `json:"status"` + } + if err := readJSON(r, &req); err != nil { + badRequest(w, "invalid body") + return + } + if req.MediaID == 0 { + badRequest(w, "media_id required") + return + } + + var err error + switch req.Status { + case "finished": + err = s.progressSvc.MarkFinished(r.Context(), userIDFromContext(r), req.MediaID) + case "not_started": + err = s.progressSvc.MarkNotStarted(r.Context(), userIDFromContext(r), req.MediaID) + default: + badRequest(w, "invalid status") + return + } + if err != nil { + handleError(w, err) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) +} + +func (s *Server) handleInProgress(w http.ResponseWriter, r *http.Request) { + if !requireService(w, s.progressSvc) { + return + } + media, err := s.progressSvc.ListInProgress(r.Context(), userIDFromContext(r)) + if err != nil { + handleError(w, err) + return + } + writeJSON(w, http.StatusOK, media) +} diff --git a/player-server/internal/api/handlers_test.go b/player-server/internal/api/handlers_test.go index 0800ba5..aa15611 100644 --- a/player-server/internal/api/handlers_test.go +++ b/player-server/internal/api/handlers_test.go @@ -1399,6 +1399,49 @@ func TestServer_Progress(t *testing.T) { }) } +func TestServer_BatchProgress(t *testing.T) { + var gotSessionID string + var gotUserID int64 + var gotUpdates []service.ProgressUpdate + ps := &service.MockProgressService{ + BatchUpdateProgressFunc: func(ctx context.Context, sessionID string, userID int64, updates []service.ProgressUpdate) error { + gotSessionID = sessionID + gotUserID = userID + gotUpdates = updates + return nil + }, + } + 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, nil, nil, nil, nil, nil, nil, nil, ps, nil, nil) + + body := `{"updates":[{"media_id":5,"position_seconds":12.3,"observed_at":"2026-05-17T10:00:00Z"}]}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/progress/batch", strings.NewReader(body)) + req.AddCookie(addSessionCookie(t, store, sm, 1)) + req.Header.Set("Content-Type", "application/json") + rr := httptest.NewRecorder() + srv.ServeHTTP(rr, req) + if rr.Code != http.StatusOK { + t.Fatalf("expected %d, got %d", http.StatusOK, rr.Code) + } + if gotSessionID == "" { + t.Fatal("expected session id passed to progress service") + } + if gotUserID != 1 { + t.Fatalf("expected user id 1, got %d", gotUserID) + } + if len(gotUpdates) != 1 { + t.Fatalf("expected 1 update, got %d", len(gotUpdates)) + } + if gotUpdates[0].MediaID != 5 || gotUpdates[0].PositionSeconds != 12.3 { + t.Fatalf("unexpected update: %+v", gotUpdates[0]) + } + if gotUpdates[0].ObservedAt.IsZero() { + t.Fatal("expected observed_at decoded") + } +} + func TestServer_ProgressStatus(t *testing.T) { store := buildSessionStore(1) sm := auth.NewSessionManager(store, &clock.MockClock{T: time.Now()}, time.Hour) diff --git a/player-server/internal/api/server.go b/player-server/internal/api/server.go index 9f1a526..5a7998b 100644 --- a/player-server/internal/api/server.go +++ b/player-server/internal/api/server.go @@ -236,6 +236,7 @@ func (s *Server) routesNotes() { // routesProgress wires the progress API routes. func (s *Server) routesProgress() { s.handleBoth(http.MethodPost, "/api/progress", s.requireSession(s.handleProgress)) + s.handleBoth(http.MethodPost, "/api/progress/batch", s.requireSession(s.handleBatchProgress)) s.handleBoth(http.MethodPost, "/api/progress/status", s.requireSession(s.handleProgressStatus)) s.handleBoth(http.MethodGet, "/api/in-progress", s.requireSession(s.handleInProgress)) } diff --git a/player-server/internal/repository/media.go b/player-server/internal/repository/media.go index a7a9553..60a611d 100644 --- a/player-server/internal/repository/media.go +++ b/player-server/internal/repository/media.go @@ -300,7 +300,11 @@ func (s *SQLite) ListDeletedMedia(ctx context.Context) ([]model.Media, error) { // IncrementPlayCount increments the play_count of a media by 1. func (s *SQLite) IncrementPlayCount(ctx context.Context, id int64) error { - _, err := s.db.ExecContext(ctx, `UPDATE media SET play_count = play_count + 1 WHERE id = ?`, id) + return incrementPlayCount(ctx, s.db, id) +} + +func incrementPlayCount(ctx context.Context, db sqlExecer, id int64) error { + _, err := db.ExecContext(ctx, `UPDATE media SET play_count = play_count + 1 WHERE id = ?`, id) if err != nil { return fmt.Errorf("increment play count: %w", err) } diff --git a/player-server/internal/repository/playback_accumulator.go b/player-server/internal/repository/playback_accumulator.go index ea2bcf0..cbbb88f 100644 --- a/player-server/internal/repository/playback_accumulator.go +++ b/player-server/internal/repository/playback_accumulator.go @@ -10,7 +10,11 @@ import ( // UpsertAccumulator inserts or replaces a playback accumulator. func (s *SQLite) UpsertAccumulator(ctx context.Context, acc *model.PlaybackAccumulator) error { - _, err := s.db.ExecContext(ctx, + return upsertAccumulator(ctx, s.db, acc) +} + +func upsertAccumulator(ctx context.Context, db sqlExecer, acc *model.PlaybackAccumulator) error { + _, err := db.ExecContext(ctx, `INSERT OR REPLACE INTO playback_accumulator (session_id, media_id, last_position, accumulated_seconds, counted, updated_at) VALUES (?, ?, ?, ?, ?, ?)`, acc.SessionID, acc.MediaID, acc.LastPosition, acc.AccumulatedSeconds, boolToInt(acc.Counted), acc.UpdatedAt, ) @@ -22,7 +26,11 @@ func (s *SQLite) UpsertAccumulator(ctx context.Context, acc *model.PlaybackAccum // GetAccumulator retrieves a playback accumulator by session and media. func (s *SQLite) GetAccumulator(ctx context.Context, sessionID string, mediaID int64) (*model.PlaybackAccumulator, error) { - row := s.db.QueryRowContext(ctx, + return getAccumulator(ctx, s.db, sessionID, mediaID) +} + +func getAccumulator(ctx context.Context, db sqlQueryRower, sessionID string, mediaID int64) (*model.PlaybackAccumulator, error) { + row := db.QueryRowContext(ctx, `SELECT session_id, media_id, last_position, accumulated_seconds, counted, updated_at FROM playback_accumulator WHERE session_id = ? AND media_id = ?`, sessionID, mediaID, ) diff --git a/player-server/internal/repository/playback_progress.go b/player-server/internal/repository/playback_progress.go index d044f92..74b877c 100644 --- a/player-server/internal/repository/playback_progress.go +++ b/player-server/internal/repository/playback_progress.go @@ -11,7 +11,11 @@ import ( // UpsertProgress inserts or replaces playback progress. func (s *SQLite) UpsertProgress(ctx context.Context, progress *model.PlaybackProgress) error { - _, err := s.db.ExecContext(ctx, + return upsertProgress(ctx, s.db, progress) +} + +func upsertProgress(ctx context.Context, db sqlExecer, progress *model.PlaybackProgress) error { + _, err := db.ExecContext(ctx, `INSERT OR REPLACE INTO playback_progress (user_id, media_id, position_seconds, finished, updated_at) VALUES (?, ?, ?, ?, ?)`, progress.UserID, progress.MediaID, progress.PositionSeconds, progress.Finished, progress.UpdatedAt, ) @@ -23,7 +27,11 @@ func (s *SQLite) UpsertProgress(ctx context.Context, progress *model.PlaybackPro // GetProgress retrieves playback progress for a user and media. func (s *SQLite) GetProgress(ctx context.Context, userID, mediaID int64) (*model.PlaybackProgress, error) { - row := s.db.QueryRowContext(ctx, + return getProgress(ctx, s.db, userID, mediaID) +} + +func getProgress(ctx context.Context, db sqlQueryRower, userID, mediaID int64) (*model.PlaybackProgress, error) { + row := db.QueryRowContext(ctx, `SELECT user_id, media_id, position_seconds, finished, updated_at FROM playback_progress WHERE user_id = ? AND media_id = ?`, userID, mediaID, ) diff --git a/player-server/internal/repository/progress_transaction.go b/player-server/internal/repository/progress_transaction.go new file mode 100644 index 0000000..b16018f --- /dev/null +++ b/player-server/internal/repository/progress_transaction.go @@ -0,0 +1,59 @@ +package repository + +import ( + "context" + "fmt" + + "codeberg.org/snonux/player/internal/model" +) + +// WithProgressTransaction runs progress updates in one SQLite transaction. +func (s *SQLite) WithProgressTransaction(ctx context.Context, fn func(ProgressUpdateStore) error) error { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin progress transaction: %w", err) + } + + txStore := &progressTxStore{tx: tx} + if err := fn(txStore); err != nil { + if rollbackErr := tx.Rollback(); rollbackErr != nil { + return fmt.Errorf("rollback progress transaction after %w: %v", err, rollbackErr) + } + return err + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit progress transaction: %w", err) + } + return nil +} + +type progressTxStore struct { + tx sqlProgressTx +} + +type sqlProgressTx interface { + sqlExecer + sqlQueryRower +} + +func (s *progressTxStore) UpsertProgress(ctx context.Context, progress *model.PlaybackProgress) error { + return upsertProgress(ctx, s.tx, progress) +} + +func (s *progressTxStore) GetProgress(ctx context.Context, userID, mediaID int64) (*model.PlaybackProgress, error) { + return getProgress(ctx, s.tx, userID, mediaID) +} + +func (s *progressTxStore) GetAccumulator(ctx context.Context, sessionID string, mediaID int64) (*model.PlaybackAccumulator, error) { + return getAccumulator(ctx, s.tx, sessionID, mediaID) +} + +func (s *progressTxStore) UpsertAccumulator(ctx context.Context, acc *model.PlaybackAccumulator) error { + return upsertAccumulator(ctx, s.tx, acc) +} + +func (s *progressTxStore) IncrementPlayCount(ctx context.Context, id int64) error { + return incrementPlayCount(ctx, s.tx, id) +} + +var _ ProgressUpdateStore = (*progressTxStore)(nil) diff --git a/player-server/internal/repository/repository.go b/player-server/internal/repository/repository.go index a1ad9fa..487921f 100644 --- a/player-server/internal/repository/repository.go +++ b/player-server/internal/repository/repository.go @@ -296,6 +296,20 @@ type PlaybackAccumulatorRepo interface { DeleteAccumulatorByMedia(ctx context.Context, mediaID int64) error } +// ProgressUpdateStore is the narrow store needed for applying playback updates. +type ProgressUpdateStore interface { + UpsertProgress(ctx context.Context, progress *model.PlaybackProgress) error + GetProgress(ctx context.Context, userID, mediaID int64) (*model.PlaybackProgress, error) + GetAccumulator(ctx context.Context, sessionID string, mediaID int64) (*model.PlaybackAccumulator, error) + UpsertAccumulator(ctx context.Context, acc *model.PlaybackAccumulator) error + IncrementPlayCount(ctx context.Context, id int64) error +} + +// ProgressTransactionStore applies progress updates inside one database transaction. +type ProgressTransactionStore interface { + WithProgressTransaction(ctx context.Context, fn func(ProgressUpdateStore) error) error +} + // SessionRepo manages browser sessions. type SessionRepo interface { // CreateSession stores a browser session. diff --git a/player-server/internal/repository/sqlite.go b/player-server/internal/repository/sqlite.go index 8d69ebf..1fb4f61 100644 --- a/player-server/internal/repository/sqlite.go +++ b/player-server/internal/repository/sqlite.go @@ -61,11 +61,20 @@ func (s *SQLite) Ping(ctx context.Context) error { var _ Store = (*SQLite)(nil) var _ APITokenRepo = (*SQLite)(nil) var _ PodcastRepo = (*SQLite)(nil) +var _ ProgressTransactionStore = (*SQLite)(nil) type sqlScanner interface { Scan(dest ...any) error } +type sqlExecer interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) +} + +type sqlQueryRower interface { + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row +} + func boolToInt(b bool) int { if b { return 1 diff --git a/player-server/internal/service/mock.go b/player-server/internal/service/mock.go index ef844a0..dbf1166 100644 --- a/player-server/internal/service/mock.go +++ b/player-server/internal/service/mock.go @@ -18,6 +18,7 @@ var ( _ MediaNoteService = (*MockMediaService)(nil) _ MediaService = (*MockMediaService)(nil) _ AuthService = (*MockAuthService)(nil) + _ ProgressService = (*MockProgressService)(nil) ) // MockMediaService is a fake MediaService for testing. @@ -439,10 +440,11 @@ 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) + UpdateProgressFunc func(ctx context.Context, sessionID string, userID, mediaID int64, position float64) error + BatchUpdateProgressFunc func(ctx context.Context, sessionID string, userID int64, updates []ProgressUpdate) 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. @@ -453,6 +455,14 @@ func (m *MockProgressService) UpdateProgress(ctx context.Context, sessionID stri return nil } +// BatchUpdateProgress calls BatchUpdateProgressFunc or returns nil. +func (m *MockProgressService) BatchUpdateProgress(ctx context.Context, sessionID string, userID int64, updates []ProgressUpdate) error { + if m.BatchUpdateProgressFunc != nil { + return m.BatchUpdateProgressFunc(ctx, sessionID, userID, updates) + } + return nil +} + // MarkFinished calls MarkFinishedFunc or returns nil. func (m *MockProgressService) MarkFinished(ctx context.Context, userID, mediaID int64) error { if m.MarkFinishedFunc != nil { diff --git a/player-server/internal/service/progress.go b/player-server/internal/service/progress.go index 4e79aa5..6b2fe84 100644 --- a/player-server/internal/service/progress.go +++ b/player-server/internal/service/progress.go @@ -4,6 +4,8 @@ import ( "context" "errors" "fmt" + "sort" + "time" "codeberg.org/snonux/player/internal/clock" "codeberg.org/snonux/player/internal/model" @@ -34,18 +36,77 @@ func (s *progressService) UpdateProgress(ctx context.Context, sessionID string, return errors.New("media_id required") } + return s.applyProgress(ctx, s.store, sessionID, userID, mediaID, position, s.clock.Now()) +} + +func (s *progressService) BatchUpdateProgress(ctx context.Context, sessionID string, userID int64, updates []ProgressUpdate) error { + if sessionID == "" { + return errors.New("session_id required") + } + + normalized := make([]ProgressUpdate, len(updates)) now := s.clock.Now() + for i, update := range updates { + if update.MediaID == 0 { + return fmt.Errorf("updates[%d].media_id required", i) + } + if update.ObservedAt.IsZero() { + update.ObservedAt = now + } + normalized[i] = update + } + sort.SliceStable(normalized, func(i, j int) bool { + return normalized[i].ObservedAt.Before(normalized[j].ObservedAt) + }) - if err := s.store.UpsertProgress(ctx, &model.PlaybackProgress{ + apply := func(store repository.ProgressUpdateStore) error { + for _, update := range normalized { + if err := s.applyProgress( + ctx, + store, + sessionID, + userID, + update.MediaID, + update.PositionSeconds, + update.ObservedAt, + ); err != nil { + return fmt.Errorf("apply progress media_id %d: %w", update.MediaID, err) + } + } + return nil + } + if txStore, ok := s.store.(repository.ProgressTransactionStore); ok { + return txStore.WithProgressTransaction(ctx, apply) + } + return apply(s.store) +} + +func (s *progressService) applyProgress( + ctx context.Context, + store repository.ProgressUpdateStore, + sessionID string, + userID, mediaID int64, + position float64, + observedAt time.Time, +) error { + existing, err := store.GetProgress(ctx, userID, mediaID) + if err != nil { + return fmt.Errorf("get progress: %w", err) + } + if existing != nil && existing.UpdatedAt.After(observedAt) { + return nil + } + + if err := store.UpsertProgress(ctx, &model.PlaybackProgress{ UserID: userID, MediaID: mediaID, PositionSeconds: position, - UpdatedAt: now, + UpdatedAt: observedAt, }); err != nil { return fmt.Errorf("upsert progress: %w", err) } - acc, err := s.store.GetAccumulator(ctx, sessionID, mediaID) + acc, err := store.GetAccumulator(ctx, sessionID, mediaID) if err != nil { return fmt.Errorf("get accumulator: %w", err) } @@ -56,7 +117,7 @@ func (s *progressService) UpdateProgress(ctx context.Context, sessionID string, LastPosition: 0, AccumulatedSeconds: 0, Counted: false, - UpdatedAt: now, + UpdatedAt: observedAt, } } @@ -69,16 +130,16 @@ func (s *progressService) UpdateProgress(ctx context.Context, sessionID string, } acc.AccumulatedSeconds += delta acc.LastPosition = position - acc.UpdatedAt = now + acc.UpdatedAt = observedAt if acc.AccumulatedSeconds >= 60 && !acc.Counted { - if err := s.store.IncrementPlayCount(ctx, mediaID); err != nil { + if err := store.IncrementPlayCount(ctx, mediaID); err != nil { return fmt.Errorf("increment play count: %w", err) } acc.Counted = true } - if err := s.store.UpsertAccumulator(ctx, acc); err != nil { + if err := store.UpsertAccumulator(ctx, acc); err != nil { return fmt.Errorf("upsert accumulator: %w", err) } diff --git a/player-server/internal/service/progress_test.go b/player-server/internal/service/progress_test.go index a2b36db..6e6a1a0 100644 --- a/player-server/internal/service/progress_test.go +++ b/player-server/internal/service/progress_test.go @@ -4,7 +4,9 @@ import ( "context" "errors" "testing" + "time" + "codeberg.org/snonux/player/internal/clock" "codeberg.org/snonux/player/internal/model" "codeberg.org/snonux/player/internal/repository" ) @@ -198,6 +200,115 @@ func TestProgressService_UpdateProgress(t *testing.T) { } } +func TestProgressService_BatchUpdateProgress_OrdersByObservedAt(t *testing.T) { + ctx := context.Background() + observedBase := time.Date(2026, 5, 17, 10, 0, 0, 0, time.UTC) + progressByMedia := make(map[int64]*model.PlaybackProgress) + accByMedia := make(map[int64]*model.PlaybackAccumulator) + var positions []float64 + + store := &repository.MockStore{ + PlaybackProgressRepo: repository.MockPlaybackProgressRepo{ + GetProgressFunc: func(ctx context.Context, userID, mediaID int64) (*model.PlaybackProgress, error) { + return progressByMedia[mediaID], nil + }, + UpsertProgressFunc: func(ctx context.Context, progress *model.PlaybackProgress) error { + cp := *progress + progressByMedia[progress.MediaID] = &cp + positions = append(positions, progress.PositionSeconds) + return nil + }, + }, + PlaybackAccumulatorRepo: repository.MockPlaybackAccumulatorRepo{ + GetAccumulatorFunc: func(ctx context.Context, sessionID string, mediaID int64) (*model.PlaybackAccumulator, error) { + return accByMedia[mediaID], nil + }, + UpsertAccumulatorFunc: func(ctx context.Context, acc *model.PlaybackAccumulator) error { + cp := *acc + accByMedia[acc.MediaID] = &cp + return nil + }, + }, + } + + svc := NewProgressService(store, &clock.MockClock{T: observedBase}) + err := svc.BatchUpdateProgress(ctx, "sess", 1, []ProgressUpdate{ + {MediaID: 10, PositionSeconds: 30, ObservedAt: observedBase.Add(2 * time.Minute)}, + {MediaID: 10, PositionSeconds: 10, ObservedAt: observedBase}, + {MediaID: 11, PositionSeconds: 20, ObservedAt: observedBase.Add(time.Minute)}, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + wantPositions := []float64{10, 20, 30} + if len(positions) != len(wantPositions) { + t.Fatalf("expected %d progress writes, got %d", len(wantPositions), len(positions)) + } + for i, want := range wantPositions { + if positions[i] != want { + t.Fatalf("position call %d: expected %v, got %v", i, want, positions[i]) + } + } + if got := progressByMedia[10].PositionSeconds; got != 30 { + t.Fatalf("expected latest media 10 position to win, got %v", got) + } +} + +func TestProgressService_BatchUpdateProgress_TransactionRollback(t *testing.T) { + ctx := context.Background() + store, err := repository.Open(":memory:") + if err != nil { + t.Fatalf("open store: %v", err) + } + defer store.Close() + + now := time.Date(2026, 5, 17, 10, 0, 0, 0, time.UTC) + userID, err := store.CreateUser(ctx, &model.User{Username: "alice", PasswordHash: "hash", CreatedAt: now}) + if err != nil { + t.Fatalf("create user: %v", err) + } + setID, err := store.CreateSet(ctx, &model.Set{Name: "set", RootPath: "/media/set", CreatedAt: now}) + if err != nil { + t.Fatalf("create set: %v", err) + } + mediaID, err := store.CreateMedia(ctx, &model.Media{ + SetID: setID, + RelPath: "one.mp4", + FileName: "one.mp4", + AbsPath: "/media/set/one.mp4", + Type: model.MediaTypeVideo, + CreatedAt: now, + }) + if err != nil { + t.Fatalf("create media: %v", err) + } + if err := store.CreateSession(ctx, &model.Session{ + ID: "sess", + UserID: userID, + ExpiresAt: now.Add(time.Hour), + CreatedAt: now, + }); err != nil { + t.Fatalf("create session: %v", err) + } + + svc := NewProgressService(store, &clock.MockClock{T: now}) + err = svc.BatchUpdateProgress(ctx, "sess", userID, []ProgressUpdate{ + {MediaID: mediaID, PositionSeconds: 10, ObservedAt: now}, + {MediaID: 9999, PositionSeconds: 20, ObservedAt: now.Add(time.Second)}, + }) + if err == nil { + t.Fatal("expected batch update error") + } + progress, err := store.GetProgress(ctx, userID, mediaID) + if err != nil { + t.Fatalf("get progress: %v", err) + } + if progress != nil { + t.Fatalf("expected transaction rollback to remove first progress update, got %+v", progress) + } +} + func TestProgressService_MarkFinished(t *testing.T) { ctx := context.Background() var saved *model.PlaybackProgress diff --git a/player-server/internal/service/service.go b/player-server/internal/service/service.go index c4914b6..84525e5 100644 --- a/player-server/internal/service/service.go +++ b/player-server/internal/service/service.go @@ -237,6 +237,8 @@ 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 + // BatchUpdateProgress stores playback positions ordered by observation time. + BatchUpdateProgress(ctx context.Context, sessionID string, userID int64, updates []ProgressUpdate) 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. @@ -245,6 +247,13 @@ type ProgressService interface { ListInProgress(ctx context.Context, userID int64) ([]model.Media, error) } +// ProgressUpdate is one observed playback position from a client. +type ProgressUpdate struct { + MediaID int64 + PositionSeconds float64 + ObservedAt time.Time +} + // MediaStreamer prepares authorized file results for HTTP streaming. type MediaStreamer interface { // Open opens a file result and returns the headers/reader needed by the API. |
