diff --git a/.golangci.yml b/.golangci.yml index 7df4f8ec3..923c06c7a 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -9,6 +9,7 @@ linters: - asciicheck - bidichk - bodyclose + - containedctx - copyloopvar - dogsled - durationcheck @@ -57,6 +58,9 @@ linters: - gosec path: _test\.go text: "G703" + - path: _test\.go + linters: + - containedctx - path-except: 'db/migrations/' linters: - forbidigo diff --git a/adapters/lastfm/agent_test.go b/adapters/lastfm/agent_test.go index b9fb786c3..0f68bd46b 100644 --- a/adapters/lastfm/agent_test.go +++ b/adapters/lastfm/agent_test.go @@ -357,7 +357,7 @@ var _ = Describe("lastfmAgent", func() { var httpClient *tests.FakeHttpClient var track *model.MediaFile BeforeEach(func() { - _ = ds.UserProps(ctx).Put("user-1", sessionKeyProperty, "SK-1") + _ = ds.UserProps().Put(ctx, "user-1", sessionKeyProperty, "SK-1") httpClient = &tests.FakeHttpClient{} client := newClient("API_KEY", "SECRET", httpClient) agent = lastFMConstructor(ds) diff --git a/adapters/lastfm/auth_router_test.go b/adapters/lastfm/auth_router_test.go index 1f65c059e..476daf443 100644 --- a/adapters/lastfm/auth_router_test.go +++ b/adapters/lastfm/auth_router_test.go @@ -50,7 +50,7 @@ var _ = Describe("auth_router", func() { }) storedSessionKey := func(userID string) string { - key, _ := userProps.Get(userID, sessionKeyProperty) + key, _ := userProps.Get(GinkgoT().Context(), userID, sessionKeyProperty) return key } diff --git a/adapters/listenbrainz/agent_test.go b/adapters/listenbrainz/agent_test.go index a201b7c3a..d2aa07f1a 100644 --- a/adapters/listenbrainz/agent_test.go +++ b/adapters/listenbrainz/agent_test.go @@ -30,7 +30,7 @@ var _ = Describe("listenBrainzAgent", func() { BeforeEach(func() { ds = &tests.MockDataStore{} ctx = context.Background() - _ = ds.UserProps(ctx).Put("user-1", sessionKeyProperty, "SK-1") + _ = ds.UserProps().Put(ctx, "user-1", sessionKeyProperty, "SK-1") httpClient = &tests.FakeHttpClient{} agent = listenBrainzConstructor(ds) agent.client = newClient("http://localhost:8080", httpClient) diff --git a/cmd/artwork.go b/cmd/artwork.go index 9a64cd8c3..72685678e 100644 --- a/cmd/artwork.go +++ b/cmd/artwork.go @@ -174,28 +174,28 @@ func queueTotal(stats []model.ArtworkQueueStat) int64 { } func collectStatus(ctx context.Context, ds model.DataStore) (statusReport, error) { - q := ds.ArtworkQueue(ctx) + q := ds.ArtworkQueue() var rep statusReport var err error - if rep.queue, err = q.CountQueued(nil, nil); err != nil { + if rep.queue, err = q.CountQueued(ctx, nil, nil); err != nil { return rep, fmt.Errorf("breaking the artwork queue down by kind: %w", err) } for _, k := range artwork.ReprocessKinds { - sources, err := q.SourcesInUse(k) + sources, err := q.SourcesInUse(ctx, k) if err != nil { return rep, fmt.Errorf("listing the sources in use by %s artwork: %w", k, err) } slices.Sort(sources) for _, s := range sources { - n, err := q.CountBySource(k, []string{s}) + n, err := q.CountBySource(ctx, k, []string{s}) if err != nil { return rep, fmt.Errorf("counting %s artwork resolved from %s: %w", k, displaySource(s), err) } rep.sources = append(rep.sources, sourceCount{kind: k, source: s, count: n}) // An absent state is exactly a row with no source, so it needs no second query. if s == "" { - failed, err := q.CountBySource(k, []string{model.ArtworkSourceFailed}) + failed, err := q.CountBySource(ctx, k, []string{model.ArtworkSourceFailed}) if err != nil { return rep, fmt.Errorf("counting failed %s artwork: %w", k, err) } @@ -205,7 +205,7 @@ func collectStatus(ctx context.Context, ds model.DataStore) (statusReport, error } rep.current, rep.inputs = artwork.ConfigFingerprint(), artwork.FingerprintInputs() - if rep.stored, err = ds.Property(ctx).DefaultGet(consts.ArtConfFingerprintPropertyKey, ""); err != nil { + if rep.stored, err = ds.Property().DefaultGet(ctx, consts.ArtConfFingerprintPropertyKey, ""); err != nil { return rep, fmt.Errorf("reading the stored artwork fingerprint: %w", err) } return rep, nil @@ -442,13 +442,13 @@ func promptConfirm(in io.Reader, verb string) confirmFunc { // validateSources rejects a typo'd source: matching nothing silently reads as "nothing to do" when // it means the filter was wrong. Checked table-wide, so a filter is never a typo for one --kind only. -func validateSources(q model.ArtworkQueueRepository, sources []string) error { +func validateSources(ctx context.Context, q model.ArtworkQueueRepository, sources []string) error { if len(sources) == 0 { return nil } var inUse []string for _, k := range artwork.ReprocessKinds { - found, err := q.SourcesInUse(k) + found, err := q.SourcesInUse(ctx, k) if err != nil { return fmt.Errorf("listing the sources in use by %s artwork: %w", k, err) } @@ -475,8 +475,8 @@ func validateSources(q model.ArtworkQueueRepository, sources []string) error { // actually inserted; the two differ because an already-queued row is left untouched. func reprocessArtwork(ctx context.Context, ds model.DataStore, kinds []model.Kind, sources []string, imageAgents artwork.ImageAgentCount, dryRun bool, confirm confirmFunc, out io.Writer) error { - q := ds.ArtworkQueue(ctx) - if err := validateSources(q, sources); err != nil { + q := ds.ArtworkQueue() + if err := validateSources(ctx, q, sources); err != nil { return err } @@ -495,7 +495,7 @@ func reprocessArtwork(ctx context.Context, ds model.DataStore, kinds []model.Kin matched := make([]int64, len(kinds)) var total, external int64 for i, k := range kinds { - n, err := q.CountBySource(k, sources) + n, err := q.CountBySource(ctx, k, sources) if err != nil { return fmt.Errorf("counting %s artwork: %w", k, err) } @@ -523,7 +523,7 @@ func reprocessArtwork(ctx context.Context, ds model.DataStore, kinds []model.Kin if matched[i] == 0 { continue } - n, err := q.EnqueueBySource(k, sources, model.ArtworkPriorityRecheck) + n, err := q.EnqueueBySource(ctx, k, sources, model.ArtworkPriorityRecheck) if err != nil { return fmt.Errorf("queueing %s artwork: %w", k, err) } @@ -590,8 +590,8 @@ func parseAll[T comparable](values []string, parse func(string) (T, error)) ([]T func cancelArtwork(ctx context.Context, ds model.DataStore, kinds []model.Kind, priorities []int, dryRun bool, confirm confirmFunc, out io.Writer) error { - q := ds.ArtworkQueue(ctx) - matched, err := q.CountQueued(kinds, priorities) + q := ds.ArtworkQueue() + matched, err := q.CountQueued(ctx, kinds, priorities) if err != nil { return fmt.Errorf("counting queued artwork: %w", err) } @@ -612,7 +612,7 @@ func cancelArtwork(ctx context.Context, ds model.DataStore, kinds []model.Kind, return nil } - cancelled, err := q.PurgeQueued(kinds, priorities) + cancelled, err := q.PurgeQueued(ctx, kinds, priorities) if err != nil { return fmt.Errorf("cancelling queued artwork: %w", err) } @@ -984,11 +984,11 @@ func runExplain(ctx context.Context, args []string) { } rep := explainReport{kind: kind, id: id, name: name} if artwork.KeepsState(kind) { - rep.stored, err = ds.Artwork(ctx).GetItemArtwork(kind, id, model.ImageTypePrimary) + rep.stored, err = ds.Artwork().GetItemArtwork(ctx, kind, id, model.ImageTypePrimary) if err != nil && !errors.Is(err, model.ErrNotFound) { log.Fatal(ctx, "Failed to read artwork state", "kind", kind, "id", id, err) } - rep.queued, err = ds.ArtworkQueue(ctx).Get(kind, id, model.ImageTypePrimary) + rep.queued, err = ds.ArtworkQueue().Get(ctx, kind, id, model.ImageTypePrimary) if err != nil && !errors.Is(err, model.ErrNotFound) { log.Fatal(ctx, "Failed to read the artwork queue", "kind", kind, "id", id, err) } diff --git a/cmd/artwork_test.go b/cmd/artwork_test.go index 71b450914..df0ea0665 100644 --- a/cmd/artwork_test.go +++ b/cmd/artwork_test.go @@ -491,17 +491,17 @@ var _ = Describe("explain/reprocess source round trip", func() { It("names the absent state as reprocess --source accepts it", func() { ds := &tests.MockDataStore{} - art := ds.Artwork(ctx).(*tests.MockArtworkRepo) - Expect(art.PutItemArtwork(&model.ItemArtwork{ItemKind: model.KindArtistArtwork.Prefix(), + art := ds.Artwork().(*tests.MockArtworkRepo) + Expect(art.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: model.KindArtistArtwork.Prefix(), ItemID: "ar-1", ImageType: model.ImageTypePrimary})).To(Succeed()) shown := storedSource(formatExplain(explainReport{kind: model.KindArtistArtwork, id: "ar-1", stored: &model.ItemArtwork{AttemptedAt: time.Now()}})) - q := ds.ArtworkQueue(ctx) - Expect(validateSources(q, repositorySources([]string{shown}))).To(Succeed(), + q := ds.ArtworkQueue() + Expect(validateSources(ctx, q, repositorySources([]string{shown}))).To(Succeed(), "explain's spelling of a source must be pasteable into --source") - Expect(validateSources(q, repositorySources([]string{"(" + shown + ")"}))).ToNot(Succeed(), + Expect(validateSources(ctx, q, repositorySources([]string{"(" + shown + ")"}))).ToNot(Succeed(), "a parenthesised name would be rejected, so explain must not print one") }) }) @@ -573,7 +573,7 @@ var _ = Describe("reprocessArtwork", func() { decline := func(io.Writer, int64, int64) bool { return false } put := func(kind model.Kind, id, source string) { - Expect(art.PutItemArtwork(&model.ItemArtwork{ItemKind: kind.Prefix(), ItemID: id, + Expect(art.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: kind.Prefix(), ItemID: id, ImageType: model.ImageTypePrimary, Hash: "h" + id, Source: source})).To(Succeed()) } @@ -584,8 +584,8 @@ var _ = Describe("reprocessArtwork", func() { conf.Server.EnableM3UExternalAlbumArt = false imageAgents = artwork.ImageAgentCount{Artist: 1, Album: 1} ds = &tests.MockDataStore{} - art = ds.Artwork(ctx).(*tests.MockArtworkRepo) - queue = ds.ArtworkQueue(ctx).(*tests.MockArtworkQueueRepo) + art = ds.Artwork().(*tests.MockArtworkRepo) + queue = ds.ArtworkQueue().(*tests.MockArtworkQueueRepo) out.Reset() put(model.KindArtistArtwork, "ar-1", "external:deezer") put(model.KindArtistArtwork, "ar-2", "") @@ -601,19 +601,19 @@ var _ = Describe("reprocessArtwork", func() { Expect(out.String()).To(ContainSubstring("album")) Expect(out.String()).To(ContainSubstring("TOTAL")) Expect(out.String()).To(ContainSubstring("Dry run")) - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) }) It("queues nothing when the operator declines", func() { Expect(reprocessArtwork(ctx, ds, kinds, nil, imageAgents, false, decline, &out)).To(Succeed()) Expect(out.String()).To(ContainSubstring("Aborted")) - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) }) DescribeTable("records the applied config only for a run that leaves nothing on the old one", func(selected []model.Kind, sources []string, dryRun, applied bool) { - Expect(ds.Property(ctx).Put(consts.ArtConfFingerprintPropertyKey, "stale-fingerprint")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.ArtConfFingerprintPropertyKey, "stale-fingerprint")).To(Succeed()) Expect(reprocessArtwork(ctx, ds, selected, sources, imageAgents, dryRun, accept, &out)).To(Succeed()) @@ -621,7 +621,7 @@ var _ = Describe("reprocessArtwork", func() { if applied { want = artwork.ConfigFingerprint() } - Expect(ds.Property(ctx).Get(consts.ArtConfFingerprintPropertyKey)).To(Equal(want)) + Expect(ds.Property().Get(ctx, consts.ArtConfFingerprintPropertyKey)).To(Equal(want)) }, Entry("every kind, unfiltered", artwork.ReprocessKinds, nil, false, true), Entry("every kind, but nothing matched", artwork.ReprocessKinds, []string{}, false, true), @@ -633,14 +633,14 @@ var _ = Describe("reprocessArtwork", func() { It("queues the matching items at recheck priority, leaving their artwork state alone", func() { Expect(reprocessArtwork(ctx, ds, kinds, []string{"external:deezer"}, imageAgents, false, accept, &out)).To(Succeed()) - Expect(queue.Count()).To(Equal(int64(2))) - queued, err := queue.Get(model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) + Expect(queue.Count(ctx)).To(Equal(int64(2))) + queued, err := queue.Get(ctx, model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(queued.Priority).To(Equal(model.ArtworkPriorityRecheck)) - _, err = queue.Get(model.KindAlbumArtwork, "al-2", model.ImageTypePrimary) + _, err = queue.Get(ctx, model.KindAlbumArtwork, "al-2", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "a non-matching source must not be queued") - stored, err := art.GetItemArtwork(model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) + stored, err := art.GetItemArtwork(ctx, model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(stored.Hash).To(Equal("hal-1"), "bulk reprocessing must not blank the current artwork") }) @@ -648,20 +648,20 @@ var _ = Describe("reprocessArtwork", func() { It("targets the absent state", func() { Expect(reprocessArtwork(ctx, ds, kinds, []string{""}, imageAgents, false, accept, &out)).To(Succeed()) - Expect(queue.Count()).To(Equal(int64(1))) - _, err := queue.Get(model.KindArtistArtwork, "ar-2", model.ImageTypePrimary) + Expect(queue.Count(ctx)).To(Equal(int64(1))) + _, err := queue.Get(ctx, model.KindArtistArtwork, "ar-2", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) }) It("reports matched and queued separately when part of the set is already queued", func() { - Expect(queue.Enqueue(model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar-1", + Expect(queue.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar-1", ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBump})).To(Succeed()) Expect(reprocessArtwork(ctx, ds, kinds, []string{"external:deezer"}, imageAgents, false, accept, &out)).To(Succeed()) Expect(out.String()).To(ContainSubstring("Queued 1 of 2 matched items")) Expect(out.String()).To(ContainSubstring("Already queued, left unchanged: 1")) - queued, err := queue.Get(model.KindArtistArtwork, "ar-1", model.ImageTypePrimary) + queued, err := queue.Get(ctx, model.KindArtistArtwork, "ar-1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(queued.Priority).To(Equal(model.ArtworkPriorityBump), "an already-queued row keeps its priority and backoff") @@ -675,7 +675,7 @@ var _ = Describe("reprocessArtwork", func() { }, &out)).To(Succeed()) Expect(out.String()).To(ContainSubstring("Nothing")) - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) }) It("reports an empty selection as a dry run when one was asked for", func() { @@ -767,7 +767,7 @@ var _ = Describe("reprocessArtwork", func() { Expect(err.Error()).To(ContainSubstring("external:deezer")) Expect(err.Error()).To(ContainSubstring("folder")) Expect(err.Error()).To(ContainSubstring("absent"), "the empty source prints under its user-facing name") - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) }) It("accepts the absent filter with nothing absent, still rejecting a typo", func() { @@ -777,7 +777,7 @@ var _ = Describe("reprocessArtwork", func() { imageAgents, false, accept, &out)).To(Succeed(), "a reserved source must stay valid once the library has none of it") Expect(out.String()).To(ContainSubstring("Nothing matches")) - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) Expect(reprocessArtwork(ctx, ds, kinds, repositorySources([]string{"absnt"}), imageAgents, true, accept, &out)).ToNot(Succeed(), "a typo must still be rejected") @@ -797,7 +797,7 @@ var _ = Describe("reprocessArtwork", func() { Expect(out.String()).To(ContainSubstring("Nothing matches"), "a well-formed filter must not be reported as a typo because of the kinds selected") - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) }) }) @@ -816,19 +816,19 @@ var _ = Describe("collectStatus", func() { BeforeEach(func() { ds = &tests.MockDataStore{} - art = ds.Artwork(ctx).(*tests.MockArtworkRepo) - queue = ds.ArtworkQueue(ctx).(*tests.MockArtworkQueueRepo) + art = ds.Artwork().(*tests.MockArtworkRepo) + queue = ds.ArtworkQueue().(*tests.MockArtworkQueueRepo) put := func(kind model.Kind, id, source, hash string, attempted time.Time) { - Expect(art.PutItemArtwork(&model.ItemArtwork{ItemKind: kind.Prefix(), ItemID: id, + Expect(art.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: kind.Prefix(), ItemID: id, ImageType: model.ImageTypePrimary, Source: source, Hash: hash, AttemptedAt: attempted})).To(Succeed()) } put(model.KindArtistArtwork, "ar-1", "external:deezer", "h1", time.Now()) put(model.KindArtistArtwork, "ar-2", "", "", time.Now().Add(-24*time.Hour)) // ar-3 is absent because it gave up, so the two absent artists split across the columns. - Expect(art.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "ar-3", + Expect(art.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "ar-3", ImageType: model.ImageTypePrimary, LastFailure: "[]", AttemptedAt: time.Now()})).To(Succeed()) put(model.KindAlbumArtwork, "al-1", "folder", "h2", time.Now()) - Expect(queue.Enqueue(model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar-9", + Expect(queue.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar-9", ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBackfill})).To(Succeed()) }) @@ -848,7 +848,7 @@ var _ = Describe("collectStatus", func() { }) It("compares the stored fingerprint against the current one", func() { - Expect(ds.Property(ctx).Put(consts.ArtConfFingerprintPropertyKey, "old-fingerprint")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.ArtConfFingerprintPropertyKey, "old-fingerprint")).To(Succeed()) rep, err := collectStatus(ctx, ds) Expect(err).ToNot(HaveOccurred()) @@ -860,7 +860,7 @@ var _ = Describe("collectStatus", func() { It("queues nothing", func() { _, err := collectStatus(ctx, ds) Expect(err).ToNot(HaveOccurred()) - Expect(queue.Count()).To(Equal(int64(1)), "status must not enqueue anything") + Expect(queue.Count(ctx)).To(Equal(int64(1)), "status must not enqueue anything") }) }) @@ -973,21 +973,21 @@ var _ = Describe("refreshItems", func() { albums := tests.CreateMockAlbumRepo() albums.SetData(model.Albums{{ID: "al-1"}, {ID: "al-3"}}) ds = &tests.MockDataStore{MockedAlbum: albums} - art = ds.Artwork(ctx).(*tests.MockArtworkRepo) - queue = ds.ArtworkQueue(ctx).(*tests.MockArtworkQueueRepo) + art = ds.Artwork().(*tests.MockArtworkRepo) + queue = ds.ArtworkQueue().(*tests.MockArtworkQueueRepo) out.Reset() }) It("clears the stored state and queues each id at Bump priority", func() { - Expect(art.PutItemArtwork(&model.ItemArtwork{ItemKind: model.KindAlbumArtwork.Prefix(), + Expect(art.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: model.KindAlbumArtwork.Prefix(), ItemID: "al-1", ImageType: model.ImageTypePrimary, Hash: "abc123"})).To(Succeed()) Expect(refreshItems(ctx, ds, []model.ArtworkID{ {Kind: model.KindAlbumArtwork, ID: "al-1"}, {Kind: model.KindAlbumArtwork, ID: "al-3"}}, &out)).To(BeZero()) - _, err := art.GetItemArtwork(model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) + _, err := art.GetItemArtwork(ctx, model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) - queued, err := queue.Get(model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) + queued, err := queue.Get(ctx, model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(queued.Priority).To(Equal(model.ArtworkPriorityBump)) Expect(out.String()).To(Equal("al/al-1: queued\nal/al-3: queued\n")) @@ -996,7 +996,7 @@ var _ = Describe("refreshItems", func() { It("skips an id that does not exist instead of queuing it", func() { Expect(refreshItems(ctx, ds, []model.ArtworkID{{Kind: model.KindAlbumArtwork, ID: "al-2"}}, &out)).To(Equal(1)) - _, err := queue.Get(model.KindAlbumArtwork, "al-2", model.ImageTypePrimary) + _, err := queue.Get(ctx, model.KindAlbumArtwork, "al-2", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "a typo must not leave an orphan queue row") Expect(out.String()).To(BeEmpty()) }) @@ -1145,9 +1145,9 @@ var _ = Describe("cancelArtwork", func() { BeforeEach(func() { ds = &tests.MockDataStore{} - queue = ds.ArtworkQueue(ctx).(*tests.MockArtworkQueueRepo) + queue = ds.ArtworkQueue().(*tests.MockArtworkQueueRepo) out.Reset() - Expect(queue.Enqueue( + Expect(queue.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar-1", ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBackfill}, model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar-2", ImageType: model.ImageTypePrimary, @@ -1164,28 +1164,28 @@ var _ = Describe("cancelArtwork", func() { Expect(out.String()).To(ContainSubstring("backfill")) Expect(out.String()).To(ContainSubstring("TOTAL")) Expect(out.String()).To(ContainSubstring("Dry run")) - Expect(queue.Count()).To(BeNumerically("==", 3)) + Expect(queue.Count(ctx)).To(BeNumerically("==", 3)) }) It("cancels nothing when the operator declines", func() { Expect(cancelArtwork(ctx, ds, nil, nil, false, decline, &out)).To(Succeed()) Expect(out.String()).To(ContainSubstring("Aborted")) - Expect(queue.Count()).To(BeNumerically("==", 3)) + Expect(queue.Count(ctx)).To(BeNumerically("==", 3)) }) It("deletes the selected rows and leaves the rest queued", func() { Expect(cancelArtwork(ctx, ds, nil, []int{model.ArtworkPriorityBackfill}, false, accept, &out)).To(Succeed()) - Expect(queue.Count()).To(BeNumerically("==", 1)) - _, err := queue.Get(model.KindArtistArtwork, "ar-2", model.ImageTypePrimary) + Expect(queue.Count(ctx)).To(BeNumerically("==", 1)) + _, err := queue.Get(ctx, model.KindArtistArtwork, "ar-2", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred(), "a non-matching priority must stay queued") Expect(out.String()).To(ContainSubstring("Cancelled 2 of 2 matched items.")) }) It("cancels every kind and priority when neither filter is given", func() { Expect(cancelArtwork(ctx, ds, nil, nil, false, accept, &out)).To(Succeed()) - Expect(queue.Count()).To(BeZero()) + Expect(queue.Count(ctx)).To(BeZero()) }) It("stops at a selection that matches nothing instead of prompting", func() { @@ -1196,7 +1196,7 @@ var _ = Describe("cancelArtwork", func() { Expect(cancelArtwork(ctx, ds, []model.Kind{model.KindPlaylistArtwork}, nil, false, refuse, &out)).To(Succeed()) Expect(out.String()).To(ContainSubstring("Nothing matches this selection.")) - Expect(queue.Count()).To(BeNumerically("==", 3)) + Expect(queue.Count(ctx)).To(BeNumerically("==", 3)) }) It("reports a queue read failure instead of reporting nothing to cancel", func() { diff --git a/cmd/missing.go b/cmd/missing.go index 65f571661..ce95e39f0 100644 --- a/cmd/missing.go +++ b/cmd/missing.go @@ -72,7 +72,7 @@ func runMissingList(ctx context.Context) { } ds, ctx := getAdminContext(ctx) - mfs, err := ds.MediaFile(ctx).GetCursor(model.QueryOptions{ + mfs, err := ds.MediaFile().GetCursor(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, Sort: "path", }) @@ -128,7 +128,7 @@ func runMissingFix(ctx context.Context, missingRef, targetRef string) { // resolveMediaFile looks up a media file by ID first, then by path (optionally libraryID:path). func resolveMediaFile(ctx context.Context, ds model.DataStore, ref string) *model.MediaFile { - mf, err := ds.MediaFile(ctx).Get(ref) + mf, err := ds.MediaFile().Get(ctx, ref) if err == nil { return mf } @@ -136,7 +136,7 @@ func resolveMediaFile(ctx context.Context, ds model.DataStore, ref string) *mode log.Fatal(ctx, "Error looking up media file", "ref", ref, err) } - mfs, err := ds.MediaFile(ctx).FindByPaths([]string{ref}) + mfs, err := ds.MediaFile().FindByPaths(ctx, []string{ref}) if err != nil { log.Fatal(ctx, "Error looking up media file by path", "ref", ref, err) } diff --git a/cmd/pls.go b/cmd/pls.go index 93b411483..bf16ea420 100644 --- a/cmd/pls.go +++ b/cmd/pls.go @@ -109,7 +109,7 @@ func fetchPlaylists(ctx context.Context, ds model.DataStore, sort string) model. } options.Filters = squirrel.Eq{"owner_id": user.ID} } - pls, err := ds.Playlist(ctx).GetAll(options) + pls, err := ds.Playlist().GetAll(ctx, options) if err != nil { log.Fatal(ctx, "Failed to retrieve playlists", err) } @@ -117,17 +117,17 @@ func fetchPlaylists(ctx context.Context, ds model.DataStore, sort string) model. } func findPlaylist(ctx context.Context, ds model.DataStore, nameOrID string) *model.Playlist { - playlist, err := ds.Playlist(ctx).GetWithTracks(nameOrID, true, false) + playlist, err := ds.Playlist().GetWithTracks(ctx, nameOrID, true, false) if err != nil && !errors.Is(err, model.ErrNotFound) { log.Fatal("Error retrieving playlist", "name", nameOrID, err) } if errors.Is(err, model.ErrNotFound) { - playlists, err := ds.Playlist(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"playlist.name": nameOrID}}) + playlists, err := ds.Playlist().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"playlist.name": nameOrID}}) if err != nil { log.Fatal("Error retrieving playlist", "name", nameOrID, err) } if len(playlists) > 0 { - playlist, err = ds.Playlist(ctx).GetWithTracks(playlists[0].ID, true, false) + playlist, err = ds.Playlist().GetWithTracks(ctx, playlists[0].ID, true, false) if err != nil { log.Fatal("Error retrieving playlist", "name", nameOrID, err) } @@ -194,7 +194,7 @@ func runExport(ctx context.Context) { exported := 0 for _, pls := range allPls { - plsWithTracks, err := ds.Playlist(ctx).GetWithTracks(pls.ID, true, false) + plsWithTracks, err := ds.Playlist().GetWithTracks(ctx, pls.ID, true, false) if err != nil { log.Error("Error loading playlist tracks", "playlist", pls.Name, err) continue diff --git a/cmd/plugin.go b/cmd/plugin.go index ded28e969..7b8a9a393 100644 --- a/cmd/plugin.go +++ b/cmd/plugin.go @@ -243,7 +243,7 @@ func runPluginInfo(ctx context.Context, arg string) { } requirePluginsEnabled(ctx) ds, ctx := getAdminContext(ctx) - p, err := ds.Plugin(ctx).Get(arg) + p, err := ds.Plugin().Get(ctx, arg) if err != nil { log.Fatal(ctx, "Plugin not found", "id", arg, err) } @@ -264,7 +264,7 @@ func runPluginValidate(ctx context.Context, arg string) { } requirePluginsEnabled(ctx) ds, ctx := getAdminContext(ctx) - p, err := ds.Plugin(ctx).Get(arg) + p, err := ds.Plugin().Get(ctx, arg) if err != nil { log.Fatal(ctx, "Plugin not found", "id", arg, err) } @@ -329,7 +329,7 @@ func formatPluginList(list model.Plugins, format string) (string, error) { func runPluginList(ctx context.Context) { requirePluginsEnabled(ctx) ds, ctx := getAdminContext(ctx) - list, err := ds.Plugin(ctx).GetAll() + list, err := ds.Plugin().GetAll(ctx) if err != nil { log.Fatal(ctx, "Failed to list plugins", err) } @@ -372,7 +372,7 @@ var pluginEditCmd = &cobra.Command{ Run: func(cmd *cobra.Command, args []string) { requirePluginsEnabled(cmd.Context()) ds, ctx := getAdminContext(cmd.Context()) - cur, err := ds.Plugin(ctx).Get(args[0]) + cur, err := ds.Plugin().Get(ctx, args[0]) if err != nil { log.Fatal(ctx, "Plugin not found", "id", args[0], err) } diff --git a/cmd/root.go b/cmd/root.go index 02cd30240..b23674441 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -184,11 +184,11 @@ func schedulePeriodicScan(ctx context.Context) func() error { } func pidHashChanged(ds model.DataStore) (bool, error) { - pidAlbum, err := ds.Property(context.Background()).DefaultGet(consts.PIDAlbumKey, "") + pidAlbum, err := ds.Property().DefaultGet(context.Background(), consts.PIDAlbumKey, "") if err != nil { return false, err } - pidTrack, err := ds.Property(context.Background()).DefaultGet(consts.PIDTrackKey, "") + pidTrack, err := ds.Property().DefaultGet(context.Background(), consts.PIDTrackKey, "") if err != nil { return false, err } @@ -199,11 +199,11 @@ func pidHashChanged(ds model.DataStore) (bool, error) { func runInitialScan(ctx context.Context) func() error { return func() error { ds := CreateDataStore() - fullScanRequired, err := ds.Property(ctx).DefaultGet(consts.FullScanAfterMigrationFlagKey, "0") + fullScanRequired, err := ds.Property().DefaultGet(ctx, consts.FullScanAfterMigrationFlagKey, "0") if err != nil { return err } - inProgress, err := ds.Library(ctx).ScanInProgress() + inProgress, err := ds.Library().ScanInProgress(ctx) if err != nil { return err } @@ -219,7 +219,7 @@ func runInitialScan(ctx context.Context) func() error { switch { case fullScanRequired == "1": log.Warn(ctx, "Full scan required after migration") - _ = ds.Property(ctx).Delete(consts.FullScanAfterMigrationFlagKey) + _ = ds.Property().Delete(ctx, consts.FullScanAfterMigrationFlagKey) case pidHasChanged: log.Warn(ctx, "PID config changed, performing full scan") fullScanRequired = "1" diff --git a/cmd/svc.go b/cmd/svc.go index 7fec708ff..c71f5ef2b 100644 --- a/cmd/svc.go +++ b/cmd/svc.go @@ -44,7 +44,7 @@ var svcCmd = &cobra.Command{ } type svcControl struct { - ctx context.Context + ctx context.Context //nolint:containedctx // service lifecycle ctx, cancelled by Stop cancel context.CancelFunc done chan struct{} } diff --git a/cmd/user.go b/cmd/user.go index 1abf157b7..eb64e69fe 100644 --- a/cmd/user.go +++ b/cmd/user.go @@ -183,7 +183,7 @@ func runCreateUser(ctx context.Context) { ds, ctx := getAdminContext(ctx) err := ds.WithTx(func(tx model.DataStore) error { - existingUser, err := tx.User(ctx).FindByUsername(userID) + existingUser, err := tx.User().FindByUsername(ctx, userID) if existingUser != nil { return fmt.Errorf("existing user '%s'", userID) } @@ -193,7 +193,7 @@ func runCreateUser(ctx context.Context) { } if len(libraryIds) > 0 && !setAdmin { - user.Libraries, err = tx.Library(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"id": libraryIds}}) + user.Libraries, err = tx.Library().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"id": libraryIds}}) if err != nil { return err } @@ -202,13 +202,13 @@ func runCreateUser(ctx context.Context) { return libraryError(user.Libraries) } } else { - user.Libraries, err = tx.Library(ctx).GetAll() + user.Libraries, err = tx.Library().GetAll(ctx) if err != nil { return err } } - err = tx.User(ctx).Put(&user) + err = tx.User().Put(ctx, &user) if err != nil { return err } @@ -218,7 +218,7 @@ func runCreateUser(ctx context.Context) { updatedIds[idx] = lib.ID } - err = tx.User(ctx).SetUserLibraries(user.ID, updatedIds) + err = tx.User().SetUserLibraries(ctx, user.ID, updatedIds) return err }) @@ -236,7 +236,7 @@ func runDeleteUser(ctx context.Context) { var user *model.User err = ds.WithTx(func(tx model.DataStore) error { - count, err := tx.User(ctx).CountAll() + count, err := tx.User().CountAll(ctx) if err != nil { return err } @@ -250,7 +250,7 @@ func runDeleteUser(ctx context.Context) { return err } - return tx.User(ctx).Delete(user.ID) + return tx.User().Delete(ctx, user.ID) }) if err != nil { @@ -276,7 +276,7 @@ func runUserEdit(ctx context.Context) { } if len(libraryIds) > 0 && !setAdmin { - libraries, err := tx.Library(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"id": libraryIds}}) + libraries, err := tx.Library().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"id": libraryIds}}) if err != nil { return err @@ -291,7 +291,7 @@ func runUserEdit(ctx context.Context) { } if setAdmin && !user.IsAdmin { - libraries, err := tx.Library(ctx).GetAll() + libraries, err := tx.Library().GetAll(ctx) if err != nil { return err } @@ -337,7 +337,7 @@ func runUserEdit(ctx context.Context) { return nil } - err := tx.User(ctx).Put(user) + err := tx.User().Put(ctx, user) if err != nil { return err } @@ -348,7 +348,7 @@ func runUserEdit(ctx context.Context) { updatedIds[idx] = lib.ID } - err := tx.User(ctx).SetUserLibraries(user.ID, updatedIds) + err := tx.User().SetUserLibraries(ctx, user.ID, updatedIds) if err != nil { return err } @@ -393,13 +393,11 @@ func runUserList(ctx context.Context) { ds, ctx := getAdminContext(ctx) - users, err := ds.User(ctx).ReadAll() + userList, err := ds.User().ReadAll(ctx) if err != nil { log.Fatal(ctx, "Failed to retrieve users", err) } - userList := users.(model.Users) - if outputFormat == "csv" { w := csv.NewWriter(os.Stdout) _ = w.Write([]string{ diff --git a/cmd/utils.go b/cmd/utils.go index a67f720ac..72ec67f90 100644 --- a/cmd/utils.go +++ b/cmd/utils.go @@ -52,14 +52,14 @@ func getAdminContext(ctx context.Context) (model.DataStore, context.Context) { } func getUser(ctx context.Context, id string, ds model.DataStore) (*model.User, error) { - user, err := ds.User(ctx).FindByUsername(id) + user, err := ds.User().FindByUsername(ctx, id) if err != nil && !errors.Is(err, model.ErrNotFound) { return nil, fmt.Errorf("finding user by name: %w", err) } if errors.Is(err, model.ErrNotFound) { - user, err = ds.User(ctx).Get(id) + user, err = ds.User().Get(ctx, id) if err != nil { return nil, fmt.Errorf("finding user by id: %w", err) } diff --git a/core/agents/local_agent.go b/core/agents/local_agent.go index 1cb9060a1..c777ab46d 100644 --- a/core/agents/local_agent.go +++ b/core/agents/local_agent.go @@ -24,7 +24,7 @@ func (p *localAgent) AgentName() string { } func (p *localAgent) GetArtistTopSongs(ctx context.Context, id, artistName, mbid string, count int) ([]Song, error) { - top, err := p.ds.MediaFile(ctx).GetAll(model.QueryOptions{ + top, err := p.ds.MediaFile().GetAll(ctx, model.QueryOptions{ Sort: "playCount", Order: "desc", Max: count, @@ -43,7 +43,7 @@ func (p *localAgent) GetArtistTopSongs(ctx context.Context, id, artistName, mbid } func (p *localAgent) GetSimilarSongsByTrack(ctx context.Context, id, name, artist, mbid string, count int) ([]Song, error) { - seed, err := p.ds.MediaFile(ctx).Get(id) + seed, err := p.ds.MediaFile().Get(ctx, id) if err != nil { return nil, err } @@ -53,7 +53,7 @@ func (p *localAgent) GetSimilarSongsByTrack(ctx context.Context, id, name, artis return nil, nil } // Ask for extra so we can drop the seed itself and still fill the count. - candidates, err := p.ds.MediaFile(ctx).GetRandom(model.QueryOptions{ + candidates, err := p.ds.MediaFile().GetRandom(ctx, model.QueryOptions{ Filters: squirrel.And{ persistence.SongGenres.ByID(genreIDs), squirrel.Eq{"missing": false}, diff --git a/core/agents/session_keys.go b/core/agents/session_keys.go index cea6005ff..1eb414b15 100644 --- a/core/agents/session_keys.go +++ b/core/agents/session_keys.go @@ -13,13 +13,13 @@ type SessionKeys struct { } func (sk *SessionKeys) Put(ctx context.Context, userId, sessionKey string) error { - return sk.DataStore.UserProps(ctx).Put(userId, sk.KeyName, sessionKey) + return sk.DataStore.UserProps().Put(ctx, userId, sk.KeyName, sessionKey) } func (sk *SessionKeys) Get(ctx context.Context, userId string) (string, error) { - return sk.DataStore.UserProps(ctx).Get(userId, sk.KeyName) + return sk.DataStore.UserProps().Get(ctx, userId, sk.KeyName) } func (sk *SessionKeys) Delete(ctx context.Context, userId string) error { - return sk.DataStore.UserProps(ctx).Delete(userId, sk.KeyName) + return sk.DataStore.UserProps().Delete(ctx, userId, sk.KeyName) } diff --git a/core/archiver.go b/core/archiver.go index 33236d889..6f362322a 100644 --- a/core/archiver.go +++ b/core/archiver.go @@ -62,7 +62,7 @@ func (a *archiver) ZipArtist(ctx context.Context, id string, format string, bitr // rootArt, when set, is added to the archive root. func (a *archiver) zipAlbums(ctx context.Context, id string, format string, bitrate int, out io.Writer, filters squirrel.Sqlizer, rootArt model.ArtworkID) error { - mfs, err := a.ds.MediaFile(ctx).GetAll(model.QueryOptions{Filters: filters, Sort: "album"}) + mfs, err := a.ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: filters, Sort: "album"}) if err != nil { log.Error(ctx, "Error loading mediafiles from artist", "id", id, err) return err @@ -189,7 +189,7 @@ func (a *archiver) ZipShare(ctx context.Context, s *model.Share, out io.Writer) } func (a *archiver) ZipPlaylist(ctx context.Context, id string, format string, bitrate int, out io.Writer) error { - pls, err := a.ds.Playlist(ctx).GetWithTracks(id, true, false) + pls, err := a.ds.Playlist().GetWithTracks(ctx, id, true, false) if err != nil { log.Error(ctx, "Error loading mediafiles from playlist", "id", id, err) return err diff --git a/core/archiver_test.go b/core/archiver_test.go index 9dab44cef..4e00ce78c 100644 --- a/core/archiver_test.go +++ b/core/archiver_test.go @@ -52,7 +52,7 @@ var _ = Describe("Archiver", func() { Sort: "album", }}).Return(mfs, nil) - ds.On("MediaFile", mock.Anything).Return(mfRepo) + ds.On("MediaFile").Return(mfRepo) ms.On("NewStream", mock.Anything, mock.Anything, stream.Request{Format: "mp3", BitRate: 128}).Return(io.NopCloser(strings.NewReader("test")), nil).Times(3) out := new(bytes.Buffer) @@ -84,7 +84,7 @@ var _ = Describe("Archiver", func() { Sort: "album", }}).Return(mfs, nil) - ds.On("MediaFile", mock.Anything).Return(mfRepo) + ds.On("MediaFile").Return(mfRepo) ms.On("NewStream", mock.Anything, mock.Anything, stream.Request{Format: "mp3", BitRate: 128}).Return(io.NopCloser(strings.NewReader("test")), nil).Times(2) out := new(bytes.Buffer) @@ -246,7 +246,7 @@ var _ = Describe("Archiver", func() { Filters: squirrel.Eq{"album_id": "1"}, Sort: "album", }}).Return(mfs, nil) - ds.On("MediaFile", mock.Anything).Return(mfRepo) + ds.On("MediaFile").Return(mfRepo) ms.On("NewStream", mock.Anything, mock.Anything, stream.Request{Format: "mp3", BitRate: 128}). Return(nil, stream.ErrTooManyTranscodes).Once() @@ -310,7 +310,7 @@ var _ = Describe("Archiver", func() { plRepo := &mockPlaylistRepository{} plRepo.On("GetWithTracks", "1", true, false).Return(pls, nil) - ds.On("Playlist", mock.Anything).Return(plRepo) + ds.On("Playlist").Return(plRepo) ms.On("NewStream", mock.Anything, mock.Anything, stream.Request{Format: "mp3", BitRate: 128}).Return(io.NopCloser(strings.NewReader("test")), nil).Times(2) out := new(bytes.Buffer) @@ -515,17 +515,17 @@ type mockDataStore struct { model.DataStore } -func (m *mockDataStore) MediaFile(ctx context.Context) model.MediaFileRepository { - args := m.Called(ctx) +func (m *mockDataStore) MediaFile() model.MediaFileRepository { + args := m.Called() return args.Get(0).(model.MediaFileRepository) } -func (m *mockDataStore) Playlist(ctx context.Context) model.PlaylistRepository { - args := m.Called(ctx) +func (m *mockDataStore) Playlist() model.PlaylistRepository { + args := m.Called() return args.Get(0).(model.PlaylistRepository) } -func (m *mockDataStore) Library(context.Context) model.LibraryRepository { +func (m *mockDataStore) Library() model.LibraryRepository { return &mockLibraryRepository{} } @@ -534,7 +534,7 @@ type mockLibraryRepository struct { model.LibraryRepository } -func (m *mockLibraryRepository) GetPath(id int) (string, error) { +func (m *mockLibraryRepository) GetPath(_ context.Context, id int) (string, error) { return "/music", nil } @@ -543,7 +543,7 @@ type mockMediaFileRepository struct { model.MediaFileRepository } -func (m *mockMediaFileRepository) GetAll(options ...model.QueryOptions) (model.MediaFiles, error) { +func (m *mockMediaFileRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { args := m.Called(options) return args.Get(0).(model.MediaFiles), args.Error(1) } @@ -553,7 +553,7 @@ type mockPlaylistRepository struct { model.PlaylistRepository } -func (m *mockPlaylistRepository) GetWithTracks(id string, refreshSmartPlaylists, includeMissing bool) (*model.Playlist, error) { +func (m *mockPlaylistRepository) GetWithTracks(_ context.Context, id string, refreshSmartPlaylists, includeMissing bool) (*model.Playlist, error) { args := m.Called(id, refreshSmartPlaylists, includeMissing) return args.Get(0).(*model.Playlist), args.Error(1) } diff --git a/core/artwork/artwork.go b/core/artwork/artwork.go index e27fa118e..2a9ebe771 100644 --- a/core/artwork/artwork.go +++ b/core/artwork/artwork.go @@ -59,21 +59,21 @@ func entityExists(ctx context.Context, ds model.DataStore, artID model.ArtworkID var err error switch artID.Kind { case model.KindArtistArtwork: - found, err = ds.Artist(ctx).Exists(artID.ID) + found, err = ds.Artist().Exists(ctx, artID.ID) case model.KindAlbumArtwork: - found, err = ds.Album(ctx).Exists(artID.ID) + found, err = ds.Album().Exists(ctx, artID.ID) case model.KindMediaFileArtwork: - found, err = ds.MediaFile(ctx).Exists(artID.ID) + found, err = ds.MediaFile().Exists(ctx, artID.ID) case model.KindPlaylistArtwork: - found, err = ds.Playlist(ctx).Exists(artID.ID) + found, err = ds.Playlist().Exists(ctx, artID.ID) case model.KindRadioArtwork: - found, err = ds.Radio(ctx).Exists(artID.ID) + found, err = ds.Radio().Exists(ctx, artID.ID) case model.KindDiscArtwork: albumID, _, perr := model.ParseDiscArtworkID(artID.ID) if perr != nil { return false } - found, err = ds.Album(ctx).Exists(albumID) + found, err = ds.Album().Exists(ctx, albumID) default: return false } @@ -119,7 +119,7 @@ func (s *service) Get(ctx context.Context, artID model.ArtworkID, size int, squa } func (s *service) serveEntity(ctx context.Context, artID model.ArtworkID, size int, square bool) (*Image, error) { - ia, err := s.ds.Artwork(ctx).GetItemArtwork(artID.Kind, artID.ID, model.ImageTypePrimary) + ia, err := s.ds.Artwork().GetItemArtwork(ctx, artID.Kind, artID.ID, model.ImageTypePrimary) switch { case errors.Is(err, model.ErrNotFound): return s.provisional(ctx, artID, size, square) @@ -173,7 +173,7 @@ func (s *service) serveHash(ctx context.Context, artID model.ArtworkID, ia *mode log.Warn(ctx, "Artwork: Stored source is not an image file, re-resolving", "artID", artID, "path", ia.SourcePath) return s.dangling(ctx, artID) } - art, err := s.ds.Artwork(ctx).GetImage(ia.Hash) + art, err := s.ds.Artwork().GetImage(ctx, ia.Hash) if err != nil { if errors.Is(err, model.ErrNotFound) { return s.dangling(ctx, artID) @@ -264,13 +264,13 @@ func (s *service) serveMediaFile(ctx context.Context, artID model.ArtworkID, siz // The setting is not in the config fingerprint, so honor it at serve time: a direct mf- URL // must fall back to disc/album instead of serving stale persisted embedded art. if !conf.Server.EnableMediaFileCoverArt { - mf, err := s.ds.MediaFile(ctx).Get(artID.ID) + mf, err := s.ds.MediaFile().Get(ctx, artID.ID) if err != nil { return nil, err } return s.Get(ctx, mf.DiscCoverArtID(), size, square) } - ia, err := s.ds.Artwork(ctx).GetItemArtwork(model.KindMediaFileArtwork, artID.ID, model.ImageTypePrimary) + ia, err := s.ds.Artwork().GetItemArtwork(ctx, model.KindMediaFileArtwork, artID.ID, model.ImageTypePrimary) switch { case err == nil && ia.Hash != "": return s.serveHash(ctx, artID, ia, size, square) @@ -283,7 +283,7 @@ func (s *service) serveMediaFile(ctx context.Context, artID model.ArtworkID, siz } noRow := errors.Is(err, model.ErrNotFound) - mf, err := s.ds.MediaFile(ctx).Get(artID.ID) + mf, err := s.ds.MediaFile().Get(ctx, artID.ID) if err != nil { return nil, err } @@ -342,7 +342,7 @@ func (s *service) dangling(ctx context.Context, artID model.ArtworkID) (*Image, } func (s *service) enqueue(ctx context.Context, artID model.ArtworkID, priority int) { - err := s.ds.ArtworkQueue(ctx).EnqueuePreservingBackoff(model.ArtworkQueueItem{ + err := s.ds.ArtworkQueue().EnqueuePreservingBackoff(ctx, model.ArtworkQueueItem{ ItemKind: artID.Kind.Prefix(), ItemID: artID.ID, ImageType: model.ImageTypePrimary, diff --git a/core/artwork/artwork_suite_test.go b/core/artwork/artwork_suite_test.go index 93aacc1fa..921a77b89 100644 --- a/core/artwork/artwork_suite_test.go +++ b/core/artwork/artwork_suite_test.go @@ -1,6 +1,7 @@ package artwork import ( + "context" "io/fs" "net/netip" "net/url" @@ -108,15 +109,15 @@ type fakeFolderRepo struct { otherAudioErr error } -func (f *fakeFolderRepo) GetAll(...model.QueryOptions) ([]model.Folder, error) { +func (f *fakeFolderRepo) GetAll(context.Context, ...model.QueryOptions) ([]model.Folder, error) { return f.result, f.err } -func (f *fakeFolderRepo) HasAudioOutsideFolders(model.Folder, []string) (bool, error) { +func (f *fakeFolderRepo) HasAudioOutsideFolders(context.Context, model.Folder, []string) (bool, error) { return f.hasOtherAudio, f.otherAudioErr } -func (f *fakeFolderRepo) Get(string) (*model.Folder, error) { +func (f *fakeFolderRepo) Get(context.Context, string) (*model.Folder, error) { f.getCallCount++ if f.getErr != nil { return nil, f.getErr diff --git a/core/artwork/artwork_test.go b/core/artwork/artwork_test.go index 8b35a872e..2e8ed070e 100644 --- a/core/artwork/artwork_test.go +++ b/core/artwork/artwork_test.go @@ -45,8 +45,8 @@ var _ = Describe("Artwork", func() { hash, err := hashImage(bytes.NewReader(imgBytes)) Expect(err).ToNot(HaveOccurred()) Expect(store.Write(hash, "image/jpeg", bytes.NewReader(imgBytes))).To(Succeed()) - Expect(artRepo.PutImage(&model.Artwork{Hash: hash, Mime: "image/jpeg"})).To(Succeed()) - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: kind, ItemID: id, Hash: hash, Source: "external"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: hash, Mime: "image/jpeg"})).To(Succeed()) + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: kind, ItemID: id, Hash: hash, Source: "external"})).To(Succeed()) seedEntity(kind, id) return hash } @@ -56,9 +56,9 @@ var _ = Describe("Artwork", func() { GinkgoHelper() switch kind { case "al": - Expect(albumRepo.Put(&model.Album{ID: id, Name: "Album"})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{ID: id, Name: "Album"})).To(Succeed()) case "mf": - Expect(mfRepo.Put(&model.MediaFile{ID: id})).To(Succeed()) + Expect(mfRepo.Put(ctx, &model.MediaFile{ID: id})).To(Succeed()) } } @@ -147,9 +147,9 @@ var _ = Describe("Artwork", func() { imgPath := filepath.Join(dir, "cover.jpg") Expect(os.WriteFile(imgPath, coverBytes, 0600)).To(Succeed()) mtime := fileMtime(imgPath) - Expect(artRepo.PutImage(&model.Artwork{Hash: "aaaaaaaaaaaaaaaa", Mime: "image/jpeg"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "aaaaaaaaaaaaaaaa", Mime: "image/jpeg"})).To(Succeed()) seedEntity("al", "al2") - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al2", Hash: "aaaaaaaaaaaaaaaa", Source: "folder", SourcePath: imgPath, RefMtime: mtime, })).To(Succeed()) @@ -163,9 +163,9 @@ var _ = Describe("Artwork", func() { dir := GinkgoT().TempDir() secretPath := filepath.Join(dir, "config.ini") Expect(os.WriteFile(secretPath, []byte("password=secret"), 0600)).To(Succeed()) - Expect(artRepo.PutImage(&model.Artwork{Hash: "dddddddddddddddd", Mime: "image/jpeg"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "dddddddddddddddd", Mime: "image/jpeg"})).To(Succeed()) seedEntity("al", "alni") - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "alni", Hash: "dddddddddddddddd", Source: "folder", SourcePath: secretPath, RefMtime: fileMtime(secretPath), })).To(Succeed()) @@ -180,9 +180,9 @@ var _ = Describe("Artwork", func() { dir := GinkgoT().TempDir() secretPath := filepath.Join(dir, "config.ini") Expect(os.WriteFile(secretPath, secret, 0600)).To(Succeed()) - Expect(artRepo.PutImage(&model.Artwork{Hash: "eeeeeeeeeeeeeeee", Mime: "image/jpeg"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "eeeeeeeeeeeeeeee", Mime: "image/jpeg"})).To(Succeed()) seedEntity("al", "alnic") - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "alnic", Hash: "eeeeeeeeeeeeeeee", Source: "folder", SourcePath: secretPath, RefMtime: fileMtime(secretPath), })).To(Succeed()) @@ -210,9 +210,9 @@ var _ = Describe("Artwork", func() { dir := GinkgoT().TempDir() imgPath := filepath.Join(dir, "cover.jpg") Expect(os.WriteFile(imgPath, coverBytes, 0600)).To(Succeed()) - Expect(artRepo.PutImage(&model.Artwork{Hash: "bbbbbbbbbbbbbbbb", Mime: "image/jpeg"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "bbbbbbbbbbbbbbbb", Mime: "image/jpeg"})).To(Succeed()) seedEntity("al", "al3") - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al3", Hash: "bbbbbbbbbbbbbbbb", Source: "folder", SourcePath: imgPath, RefMtime: fileMtime(imgPath) + 999, })).To(Succeed()) @@ -220,7 +220,7 @@ var _ = Describe("Artwork", func() { _, err := svc.Get(ctx, model.MustParseArtworkID("al-al3"), 0, false) Expect(err).To(MatchError(ErrUnavailable)) Expect(queueRepo.Data[primaryKey("al", "al3")].Priority).To(Equal(model.ArtworkPriorityScan)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al3", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al3", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(Equal("bbbbbbbbbbbbbbbb")) }) @@ -229,9 +229,9 @@ var _ = Describe("Artwork", func() { dir := GinkgoT().TempDir() imgPath := filepath.Join(dir, "cover.jpg") Expect(os.WriteFile(imgPath, coverBytes, 0600)).To(Succeed()) - Expect(artRepo.PutImage(&model.Artwork{Hash: "cccccccccccccccc", Mime: "image/jpeg"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "cccccccccccccccc", Mime: "image/jpeg"})).To(Succeed()) seedEntity("al", "al3b") - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al3b", Hash: "cccccccccccccccc", Source: "folder", SourcePath: imgPath, RefMtime: fileMtime(imgPath) + 999, })).To(Succeed()) @@ -252,7 +252,7 @@ var _ = Describe("Artwork", func() { }) It("never re-enqueues an absent state on view, however old", func() { - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al4", AttemptedAt: time.Now().Add(-365 * 24 * time.Hour), })).To(Succeed()) @@ -272,7 +272,7 @@ var _ = Describe("Artwork", func() { Expect(readAll(img)).To(Equal(coverBytes)) Expect(queueRepo.Data[primaryKey("al", "al5")].Priority).To(Equal(model.ArtworkPriorityBump)) - _, err = artRepo.GetItemArtwork(model.KindAlbumArtwork, "al5", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al5", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -283,7 +283,7 @@ var _ = Describe("Artwork", func() { _, err := svc.Get(ctx, model.MustParseArtworkID("al-al6"), 0, false) Expect(err).To(MatchError(ErrUnavailable)) Expect(queueRepo.Data[primaryKey("al", "al6")].Priority).To(Equal(model.ArtworkPriorityBump)) - _, err = artRepo.GetItemArtwork(model.KindAlbumArtwork, "al6", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al6", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) }) @@ -310,7 +310,7 @@ var _ = Describe("Artwork", func() { It("delegates to the album when the track's state is absent", func() { seedFoundStore("al", "albm", coverBytes) - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "mf", ItemID: "mf2"})).To(Succeed()) + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "mf", ItemID: "mf2"})).To(Succeed()) mfRepo.SetData(model.MediaFiles{{ID: "mf2", AlbumID: "albm"}}) img, err := svc.Get(ctx, model.MustParseArtworkID("mf-mf2"), 0, false) @@ -343,7 +343,7 @@ var _ = Describe("Artwork", func() { Expect(err).ToNot(HaveOccurred()) Expect(len(readAll(img))).To(BeNumerically(">", 0)) Expect(queueRepo.Data[primaryKey("mf", "mf4")].Priority).To(Equal(model.ArtworkPriorityBump)) - _, err = artRepo.GetItemArtwork(model.KindMediaFileArtwork, "mf4", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindMediaFileArtwork, "mf4", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -494,7 +494,7 @@ var _ = Describe("Artwork", func() { }) It("falls back to the artist placeholder for an absent artist", func() { - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "arph"})).To(Succeed()) + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "arph"})).To(Succeed()) img, err := svc.GetOrPlaceholder(ctx, "ar-arph", 300, false) Expect(err).ToNot(HaveOccurred()) @@ -535,7 +535,7 @@ var _ = Describe("EntityExists", func() { artistRepo := tests.CreateMockArtistRepo() artistRepo.SetData(model.Artists{{ID: "ar1"}}) radioRepo := tests.CreateMockedRadioRepo() - Expect(radioRepo.Put(&model.Radio{ID: "ra1", Name: "R"})).To(Succeed()) + Expect(radioRepo.Put(ctx, &model.Radio{ID: "ra1", Name: "R"})).To(Succeed()) ds = &tests.MockDataStore{MockedAlbum: albumRepo, MockedArtist: artistRepo, MockedRadio: radioRepo} }) diff --git a/core/artwork/disc.go b/core/artwork/disc.go index acd8a3740..a050f6685 100644 --- a/core/artwork/disc.go +++ b/core/artwork/disc.go @@ -45,7 +45,7 @@ func newDiscArtworkReader(ctx context.Context, ds model.DataStore, artID model.A return nil, fmt.Errorf("invalid disc artwork id '%s': %w", artID.ID, err) } - al, err := ds.Album(ctx).Get(albumID) + al, err := ds.Album().Get(ctx, albumID) if err != nil { return nil, err } @@ -61,7 +61,7 @@ func newDiscArtworkReader(ctx context.Context, ds model.DataStore, artID model.A } // Query mediafiles for this album + disc to find folder associations and first track - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Sort: "track_number", Order: "ASC", Filters: squirrel.Eq{"album_id": albumID, "disc_number": discNumber}, @@ -88,7 +88,7 @@ func newDiscArtworkReader(ctx context.Context, ds model.DataStore, artID model.A // Resolve folder IDs to library-relative paths discFoldersRel := make(map[string]bool) if len(folderIDs) > 0 { - folders, err := ds.Folder(ctx).GetAll(model.QueryOptions{ + folders, err := ds.Folder().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"folder.id": folderIDs}, }) if err != nil { diff --git a/core/artwork/e2e/acquire_serve_test.go b/core/artwork/e2e/acquire_serve_test.go index 34dfb2cac..8cfa54113 100644 --- a/core/artwork/e2e/acquire_serve_test.go +++ b/core/artwork/e2e/acquire_serve_test.go @@ -44,20 +44,20 @@ var _ = Describe("Acquisition → serve loop", func() { itemFound := func(kind model.Kind, id string) func() bool { return func() bool { - ia, err := artRepo.GetItemArtwork(kind, id, model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, kind, id, model.ImageTypePrimary) return err == nil && ia.Hash != "" } } itemAbsent := func(kind model.Kind, id string) func() bool { return func() bool { - ia, err := artRepo.GetItemArtwork(kind, id, model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, kind, id, model.ImageTypePrimary) return err == nil && ia.Hash == "" } } // Enqueues the way the serving paths do, so the drain is driven by a plain queue row. bump := func(kind, id string) { GinkgoHelper() - Expect(ds.ArtworkQueue(ctx).EnqueuePreservingBackoff(model.ArtworkQueueItem{ + Expect(ds.ArtworkQueue().EnqueuePreservingBackoff(ctx, model.ArtworkQueueItem{ ItemKind: kind, ItemID: id, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBump, })).To(Succeed()) @@ -141,7 +141,7 @@ var _ = Describe("Acquisition → serve loop", func() { bump("al", "al1") runWorkerUntil(ctx, worker, itemFound(model.KindAlbumArtwork, "al1")) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("folder")) @@ -158,7 +158,7 @@ var _ = Describe("Acquisition → serve loop", func() { bump("ar", "ar1") runWorkerUntil(ctx, worker, itemFound(model.KindArtistArtwork, "ar1")) - ia, err := artRepo.GetItemArtwork(model.KindArtistArtwork, "ar1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "ar1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("upload")) @@ -175,14 +175,14 @@ var _ = Describe("Acquisition → serve loop", func() { bump("pl", "pl1") runWorkerUntil(ctx, worker, itemFound(model.KindPlaylistArtwork, "pl1")) - ia, err := artRepo.GetItemArtwork(model.KindPlaylistArtwork, "pl1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindPlaylistArtwork, "pl1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("generated")) img, err := svc.Get(ctx, model.MustParseArtworkID("pl-pl1"), 0, false) Expect(err).ToNot(HaveOccurred()) Expect(img.Hash).To(Equal(ia.Hash)) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) Expect(art.Mime).To(Equal("image/png")) Expect(len(readAll(img))).To(BeNumerically(">", 0)) @@ -194,7 +194,7 @@ var _ = Describe("Acquisition → serve loop", func() { bump("ra", "ra1") runWorkerUntil(ctx, worker, itemFound(model.KindRadioArtwork, "ra1")) - ia, err := artRepo.GetItemArtwork(model.KindRadioArtwork, "ra1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "ra1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("upload")) @@ -216,12 +216,12 @@ var _ = Describe("Acquisition → serve loop", func() { provisionalBytes := readAll(provisional) Expect(len(provisionalBytes)).To(BeNumerically(">", 0)) - _, err = artRepo.GetItemArtwork(model.KindMediaFileArtwork, "mf1", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindMediaFileArtwork, "mf1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "provisional serving must not write a state row") // The provisional read enqueued a Bump; drain it. runWorkerUntil(ctx, worker, itemFound(model.KindMediaFileArtwork, "mf1")) - ia, err := artRepo.GetItemArtwork(model.KindMediaFileArtwork, "mf1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindMediaFileArtwork, "mf1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("embedded")) Expect(ia.Hash).To(Equal(provisional.Hash)) @@ -237,9 +237,9 @@ var _ = Describe("Acquisition → serve loop", func() { bump("al", "al1") runWorkerUntil(ctx, worker, itemFound(model.KindAlbumArtwork, "al1")) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) Expect(art.Mime).To(Equal("image/jpeg")) Expect(art.Width).To(BeNumerically(">", 0)) @@ -259,9 +259,9 @@ var _ = Describe("Acquisition → serve loop", func() { bump("ra", "ra1") runWorkerUntil(ctx, worker, itemFound(model.KindRadioArtwork, "ra1")) - ia, err := artRepo.GetItemArtwork(model.KindRadioArtwork, "ra1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "ra1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) Expect(art.Mime).To(Equal("image/gif")) Expect(art.Width).To(BeNumerically("==", 4)) @@ -279,9 +279,9 @@ var _ = Describe("Acquisition → serve loop", func() { return itemFound(model.KindAlbumArtwork, "al1")() && itemFound(model.KindAlbumArtwork, "al2")() }) - ia1, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al1", model.ImageTypePrimary) + ia1, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) - ia2, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al2", model.ImageTypePrimary) + ia2, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al2", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia1.Hash).To(Equal(ia2.Hash), "identical bytes must share one content hash") Expect(readAll(mustGet(svc.Get(ctx, model.MustParseArtworkID("al-al2"), 0, false)))).To(Equal(coverBytes)) @@ -293,7 +293,7 @@ var _ = Describe("Acquisition → serve loop", func() { bump("ra", "ra1") runWorkerUntil(ctx, worker, itemFound(model.KindRadioArtwork, "ra1")) - ia, err := artRepo.GetItemArtwork(model.KindRadioArtwork, "ra1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "ra1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) staleHash := ia.Hash @@ -308,7 +308,7 @@ var _ = Describe("Acquisition → serve loop", func() { // That failed read enqueued a re-resolution. runWorkerUntil(ctx, worker, func() bool { - cur, gerr := artRepo.GetItemArtwork(model.KindRadioArtwork, "ra1", model.ImageTypePrimary) + cur, gerr := artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "ra1", model.ImageTypePrimary) return gerr == nil && cur.Hash != "" && cur.Hash != staleHash }) img, err := svc.Get(ctx, model.MustParseArtworkID("ra-ra1"), 0, false) diff --git a/core/artwork/e2e/artist_test.go b/core/artwork/e2e/artist_test.go index 4a6959d95..38c28a9b2 100644 --- a/core/artwork/e2e/artist_test.go +++ b/core/artwork/e2e/artist_test.go @@ -201,7 +201,7 @@ var _ = Describe("Artist artwork resolution", func() { uploaded := ar.ID + "_upload.jpg" writeUploadedImage(consts.EntityArtist, uploaded, pngBytes("artist-uploaded")) ar.UploadedImage = uploaded - Expect(rds.Artist(rctx).Put(&ar)).To(Succeed()) + Expect(rds.Artist().Put(rctx, &ar)).To(Succeed()) ia := acquire(model.KindArtistArtwork, ar.ID) Expect(ia.Source).To(Equal("upload")) @@ -279,7 +279,7 @@ var _ = Describe("Artist artwork resolution", func() { func soleArtist() model.Artist { GinkgoHelper() - artists, err := rds.Artist(rctx).GetAll(model.QueryOptions{ + artists, err := rds.Artist().GetAll(rctx, model.QueryOptions{ Filters: squirrel.Eq{"artist.name": "Artist"}, }) Expect(err).ToNot(HaveOccurred()) diff --git a/core/artwork/e2e/e2e_suite_test.go b/core/artwork/e2e/e2e_suite_test.go index 390be14a7..d03e22228 100644 --- a/core/artwork/e2e/e2e_suite_test.go +++ b/core/artwork/e2e/e2e_suite_test.go @@ -67,13 +67,17 @@ type fakeFolderRepo struct { result []model.Folder } -func (f *fakeFolderRepo) GetAll(...model.QueryOptions) ([]model.Folder, error) { return f.result, nil } +func (f *fakeFolderRepo) GetAll(context.Context, ...model.QueryOptions) ([]model.Folder, error) { + return f.result, nil +} -func (f *fakeFolderRepo) HasAudioOutsideFolders(model.Folder, []string) (bool, error) { +func (f *fakeFolderRepo) HasAudioOutsideFolders(context.Context, model.Folder, []string) (bool, error) { return false, nil } -func (f *fakeFolderRepo) Get(string) (*model.Folder, error) { return nil, model.ErrNotFound } +func (f *fakeFolderRepo) Get(context.Context, string) (*model.Folder, error) { + return nil, model.ErrNotFound +} func writeUpload(entityType, name, srcFixture string) string { GinkgoHelper() diff --git a/core/artwork/e2e/mediafile_test.go b/core/artwork/e2e/mediafile_test.go index d26756dcd..3ddb9dc48 100644 --- a/core/artwork/e2e/mediafile_test.go +++ b/core/artwork/e2e/mediafile_test.go @@ -137,7 +137,7 @@ var _ = Describe("MediaFile artwork resolution", func() { func mediafileOn(relPath string) model.MediaFile { GinkgoHelper() - mfs, err := rds.MediaFile(rctx).GetAll(model.QueryOptions{ + mfs, err := rds.MediaFile().GetAll(rctx, model.QueryOptions{ Filters: squirrel.Like{"media_file.path": relPath}, }) Expect(err).ToNot(HaveOccurred()) diff --git a/core/artwork/e2e/playlist_test.go b/core/artwork/e2e/playlist_test.go index f5ac8a7e3..0d862b20e 100644 --- a/core/artwork/e2e/playlist_test.go +++ b/core/artwork/e2e/playlist_test.go @@ -142,13 +142,13 @@ var _ = Describe("Playlist artwork resolution", func() { }) scan() - mfs, err := rds.MediaFile(rctx).GetAll(model.QueryOptions{}) + mfs, err := rds.MediaFile().GetAll(rctx, model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(2)) pl := model.Playlist{ID: "pl-7", Name: "Mix", OwnerID: "admin-1"} pl.AddMediaFilesByID([]string{mfs[0].ID, mfs[1].ID}) - Expect(rds.Playlist(rctx).Put(&pl)).To(Succeed()) + Expect(rds.Playlist().Put(rctx, &pl)).To(Succeed()) ia := acquire(model.KindPlaylistArtwork, pl.ID) Expect(ia.Source).To(Equal("generated")) @@ -180,14 +180,14 @@ var _ = Describe("Playlist artwork resolution", func() { setLayout(layout) scan() - mfs, err := rds.MediaFile(rctx).GetAll(model.QueryOptions{}) + mfs, err := rds.MediaFile().GetAll(rctx, model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(4)) ids := slice.Map(mfs, func(mf model.MediaFile) string { return mf.ID }) pl := model.Playlist{ID: "pl-8", Name: "Four", OwnerID: "admin-1"} pl.AddMediaFilesByID(ids) - Expect(rds.Playlist(rctx).Put(&pl)).To(Succeed()) + Expect(rds.Playlist().Put(rctx, &pl)).To(Succeed()) ia := acquire(model.KindPlaylistArtwork, pl.ID) Expect(ia.Source).To(Equal("generated")) @@ -208,6 +208,6 @@ func putPlaylist(pl model.Playlist) model.Playlist { if pl.OwnerID == "" { pl.OwnerID = "admin-1" } - Expect(rds.Playlist(rctx).Put(&pl)).To(Succeed()) + Expect(rds.Playlist().Put(rctx, &pl)).To(Succeed()) return pl } diff --git a/core/artwork/e2e/radio_test.go b/core/artwork/e2e/radio_test.go index bba85224a..72af215d1 100644 --- a/core/artwork/e2e/radio_test.go +++ b/core/artwork/e2e/radio_test.go @@ -23,7 +23,7 @@ var _ = Describe("Radio artwork resolution", func() { It("returns the uploaded image bytes", func() { writeUploadedImage(consts.EntityRadio, "rd-1_logo.jpg", pngBytes("radio-logo")) rd := model.Radio{ID: "rd-1", Name: "Test Radio", StreamUrl: "https://example.com/stream", UploadedImage: "rd-1_logo.jpg"} - Expect(rds.Radio(rctx).Put(&rd)).To(Succeed()) + Expect(rds.Radio().Put(rctx, &rd)).To(Succeed()) ia := acquire(model.KindRadioArtwork, rd.ID) Expect(ia.Source).To(Equal("upload")) @@ -35,7 +35,7 @@ var _ = Describe("Radio artwork resolution", func() { // (no files on disk — the resolver has no sources to fall back to) It("settles absent", func() { rd := model.Radio{ID: "rd-2", Name: "Bare Radio", StreamUrl: "https://example.com/stream"} - Expect(rds.Radio(rctx).Put(&rd)).To(Succeed()) + Expect(rds.Radio().Put(rctx, &rd)).To(Succeed()) ia := acquire(model.KindRadioArtwork, rd.ID) Expect(ia.Hash).To(BeEmpty()) diff --git a/core/artwork/e2e/resolution_harness_test.go b/core/artwork/e2e/resolution_harness_test.go index fff62ce69..fbd89925e 100644 --- a/core/artwork/e2e/resolution_harness_test.go +++ b/core/artwork/e2e/resolution_harness_test.go @@ -99,11 +99,11 @@ func setupResolutionHarness() { rds = &tests.MockDataStore{RealDS: persistence.New(db.Db())} adminUser := model.User{ID: "admin-1", UserName: "admin", Name: "Admin", IsAdmin: true, NewPassword: "password"} - Expect(rds.User(rctx).Put(&adminUser)).To(Succeed()) + Expect(rds.User().Put(rctx, &adminUser)).To(Succeed()) lib := model.Library{ID: 1, Name: "Music", Path: fakeLibPath} - Expect(rds.Library(rctx).Put(&lib)).To(Succeed()) - Expect(rds.User(rctx).SetUserLibraries(adminUser.ID, []int{lib.ID})).To(Succeed()) + Expect(rds.Library().Put(rctx, &lib)).To(Succeed()) + Expect(rds.User().SetUserLibraries(rctx, adminUser.ID, []int{lib.ID})).To(Succeed()) loadEmbeddedFixture() @@ -140,13 +140,13 @@ func scan() { func acquire(kind model.Kind, id string) model.ItemArtwork { GinkgoHelper() // Enqueues the way the serving paths do, so the drain is driven by a plain queue row. - Expect(rds.ArtworkQueue(rctx).EnqueuePreservingBackoff(model.ArtworkQueueItem{ + Expect(rds.ArtworkQueue().EnqueuePreservingBackoff(rctx, model.ArtworkQueueItem{ ItemKind: kind.Prefix(), ItemID: id, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBump, })).To(Succeed()) var ia *model.ItemArtwork runResolutionWorkerUntil(func() bool { - got, err := rds.Artwork(rctx).GetItemArtwork(kind, id, model.ImageTypePrimary) + got, err := rds.Artwork().GetItemArtwork(rctx, kind, id, model.ImageTypePrimary) if err != nil { return false } @@ -211,7 +211,7 @@ func expectAlbumFolderCover(al model.Album, suffix string) { // A drain settles every ready item, so byte-level folder assertions must precede any acquire. func requireNoStateRow(kind model.Kind, id string) { GinkgoHelper() - _, err := rds.Artwork(rctx).GetItemArtwork(kind, id, model.ImageTypePrimary) + _, err := rds.Artwork().GetItemArtwork(rctx, kind, id, model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "assert %s %q before acquiring any other entity in this spec", kind, id) } @@ -266,7 +266,7 @@ func gridQuadrants(data []byte) [4]color.RGBA { // Store-backed sources only (embedded/generated); file-backed ones assert on ia.SourcePath. func storedBytes(ia model.ItemArtwork) []byte { GinkgoHelper() - art, err := rds.Artwork(rctx).GetImage(ia.Hash) + art, err := rds.Artwork().GetImage(rctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) r, err := rstore.Open(ia.Hash, art.Mime) Expect(err).ToNot(HaveOccurred()) @@ -345,7 +345,7 @@ func replaceWithRealMP3(relPath string) { func firstAlbum() model.Album { GinkgoHelper() - albums, err := rds.Album(rctx).GetAll(model.QueryOptions{}) + albums, err := rds.Album().GetAll(rctx, model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) Expect(albums).To(HaveLen(1), "expected exactly one album, got %d", len(albums)) return albums[0] @@ -353,7 +353,7 @@ func firstAlbum() model.Album { func albumByName(name string) model.Album { GinkgoHelper() - albums, err := rds.Album(rctx).GetAll(model.QueryOptions{}) + albums, err := rds.Album().GetAll(rctx, model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) for _, al := range albums { if al.Name == name { diff --git a/core/artwork/folders_album.go b/core/artwork/folders_album.go index 88f0181a3..7b423adb6 100644 --- a/core/artwork/folders_album.go +++ b/core/artwork/folders_album.go @@ -36,7 +36,7 @@ func loadAlbumFoldersPaths(ctx context.Context, ds model.DataStore, album model. } func loadFolders(ctx context.Context, ds model.DataStore, folderIDs []string) ([]model.Folder, error) { - return ds.Folder(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"folder.id": folderIDs, "missing": false}}) + return ds.Folder().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"folder.id": folderIDs, "missing": false}}) } // folderImages collects the folders' image files, sorted so files without @@ -79,7 +79,7 @@ func albumRootParent(ctx context.Context, ds model.DataStore, folders []model.Fo if len(folders) < 2 && anyFolderHasImages(folders) { return nil, nil } - parent, err := ds.Folder(ctx).Get(commonParentID) + parent, err := ds.Folder().Get(ctx, commonParentID) if errors.Is(err, model.ErrNotFound) { log.Warn(ctx, "Artwork: Parent folder not found for album cover art lookup", "parentID", commonParentID) return nil, nil @@ -91,7 +91,7 @@ func albumRootParent(ctx context.Context, ds model.DataStore, folders []model.Fo // The library root can never be an album root return nil, nil } - hasOtherAudio, err := ds.Folder(ctx).HasAudioOutsideFolders(*parent, folderIDs) + hasOtherAudio, err := ds.Folder().HasAudioOutsideFolders(ctx, *parent, folderIDs) if err != nil { return nil, err } diff --git a/core/artwork/folders_artist.go b/core/artwork/folders_artist.go index efef42f81..3403935db 100644 --- a/core/artwork/folders_artist.go +++ b/core/artwork/folders_artist.go @@ -169,14 +169,14 @@ func loadArtistFolder(ctx context.Context, ds model.DataStore, albums model.Albu } // Cleaned like the album paths; Join keeps an empty path empty, Clean would return ".". - libPath, _ := ds.Library(ctx).GetPath(libID) + libPath, _ := ds.Library().GetPath(ctx, libID) libPath = filepath.Join(libPath) folderID := model.FolderID(model.Library{ID: libID, Path: libPath}, folderPath) log.Trace(ctx, "Artwork: Calculating artist folder details", "folderPath", folderPath, "folderID", folderID, "libPath", libPath, "libID", libID, "albumPaths", paths) - folders, err := ds.Folder(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"folder.id": folderID, "missing": false}}) + folders, err := ds.Folder().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"folder.id": folderID, "missing": false}}) if err != nil || len(folders) == 0 { log.Warn(ctx, "Artwork: Could not find folder for artist", "folderPath", folderPath, "id", folderID, "libPath", libPath, "libID", libID, err) diff --git a/core/artwork/housekeeping.go b/core/artwork/housekeeping.go index 3ca452bdc..5ec701541 100644 --- a/core/artwork/housekeeping.go +++ b/core/artwork/housekeeping.go @@ -70,7 +70,7 @@ func ConfigFingerprint() string { // resolved under it. Nothing re-resolves on its own; applying a change is an explicit reprocess. func ReconcileConfigFingerprint(ctx context.Context, ds model.DataStore) error { current := ConfigFingerprint() - stored, err := ds.Property(ctx).DefaultGet(consts.ArtConfFingerprintPropertyKey, "") + stored, err := ds.Property().DefaultGet(ctx, consts.ArtConfFingerprintPropertyKey, "") if err != nil { return err } @@ -89,14 +89,14 @@ func ReconcileConfigFingerprint(ctx context.Context, ds model.DataStore) error { // MarkConfigApplied records the current fingerprint as the one the library is resolved under. func MarkConfigApplied(ctx context.Context, ds model.DataStore) error { - return ds.Property(ctx).Put(consts.ArtConfFingerprintPropertyKey, ConfigFingerprint()) + return ds.Property().Put(ctx, consts.ArtConfFingerprintPropertyKey, ConfigFingerprint()) } // enqueueMissingAll is the safety net for entities a scan never enqueued (added between scans, or scanner off). func enqueueMissingAll(ctx context.Context, ds model.DataStore) error { - queue := ds.ArtworkQueue(ctx) + queue := ds.ArtworkQueue() for _, kind := range ReprocessKinds { - if _, err := queue.EnqueueAllMissing(kind, model.ArtworkPriorityRecheck); err != nil { + if _, err := queue.EnqueueAllMissing(ctx, kind, model.ArtworkPriorityRecheck); err != nil { return err } } @@ -108,31 +108,31 @@ func enqueueMissingAll(ctx context.Context, ds model.DataStore) error { func ItemName(ctx context.Context, ds model.DataStore, kind model.Kind, id string) (string, error) { switch kind { case model.KindArtistArtwork: - ar, err := ds.Artist(ctx).Get(id) + ar, err := ds.Artist().Get(ctx, id) if err != nil { return "", err } return ar.Name, nil case model.KindAlbumArtwork: - al, err := ds.Album(ctx).Get(id) + al, err := ds.Album().Get(ctx, id) if err != nil { return "", err } return al.Name, nil case model.KindPlaylistArtwork: - pls, err := ds.Playlist(ctx).Get(id) + pls, err := ds.Playlist().Get(ctx, id) if err != nil { return "", err } return pls.Name, nil case model.KindRadioArtwork: - rd, err := ds.Radio(ctx).Get(id) + rd, err := ds.Radio().Get(ctx, id) if err != nil { return "", err } return rd.Name, nil case model.KindMediaFileArtwork: - mf, err := ds.MediaFile(ctx).Get(id) + mf, err := ds.MediaFile().Get(ctx, id) if err != nil { return "", err } @@ -148,7 +148,7 @@ func discArtworkName(ctx context.Context, ds model.DataStore, id string) (string if err != nil { return "", err } - al, err := ds.Album(ctx).Get(albumID) + al, err := ds.Album().Get(ctx, albumID) if err != nil { return "", err } @@ -162,11 +162,11 @@ func discArtworkName(ctx context.Context, ds model.DataStore, id string) (string // Refresh drops an item's resolved artwork state and re-queues it at Bump priority. func Refresh(ctx context.Context, ds model.DataStore, kind model.Kind, id string) error { - if err := ds.Artwork(ctx).DeleteForItems(kind, []string{id}); err != nil { + if err := ds.Artwork().DeleteForItems(ctx, kind, []string{id}); err != nil { return fmt.Errorf("clearing artwork state: %w", err) } item := model.ArtworkQueueItem{ItemKind: kind.Prefix(), ItemID: id, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBump} - if err := ds.ArtworkQueue(ctx).Enqueue(item); err != nil { + if err := ds.ArtworkQueue().Enqueue(ctx, item); err != nil { return fmt.Errorf("enqueuing artwork refresh: %w", err) } return nil diff --git a/core/artwork/housekeeping_test.go b/core/artwork/housekeeping_test.go index 2027cba9a..0f501d0b1 100644 --- a/core/artwork/housekeeping_test.go +++ b/core/artwork/housekeeping_test.go @@ -95,15 +95,15 @@ var _ = Describe("Housekeeping", func() { It("records the current fingerprint when none was ever stored", func() { Expect(ReconcileConfigFingerprint(ctx, ds)).To(Succeed()) - Expect(propRepo.Get(consts.ArtConfFingerprintPropertyKey)).To(Equal(ConfigFingerprint())) + Expect(propRepo.Get(ctx, consts.ArtConfFingerprintPropertyKey)).To(Equal(ConfigFingerprint())) }) It("leaves a stale fingerprint stored, so the warning survives a restart", func() { - Expect(propRepo.Put(consts.ArtConfFingerprintPropertyKey, "stale-fingerprint")).To(Succeed()) + Expect(propRepo.Put(ctx, consts.ArtConfFingerprintPropertyKey, "stale-fingerprint")).To(Succeed()) Expect(ReconcileConfigFingerprint(ctx, ds)).To(Succeed()) - Expect(propRepo.Get(consts.ArtConfFingerprintPropertyKey)).To(Equal("stale-fingerprint")) + Expect(propRepo.Get(ctx, consts.ArtConfFingerprintPropertyKey)).To(Equal("stale-fingerprint")) }) }) @@ -153,7 +153,7 @@ var _ = Describe("ItemName", func() { {ID: "al-2", Name: "Sandinista!", Discs: model.Discs{2: "Side Three"}}, }) ds = &tests.MockDataStore{MockedAlbum: albumRepo} - Expect(ds.Artist(ctx).(*tests.MockArtistRepo).Put(&model.Artist{ID: "ar-1", Name: "Radiohead"})).To(Succeed()) + Expect(ds.Artist().(*tests.MockArtistRepo).Put(ctx, &model.Artist{ID: "ar-1", Name: "Radiohead"})).To(Succeed()) }) It("returns the album name", func() { diff --git a/core/artwork/library_fs.go b/core/artwork/library_fs.go index 6e6099325..84e088c68 100644 --- a/core/artwork/library_fs.go +++ b/core/artwork/library_fs.go @@ -30,7 +30,7 @@ func (v libraryView) Abs(rel string) string { // loadLibraryView resolves the MusicFS and absolute root path in a single // library lookup. func loadLibraryView(ctx context.Context, ds model.DataStore, libID int) (libraryView, error) { - lib, err := ds.Library(ctx).Get(libID) + lib, err := ds.Library().Get(ctx, libID) if err != nil { return libraryView{}, err } diff --git a/core/artwork/library_fs_test.go b/core/artwork/library_fs_test.go index 22498e7a1..868c5e45a 100644 --- a/core/artwork/library_fs_test.go +++ b/core/artwork/library_fs_test.go @@ -24,7 +24,7 @@ var _ = Describe("loadLibraryView", Ordered, func() { }) It("returns a view for a library backed by registered storage", func() { - Expect(ds.Library(ctx).Put(&model.Library{ID: 1, Path: "fake:///music"})).To(Succeed()) + Expect(ds.Library().Put(ctx, &model.Library{ID: 1, Path: "fake:///music"})).To(Succeed()) lib, err := loadLibraryView(ctx, ds, 1) Expect(err).ToNot(HaveOccurred()) @@ -45,7 +45,7 @@ var _ = Describe("loadLibraryView", Ordered, func() { }) It("returns an error when the library path uses an unregistered scheme", func() { - Expect(ds.Library(ctx).Put(&model.Library{ID: 2, Path: "unsupported:///music"})).To(Succeed()) + Expect(ds.Library().Put(ctx, &model.Library{ID: 2, Path: "unsupported:///music"})).To(Succeed()) _, err := loadLibraryView(ctx, ds, 2) Expect(err).To(HaveOccurred()) }) diff --git a/core/artwork/processor.go b/core/artwork/processor.go index cf2176775..78e625e94 100644 --- a/core/artwork/processor.go +++ b/core/artwork/processor.go @@ -82,7 +82,7 @@ type processor struct { // acquire resolves one queue item end to end: find an image, hash/decode/ // blurhash it, place its bytes, and persist the resulting state. func (p *processor) acquire(ctx context.Context, item model.ArtworkQueueItem) (out outcome, got *acquired, retryIn time.Duration) { - repo := p.ds.Artwork(ctx) + repo := p.ds.Artwork() start := time.Now() defer func() { log.Debug(ctx, "Artwork: Acquisition finished", "kind", item.ItemKind, "id", item.ItemID, @@ -137,7 +137,7 @@ func (p *processor) acquire(ctx context.Context, item model.ArtworkQueueItem) (o log.Trace(ctx, "Artwork: Hashed image", "kind", item.ItemKind, "id", item.ItemID, "hash", hash, "bytes", len(data), "elapsed", time.Since(hashStart)) - art, err := repo.GetImage(hash) + art, err := repo.GetImage(ctx, hash) switch { case err == nil && art.Width > 0: log.Debug(ctx, "Artwork: Reusing a known image, skipping decode", "kind", item.ItemKind, @@ -195,7 +195,7 @@ func (p *processor) persist(ctx context.Context, repo model.ArtworkRepository, i if err != nil { return nil, fmt.Errorf("writing image store: %w", err) } - if err := repo.PutImage(art); err != nil { + if err := repo.PutImage(ctx, art); err != nil { return nil, fmt.Errorf("persisting artwork image: %w", err) } ia := &model.ItemArtwork{ @@ -210,7 +210,7 @@ func (p *processor) persist(ctx context.Context, repo model.ArtworkRepository, i Trace: traceFrom(ctx).encode(sourcePath), } // PutItemArtwork stamps UpdatedAt on ia, so the returned struct matches the persisted row. - if err := repo.PutItemArtwork(ia); err != nil { + if err := repo.PutItemArtwork(ctx, ia); err != nil { return nil, fmt.Errorf("persisting item artwork state: %w", err) } return ia, nil @@ -218,7 +218,7 @@ func (p *processor) persist(ctx context.Context, repo model.ArtworkRepository, i // writeAbsent records a known-absent state: every source answered definitively "no". func writeAbsent(ctx context.Context, repo model.ArtworkRepository, item model.ArtworkQueueItem) outcome { - err := repo.PutItemArtwork(&model.ItemArtwork{ + err := repo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: item.ItemKind, ItemID: item.ItemID, ImageType: item.ImageType, diff --git a/core/artwork/processor_test.go b/core/artwork/processor_test.go index 554ca08dc..8eec97cde 100644 --- a/core/artwork/processor_test.go +++ b/core/artwork/processor_test.go @@ -93,14 +93,14 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al1"}) Expect(out).To(Equal(outcomeFound)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).ToNot(BeEmpty()) Expect(ia.Source).To(Equal("folder")) Expect(filepath.ToSlash(ia.SourcePath)).To(HaveSuffix("tests/fixtures/artist/an-album/cover.jpg")) Expect(ia.RefMtime).To(BeNumerically(">", 0)) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) // Every placeholder is derived from the one shared thumbnail, so all three land together. Expect(art.BlurHash).ToNot(BeEmpty()) @@ -156,12 +156,12 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al2"}) Expect(out).To(Equal(outcomeFound)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al2", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al2", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("embedded")) Expect(filepath.ToSlash(ia.SourcePath)).To(HaveSuffix("tests/fixtures/artist/an-album/test.mp3")) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) Expect(art.BlurHash).ToNot(BeEmpty()) @@ -179,7 +179,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al3"}) Expect(out).To(Equal(outcomeAbsent)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al3", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al3", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(BeEmpty()) Expect(ia.Source).To(BeEmpty()) @@ -200,7 +200,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al-io"}) Expect(out).To(Equal(outcomeFailed)) - _, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al-io", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al-io", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "an I/O fault must not be recorded as absent") }) @@ -225,7 +225,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "ra", ItemID: "ra-io"}) Expect(out).To(Equal(outcomeFailed)) - _, err := artRepo.GetItemArtwork(model.KindRadioArtwork, "ra-io", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "ra-io", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "an unreadable upload must not be recorded as absent") }) @@ -269,7 +269,7 @@ var _ = Describe("processor.acquire", func() { Expect(out).To(Equal(outcomeFailed)) Expect(retryIn).To(BeZero(), "a plain failure asks for no particular delay") - _, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al4", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al4", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -299,7 +299,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alstale"}) Expect(out).To(Equal(outcomeFoundStale)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alstale", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alstale", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).ToNot(BeEmpty()) Expect(ia.Source).To(Equal("folder")) @@ -316,10 +316,10 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alU"}) Expect(out).To(Equal(outcomeFound)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alU", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alU", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("folder")) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) Expect(art.Width).To(BeZero()) Expect(art.BlurHash).To(BeEmpty()) @@ -336,7 +336,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alE"}) Expect(out).To(Equal(outcomeFailed)) - _, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alE", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alE", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -354,7 +354,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alX"}) Expect(out).To(Equal(outcomeFailed)) - _, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alX", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alX", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -373,12 +373,12 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alext"}) Expect(out).To(Equal(outcomeFound)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alext", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alext", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("external:deezerFake")) Expect(ia.Hash).ToNot(BeEmpty()) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) rc, err := store.Open(ia.Hash, art.Mime) Expect(err).ToNot(HaveOccurred()) @@ -397,7 +397,7 @@ var _ = Describe("processor.acquire", func() { out1, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al5"}) Expect(out1).To(Equal(outcomeFound)) - ia1, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al5", model.ImageTypePrimary) + ia1, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al5", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) // A re-decode instead of a hash dedup would overwrite this sentinel. @@ -407,11 +407,11 @@ var _ = Describe("processor.acquire", func() { out2, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al6"}) Expect(out2).To(Equal(outcomeFound)) - ia2, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al6", model.ImageTypePrimary) + ia2, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al6", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia2.Hash).To(Equal(ia1.Hash)) - reused, err := artRepo.GetImage(ia1.Hash) + reused, err := artRepo.GetImage(ctx, ia1.Hash) Expect(err).ToNot(HaveOccurred()) Expect(reused.BlurHash).To(Equal("SENTINEL")) }) @@ -437,7 +437,7 @@ var _ = Describe("processor.acquire", func() { folderRepo.result = []model.Folder{{Path: "album-a", ImageFiles: []string{"cover.jpg"}}} outN, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alA"}) Expect(outN).To(Equal(outcomeFound)) - iaA, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alA", model.ImageTypePrimary) + iaA, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alA", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(iaA.Source).To(Equal("folder")) Expect(filepath.ToSlash(iaA.SourcePath)).To(HaveSuffix("album-a/cover.jpg")) @@ -451,19 +451,19 @@ var _ = Describe("processor.acquire", func() { folderRepo.result = []model.Folder{{Path: "album-b", ImageFiles: []string{"cover.jpg"}}} outN, _, _ = proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alB"}) Expect(outN).To(Equal(outcomeFound)) - iaB, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alB", model.ImageTypePrimary) + iaB, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alB", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(iaB.Hash).To(Equal(iaA.Hash)) Expect(filepath.ToSlash(iaB.SourcePath)).To(HaveSuffix("album-b/cover.jpg")) Expect(iaB.RefMtime).To(Equal(time.Unix(2000, 0).UnixNano())) - iaAafter, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alA", model.ImageTypePrimary) + iaAafter, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alA", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(filepath.ToSlash(iaAafter.SourcePath)).To(HaveSuffix("album-a/cover.jpg")) Expect(iaAafter.RefMtime).To(Equal(time.Unix(1000, 0).UnixNano())) Expect(artRepo.Data).To(HaveLen(1)) - reused, err := artRepo.GetImage(iaA.Hash) + reused, err := artRepo.GetImage(ctx, iaA.Hash) Expect(err).ToNot(HaveOccurred()) Expect(reused.BlurHash).To(Equal("SENTINEL")) }) @@ -483,7 +483,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "ra", ItemID: "ra1"}) Expect(out).To(Equal(outcomeFailed)) - _, err := artRepo.GetItemArtwork(model.KindRadioArtwork, "ra1", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "ra1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -504,7 +504,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "ra", ItemID: "big"}) Expect(out).To(Equal(outcomeFailed)) - _, err = artRepo.GetItemArtwork(model.KindRadioArtwork, "big", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindRadioArtwork, "big", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -565,12 +565,12 @@ var _ = Describe("processor.acquire", func() { hash, err := hashImage(bytes.NewReader(imgBytes)) Expect(err).ToNot(HaveOccurred()) - Expect(artRepo.PutImage(&model.Artwork{Hash: hash, Mime: "application/octet-stream"})).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: hash, Mime: "application/octet-stream"})).To(Succeed()) out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alM"}) Expect(out).To(Equal(outcomeFound)) - upgraded, err := artRepo.GetImage(hash) + upgraded, err := artRepo.GetImage(ctx, hash) Expect(err).ToNot(HaveOccurred()) Expect(upgraded.Width).To(BeNumerically(">", 0)) Expect(upgraded.BlurHash).ToNot(BeEmpty()) @@ -590,7 +590,7 @@ var _ = Describe("processor.acquire", func() { out, _, _ := proc.acquire(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al7"}) Expect(out).To(Equal(outcomeFailed)) - _, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al7", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al7", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) }) diff --git a/core/artwork/prune.go b/core/artwork/prune.go index c79458b8f..15b4a5673 100644 --- a/core/artwork/prune.go +++ b/core/artwork/prune.go @@ -14,9 +14,9 @@ const pruneMinAge = time.Hour func prune(ctx context.Context, ds model.DataStore, store *ImageStore) error { start := time.Now() defer func() { log.Debug(ctx, "Artwork: Prune finished", "elapsed", time.Since(start)) }() - repo := ds.Artwork(ctx) + repo := ds.Artwork() - purged, err := repo.PurgeDanglingItems() + purged, err := repo.PurgeDanglingItems(ctx) if err != nil { return err } @@ -25,7 +25,7 @@ func prune(ctx context.Context, ds model.DataStore, store *ImageStore) error { } // Queue rows for deleted entities would otherwise retry forever (Get -> not found -> failed). - queuePurged, err := ds.ArtworkQueue(ctx).PurgeDangling() + queuePurged, err := ds.ArtworkQueue().PurgeDangling(ctx) if err != nil { return err } @@ -35,7 +35,7 @@ func prune(ctx context.Context, ds model.DataStore, store *ImageStore) error { // Files younger than the grace window may belong to acquisitions whose rows aren't committed yet. cutoff := time.Now().Add(-pruneMinAge) - orphans, err := repo.PurgeOrphans(cutoff) + orphans, err := repo.PurgeOrphans(ctx, cutoff) if err != nil { return err } @@ -44,7 +44,7 @@ func prune(ctx context.Context, ds model.DataStore, store *ImageStore) error { } // Read after the delete, so the sweep below reclaims the files of the rows just removed. - mimes, err := repo.GetMimeByHash() + mimes, err := repo.GetMimeByHash(ctx) if err != nil { return err } diff --git a/core/artwork/prune_test.go b/core/artwork/prune_test.go index 8b8504f89..674211dfa 100644 --- a/core/artwork/prune_test.go +++ b/core/artwork/prune_test.go @@ -18,18 +18,20 @@ type flakyGetArtworkRepo struct { *tests.MockArtworkRepo } -func (f *flakyGetArtworkRepo) GetMimeByHash() (map[string]string, error) { +func (f *flakyGetArtworkRepo) GetMimeByHash(context.Context) (map[string]string, error) { return nil, errors.New("db locked") } var _ = Describe("Prune", func() { + var ctx context.Context var ds *tests.MockDataStore var store *ImageStore var awRepo *tests.MockArtworkRepo BeforeEach(func() { + ctx = GinkgoT().Context() ds = &tests.MockDataStore{} - awRepo = ds.Artwork(context.Background()).(*tests.MockArtworkRepo) + awRepo = ds.Artwork().(*tests.MockArtworkRepo) store = NewImageStore(GinkgoT().TempDir()) }) @@ -41,9 +43,9 @@ var _ = Describe("Prune", func() { } It("purges dangling item_artwork state for gone entities, summed across kinds", func() { - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "gone-album", ImageType: model.ImageTypePrimary})).To(Succeed()) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "gone-artist", ImageType: model.ImageTypePrimary})).To(Succeed()) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "live-artist", ImageType: model.ImageTypePrimary})).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "gone-album", ImageType: model.ImageTypePrimary})).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "gone-artist", ImageType: model.ImageTypePrimary})).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "live-artist", ImageType: model.ImageTypePrimary})).To(Succeed()) awRepo.ExistingIDs = map[string]map[string]bool{ "al": {}, "ar": {"live-artist": true}, @@ -51,17 +53,17 @@ var _ = Describe("Prune", func() { Expect(prune(context.Background(), ds, store)).To(Succeed()) - _, err := awRepo.GetItemArtwork(model.KindAlbumArtwork, "gone-album", model.ImageTypePrimary) + _, err := awRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "gone-album", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) - _, err = awRepo.GetItemArtwork(model.KindArtistArtwork, "gone-artist", model.ImageTypePrimary) + _, err = awRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "gone-artist", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) - _, err = awRepo.GetItemArtwork(model.KindArtistArtwork, "live-artist", model.ImageTypePrimary) + _, err = awRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "live-artist", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) }) It("purges dangling artwork_queue rows for gone entities", func() { queueRepo := tests.CreateMockArtworkQueueRepo() - Expect(queueRepo.Enqueue( + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "gone-album", ImageType: model.ImageTypePrimary}, model.ArtworkQueueItem{ItemKind: "al", ItemID: "live-album", ImageType: model.ImageTypePrimary}, )).To(Succeed()) @@ -80,17 +82,17 @@ var _ = Describe("Prune", func() { Expect(store.Write(h, "image/jpeg", bytes.NewReader(data))).To(Succeed()) old := time.Now().Add(-2 * time.Hour) Expect(os.Chtimes(store.path(h, "image/jpeg"), old, old)).To(Succeed()) - Expect(awRepo.PutImage(&model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) ageArtwork(h, old) kept := []byte("kept-bytes") hk, _ := hashImage(bytes.NewReader(kept)) Expect(store.Write(hk, "image/jpeg", bytes.NewReader(kept))).To(Succeed()) - Expect(awRepo.PutImage(&model.Artwork{Hash: hk, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: hk, Mime: "image/jpeg"})).To(Succeed()) Expect(prune(context.Background(), ds, store)).To(Succeed()) - _, err := awRepo.GetImage(h) + _, err := awRepo.GetImage(ctx, h) Expect(err).To(MatchError(model.ErrNotFound)) _, err = store.Open(h, "image/jpeg") Expect(os.IsNotExist(err)).To(BeTrue()) @@ -103,14 +105,14 @@ var _ = Describe("Prune", func() { data := []byte("reacquired-bytes") h, _ := hashImage(bytes.NewReader(data)) Expect(store.Write(h, "image/jpeg", bytes.NewReader(data))).To(Succeed()) - Expect(awRepo.PutImage(&model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) ageArtwork(h, time.Now().Add(-2*time.Hour)) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "a1", + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "a1", ImageType: model.ImageTypePrimary, Hash: h, Source: "folder"})).To(Succeed()) Expect(prune(context.Background(), ds, store)).To(Succeed()) - _, err := awRepo.GetImage(h) + _, err := awRepo.GetImage(ctx, h) Expect(err).ToNot(HaveOccurred()) rc, err := store.Open(h, "image/jpeg") Expect(err).ToNot(HaveOccurred()) @@ -122,11 +124,11 @@ var _ = Describe("Prune", func() { h, _ := hashImage(bytes.NewReader(data)) Expect(store.Write(h, "image/jpeg", bytes.NewReader(data))).To(Succeed()) // Reacquisition refreshed created_at, so the row is unreferenced but too young to drop. - Expect(awRepo.PutImage(&model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) Expect(prune(context.Background(), ds, store)).To(Succeed()) - _, err := awRepo.GetImage(h) + _, err := awRepo.GetImage(ctx, h) Expect(err).ToNot(HaveOccurred()) rc, err := store.Open(h, "image/jpeg") Expect(err).ToNot(HaveOccurred()) @@ -137,7 +139,7 @@ var _ = Describe("Prune", func() { data := []byte("racing-bytes") h, _ := hashImage(bytes.NewReader(data)) Expect(store.Write(h, "image/jpeg", bytes.NewReader(data))).To(Succeed()) - Expect(awRepo.PutImage(&model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) ageArtwork(h, time.Now().Add(-2*time.Hour)) // The row is orphaned, but a concurrent acquisition just touched the file's mtime. @@ -170,7 +172,7 @@ var _ = Describe("Prune", func() { Expect(os.Chtimes(store.path(h, "image/png"), old, old)).To(Succeed()) Expect(os.Chtimes(store.path(h, "image/jpeg"), old, old)).To(Succeed()) // The row records the current mime; the .png file is a superseded variant. - Expect(awRepo.PutImage(&model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: h, Mime: "image/jpeg"})).To(Succeed()) Expect(prune(context.Background(), ds, store)).To(Succeed()) @@ -192,14 +194,14 @@ var _ = Describe("Prune", func() { hb, _ := hashImage(bytes.NewReader(blocked)) Expect(store.Write(hb, "image/jpeg", bytes.NewReader(blocked))).To(Succeed()) Expect(os.Chtimes(store.path(hb, "image/jpeg"), old, old)).To(Succeed()) - Expect(awRepo.PutImage(&model.Artwork{Hash: hb, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: hb, Mime: "image/jpeg"})).To(Succeed()) ageArtwork(hb, old) good := []byte("good-bytes") hg, _ := hashImage(bytes.NewReader(good)) Expect(store.Write(hg, "image/jpeg", bytes.NewReader(good))).To(Succeed()) Expect(os.Chtimes(store.path(hg, "image/jpeg"), old, old)).To(Succeed()) - Expect(awRepo.PutImage(&model.Artwork{Hash: hg, Mime: "image/jpeg"})).To(Succeed()) + Expect(awRepo.PutImage(ctx, &model.Artwork{Hash: hg, Mime: "image/jpeg"})).To(Succeed()) ageArtwork(hg, old) // A read-only shard directory makes os.Remove fail (EACCES) for hb's file only. @@ -210,13 +212,13 @@ var _ = Describe("Prune", func() { Expect(prune(context.Background(), ds, store)).To(Succeed()) - _, err := awRepo.GetImage(hg) + _, err := awRepo.GetImage(ctx, hg) Expect(err).To(MatchError(model.ErrNotFound)) _, err = store.Open(hg, "image/jpeg") Expect(os.IsNotExist(err)).To(BeTrue()) // The row purge does not depend on file removal, so only the file survives. - _, err = awRepo.GetImage(hb) + _, err = awRepo.GetImage(ctx, hb) Expect(err).To(MatchError(model.ErrNotFound)) rc, err := store.Open(hb, "image/jpeg") Expect(err).ToNot(HaveOccurred()) diff --git a/core/artwork/resolve.go b/core/artwork/resolve.go index f4f3ef725..40baa2495 100644 --- a/core/artwork/resolve.go +++ b/core/artwork/resolve.go @@ -198,7 +198,7 @@ func (r *resolver) fetchExternalArtist(ctx context.Context, ar model.Artist) (io // resolveAlbum walks conf.Server.CoverArtPriority over the folder, embedded and external sources. func (r *resolver) resolveAlbum(ctx context.Context, albumID string) (resolution, error) { - al, err := r.ds.Album(ctx).Get(albumID) + al, err := r.ds.Album().Get(ctx, albumID) if err != nil { return resolution{}, err } @@ -243,7 +243,7 @@ func (r *resolver) resolveAlbum(ctx context.Context, albumID string) (resolution // resolveArtist tries the uploaded image first, then walks conf.Server.ArtistArtPriority. func (r *resolver) resolveArtist(ctx context.Context, artistID string) (resolution, error) { - ar, err := r.ds.Artist(ctx).Get(artistID) + ar, err := r.ds.Artist().Get(ctx, artistID) if err != nil { return resolution{}, err } @@ -259,7 +259,7 @@ func (r *resolver) resolveArtist(ctx context.Context, artistID string) (resoluti } // Only consider albums where the artist is the sole album artist. - als, err := r.ds.Album(ctx).GetAll(model.QueryOptions{Filters: persistence.SoleAlbumArtistFilter(artistID)}) + als, err := r.ds.Album().GetAll(ctx, model.QueryOptions{Filters: persistence.SoleAlbumArtistFilter(artistID)}) if err != nil { return resolution{}, err } @@ -328,7 +328,7 @@ const PlaylistGridSamples = 4 // resolvePlaylist tries the uploaded image, the sidecar and ExternalImageURL, then a generated grid. func (r *resolver) resolvePlaylist(ctx context.Context, playlistID string) (resolution, error) { - pl, err := r.ds.Playlist(ctx).Get(playlistID) + pl, err := r.ds.Playlist().Get(ctx, playlistID) if err != nil { return resolution{}, err } @@ -374,8 +374,8 @@ func (r *resolver) resolvePlaylist(ctx context.Context, playlistID string) (reso } } - albumIDs, err := r.ds.Playlist(ctx).Tracks(pl.ID, false). - GetAlbumIDs(model.QueryOptions{Max: PlaylistGridSamples, Sort: "random()"}) + albumIDs, err := r.ds.Playlist().Tracks(ctx, pl.ID, false). + GetAlbumIDs(ctx, model.QueryOptions{Max: PlaylistGridSamples, Sort: "random()"}) if err != nil { return resolution{}, err } @@ -428,7 +428,7 @@ func (r *resolver) resolvePlaylist(ctx context.Context, playlistID string) (reso // resolveRadio serves only an uploaded image; there is no fallback. func (r *resolver) resolveRadio(ctx context.Context, radioID string) (resolution, error) { - radio, err := r.ds.Radio(ctx).Get(radioID) + radio, err := r.ds.Radio().Get(ctx, radioID) if err != nil { return resolution{}, err } @@ -439,7 +439,7 @@ func (r *resolver) resolveRadio(ctx context.Context, radioID string) (resolution // resolveMediaFile resolves a track's own embedded art only, so disabled or missing cover art // is a definitive absent. func (r *resolver) resolveMediaFile(ctx context.Context, id string) (resolution, error) { - mf, err := r.ds.MediaFile(ctx).Get(id) + mf, err := r.ds.MediaFile().Get(ctx, id) if err != nil { return resolution{}, err } diff --git a/core/artwork/uploader_test.go b/core/artwork/uploader_test.go index 44f5ede26..b1420c935 100644 --- a/core/artwork/uploader_test.go +++ b/core/artwork/uploader_test.go @@ -16,12 +16,14 @@ import ( ) var _ = Describe("Uploader", func() { + var ctx context.Context var svc Uploader var tmpDir string var artRepo *tests.MockArtworkRepo var queueRepo *tests.MockArtworkQueueRepo BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) tmpDir = GinkgoT().TempDir() conf.Server.DataFolder = conf.NewDir(tmpDir) @@ -33,7 +35,6 @@ var _ = Describe("Uploader", func() { Describe("SetImage", func() { It("creates directory and saves image file", func() { - ctx := context.Background() reader := strings.NewReader("fake image data") filename, err := svc.SetImage(ctx, consts.EntityArtist, "ar-1", "Pink Floyd", "", reader, ".jpg") Expect(err).ToNot(HaveOccurred()) @@ -46,7 +47,6 @@ var _ = Describe("Uploader", func() { }) It("falls back to ID-only filename when name cleans to empty", func() { - ctx := context.Background() reader := strings.NewReader("data") filename, err := svc.SetImage(ctx, consts.EntityPlaylist, "pl-1", "!!!", "", reader, ".png") Expect(err).ToNot(HaveOccurred()) @@ -54,7 +54,6 @@ var _ = Describe("Uploader", func() { }) It("removes old image when replacing", func() { - ctx := context.Background() oldDir := filepath.Join(tmpDir, "artwork", "artist") Expect(os.MkdirAll(oldDir, 0755)).To(Succeed()) oldFile := filepath.Join(oldDir, "ar-1_old.png") @@ -70,15 +69,13 @@ var _ = Describe("Uploader", func() { }) It("ignores missing old file without error", func() { - ctx := context.Background() reader := strings.NewReader("data") _, err := svc.SetImage(ctx, consts.EntityArtist, "ar-1", "Name", "/nonexistent/path.jpg", reader, ".jpg") Expect(err).ToNot(HaveOccurred()) }) It("does not touch artwork state or the queue (that is EnqueueArtwork's job, post-Put)", func() { - ctx := context.Background() - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "ar", ItemID: "ar-1", Hash: "oldhash", Source: "external", })).To(Succeed()) @@ -87,25 +84,24 @@ var _ = Describe("Uploader", func() { // SetImage only writes the file; the state row survives and nothing is queued until // the caller has persisted the new filename and called EnqueueArtwork. - _, err = artRepo.GetItemArtwork(model.KindArtistArtwork, "ar-1", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "ar-1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) - Expect(queueRepo.DequeueBatch(1000)).To(BeEmpty()) + Expect(queueRepo.DequeueBatch(ctx, 1000)).To(BeEmpty()) }) }) Describe("EnqueueArtwork", func() { It("clears artwork state and enqueues a Bump", func() { - ctx := context.Background() - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "ar", ItemID: "ar-1", Hash: "oldhash", Source: "external", })).To(Succeed()) svc.EnqueueArtwork(ctx, consts.EntityArtist, "ar-1") - _, err := artRepo.GetItemArtwork(model.KindArtistArtwork, "ar-1", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "ar-1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) - queued, err := queueRepo.DequeueBatch(1000) + queued, err := queueRepo.DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(queued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "ar"), @@ -116,13 +112,12 @@ var _ = Describe("Uploader", func() { It("is a no-op for an unknown entity type", func() { svc.EnqueueArtwork(context.Background(), "unknown", "x-1") - Expect(queueRepo.DequeueBatch(1000)).To(BeEmpty()) + Expect(queueRepo.DequeueBatch(ctx, 1000)).To(BeEmpty()) }) }) Describe("RemoveImage", func() { It("removes the file at the given path", func() { - ctx := context.Background() dir := filepath.Join(tmpDir, "artwork", "artist") Expect(os.MkdirAll(dir, 0755)).To(Succeed()) path := filepath.Join(dir, "ar-1_test.jpg") @@ -134,13 +129,11 @@ var _ = Describe("Uploader", func() { }) It("succeeds when file does not exist", func() { - ctx := context.Background() err := svc.RemoveImage(ctx, "/nonexistent/file.jpg") Expect(err).ToNot(HaveOccurred()) }) It("succeeds with empty path", func() { - ctx := context.Background() err := svc.RemoveImage(ctx, "") Expect(err).ToNot(HaveOccurred()) }) diff --git a/core/artwork/worker.go b/core/artwork/worker.go index e8fb3a11c..28e51958c 100644 --- a/core/artwork/worker.go +++ b/core/artwork/worker.go @@ -45,7 +45,7 @@ type Worker struct { broker events.Broker pruneMu sync.RWMutex pools []*drainPool - runCtx context.Context + runCtx context.Context //nolint:containedctx // worker lifecycle ctx, set at Run paused func() bool gatesMu sync.Mutex @@ -155,7 +155,7 @@ func (w *Worker) drain(ctx context.Context, concurrency int, kinds ...string) (i } // Dequeue well past the pool size so a slow external lookup never idles the other slots. // DequeueBatch does not mark rows taken, so this is one query per pass, not per slot. - items, err := w.proc.ds.ArtworkQueue(ctx).DequeueBatch(max(16, 4*concurrency), kinds...) + items, err := w.proc.ds.ArtworkQueue().DequeueBatch(ctx, max(16, 4*concurrency), kinds...) if err != nil { return 0, err } @@ -246,12 +246,12 @@ func (w *Worker) process(ctx context.Context, item model.ArtworkQueueItem) (outc ctx = withTrace(ctx, trace) out, got, retryIn := w.proc.acquire(ctx, item) - queue := w.proc.ds.ArtworkQueue(ctx) + queue := w.proc.ds.ArtworkQueue() switch out { case outcomeFound, outcomeAbsent: // A scan that re-enqueued this row mid-flight reset its retry_at, so the row survives // here and the next drain re-resolves it. - if err := queue.DeleteIfUnchanged(item.ItemKind, item.ItemID, item.ImageType, item.RetryAt); err != nil { + if err := queue.DeleteIfUnchanged(ctx, item.ItemKind, item.ItemID, item.ImageType, item.RetryAt); err != nil { log.Warn(ctx, "Artwork: Could not delete processed queue item", "kind", item.ItemKind, "id", item.ItemID, err) } case outcomeFoundStale, outcomeFailed: @@ -260,7 +260,7 @@ func (w *Worker) process(ctx context.Context, item model.ArtworkQueueItem) (outc if retryAt.Before(item.EnqueuedAt.Add(giveUpAfter)) { // A mid-flight re-enqueue reset retry_at; stale backoff must not stomp its // fresh, immediate eligibility. - if err := queue.MarkFailedIfUnchanged(item.ItemKind, item.ItemID, item.ImageType, item.RetryAt, retryAt, encoded); err != nil { + if err := queue.MarkFailedIfUnchanged(ctx, item.ItemKind, item.ItemID, item.ImageType, item.RetryAt, retryAt, encoded); err != nil { log.Warn(ctx, "Artwork: Could not reschedule failed queue item", "kind", item.ItemKind, "id", item.ItemID, err) } log.Debug(ctx, "Artwork: Rescheduled item", "kind", item.ItemKind, "id", item.ItemID, @@ -271,7 +271,7 @@ func (w *Worker) process(ctx context.Context, item model.ArtworkQueueItem) (outc // Art already being served is kept: exhaustion means unreachable, not removed. settled := "kept previous state" if out == outcomeFailed && settlesAbsentOnGiveUp(item.ItemKind) && !w.hasResolvedArtwork(ctx, item) { - writeAbsent(ctx, w.proc.ds.Artwork(ctx), item) + writeAbsent(ctx, w.proc.ds.Artwork(), item) settled = "recorded absent" } // The queue row is about to go, taking the only record of the failure with it. This write is @@ -279,7 +279,7 @@ func (w *Worker) process(ctx context.Context, item model.ArtworkQueueItem) (outc w.recordGiveUp(ctx, item, encoded) log.Info(ctx, "Artwork: Retry budget exhausted, giving up", "kind", item.ItemKind, "id", item.ItemID, "outcome", out, "attempts", item.Attempts+1, "budget", giveUpAfter, "settled", settled) - if err := queue.DeleteIfUnchanged(item.ItemKind, item.ItemID, item.ImageType, item.RetryAt); err != nil { + if err := queue.DeleteIfUnchanged(ctx, item.ItemKind, item.ItemID, item.ImageType, item.RetryAt); err != nil { log.Warn(ctx, "Artwork: Could not remove exhausted queue item", "kind", item.ItemKind, "id", item.ItemID, err) } } @@ -293,7 +293,7 @@ func (w *Worker) recordGiveUp(ctx context.Context, item model.ArtworkQueueItem, if !ok { return } - if err := w.proc.ds.Artwork(ctx).PutLastFailure(kind, item.ItemID, item.ImageType, trace); err != nil { + if err := w.proc.ds.Artwork().PutLastFailure(ctx, kind, item.ItemID, item.ImageType, trace); err != nil { log.Warn(ctx, "Artwork: Could not record the last failure", "kind", item.ItemKind, "id", item.ItemID, err) } } @@ -303,7 +303,7 @@ func (w *Worker) hasResolvedArtwork(ctx context.Context, item model.ArtworkQueue if !ok { return false } - ia, err := w.proc.ds.Artwork(ctx).GetItemArtwork(kind, item.ItemID, item.ImageType) + ia, err := w.proc.ds.Artwork().GetItemArtwork(ctx, kind, item.ItemID, item.ImageType) return err == nil && ia.Hash != "" } diff --git a/core/artwork/worker_soak_test.go b/core/artwork/worker_soak_test.go index 803cc2dfe..9be95cf2a 100644 --- a/core/artwork/worker_soak_test.go +++ b/core/artwork/worker_soak_test.go @@ -21,6 +21,12 @@ import ( const soakCycles = 2200 var _ = Describe("Worker soak", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) + It("does not leak goroutines, heap, or fds over many acquisition cycles", func() { if testing.Short() { Skip("skipping soak test in short mode") @@ -100,9 +106,9 @@ var _ = Describe("Worker soak", func() { // Read-back exercises the surfaces a caller would use after acquisition. if out == outcomeFound { kind, _ := model.ParseKind(it.ItemKind) - ia, err := artRepo.GetItemArtwork(kind, it.ItemID, model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, kind, it.ItemID, model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred(), "cycle %d: GetItemArtwork", i) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred(), "cycle %d: GetImage", i) rc, err := store.Open(ia.Hash, art.Mime) switch { diff --git a/core/artwork/worker_test.go b/core/artwork/worker_test.go index d74e3a5ed..a6c07b763 100644 --- a/core/artwork/worker_test.go +++ b/core/artwork/worker_test.go @@ -56,8 +56,8 @@ type reenqueueOnDequeue struct { done bool } -func (r *reenqueueOnDequeue) DequeueBatch(n int, kinds ...string) ([]model.ArtworkQueueItem, error) { - items, err := r.MockArtworkQueueRepo.DequeueBatch(n, kinds...) +func (r *reenqueueOnDequeue) DequeueBatch(ctx context.Context, n int, kinds ...string) ([]model.ArtworkQueueItem, error) { + items, err := r.MockArtworkQueueRepo.DequeueBatch(ctx, n, kinds...) if !r.done && len(items) > 0 { r.done = true for k, it := range r.Data { @@ -124,18 +124,27 @@ type visibilityPlaylistDS struct { tracks model.PlaylistTrackRepository } -func (v *visibilityPlaylistDS) Playlist(ctx context.Context) model.PlaylistRepository { +func (v *visibilityPlaylistDS) Playlist() model.PlaylistRepository { repo := tests.CreateMockPlaylistRepo() repo.TracksRepo = v.tracks - if u, ok := request.UserFrom(ctx); ok && u.IsAdmin { - repo.SetData(model.Playlists{v.private}) + repo.SetData(model.Playlists{v.private}) + return &visibilityPlaylistRepo{MockPlaylistRepo: repo} +} + +type visibilityPlaylistRepo struct { + *tests.MockPlaylistRepo +} + +func (v *visibilityPlaylistRepo) Get(ctx context.Context, id string) (*model.Playlist, error) { + if u, ok := request.UserFrom(ctx); !ok || !u.IsAdmin { + return nil, model.ErrNotFound } - return repo + return v.MockPlaylistRepo.Get(ctx, id) } func adminUserRepo() *tests.MockedUserRepo { repo := tests.CreateMockUserRepo() - Expect(repo.Put(&model.User{ID: "admin", UserName: "admin", IsAdmin: true})).To(Succeed()) + Expect(repo.Put(GinkgoT().Context(), &model.User{ID: "admin", UserName: "admin", IsAdmin: true})).To(Succeed()) return repo } @@ -157,8 +166,8 @@ var _ = Describe("Worker", func() { ) BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) - ctx = context.Background() var err error repoRoot, err = os.Getwd() Expect(err).ToNot(HaveOccurred()) @@ -200,7 +209,7 @@ var _ = Describe("Worker", func() { ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{ {ID: "al1", Name: "Album", FolderIDs: []string{"f1"}}, }) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: "al1", Priority: model.ArtworkPriorityScan, })).To(Succeed()) @@ -208,11 +217,11 @@ var _ = Describe("Worker", func() { Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(1)) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("folder")) - count, err := queueRepo.Count() + count, err := queueRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeZero(), "a found item must be deleted from the queue") }) @@ -223,7 +232,7 @@ var _ = Describe("Worker", func() { ds.MockedMediaFile.(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: "mf1", LibraryID: 0, Path: "tests/fixtures/artist/an-album/test.mp3", HasCoverArt: true}, }) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "mf", ItemID: "mf1", Priority: model.ArtworkPriorityBump, })).To(Succeed()) @@ -231,12 +240,12 @@ var _ = Describe("Worker", func() { Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(1)) - ia, err := artRepo.GetItemArtwork(model.KindMediaFileArtwork, "mf1", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindMediaFileArtwork, "mf1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("embedded")) Expect(ia.Hash).ToNot(BeEmpty()) - art, err := artRepo.GetImage(ia.Hash) + art, err := artRepo.GetImage(ctx, ia.Hash) Expect(err).ToNot(HaveOccurred()) r, err := store.Open(ia.Hash, art.Mime) Expect(err).ToNot(HaveOccurred()) @@ -245,7 +254,7 @@ var _ = Describe("Worker", func() { Expect(err).ToNot(HaveOccurred()) Expect(data).ToNot(BeEmpty(), "embedded bytes must be written to the store") - count, err := queueRepo.Count() + count, err := queueRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeZero()) }) @@ -254,7 +263,7 @@ var _ = Describe("Worker", func() { conf.Server.CoverArtPriority = "external" ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al4", Name: "Album"}}) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al4"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al4"})).To(Succeed()) n, err := w.drain(ctx, 2) Expect(err).ToNot(HaveOccurred()) @@ -265,7 +274,7 @@ var _ = Describe("Worker", func() { Expect(it.Attempts).To(Equal(1)) Expect(it.RetryAt).To(BeTemporally(">", time.Now())) - _, err = artRepo.GetItemArtwork(model.KindAlbumArtwork, "al4", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al4", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "a timeout must never settle on absent") }) @@ -275,7 +284,7 @@ var _ = Describe("Worker", func() { // Well above backoff(0)'s jittered ceiling, so only the hint can produce this retry_at. const askedFor = 90 * time.Minute imageAgents(&fakeImageAgent{name: "throttledAgent", err: &agents.RetryLaterError{RetryIn: askedFor}}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al9"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al9"})).To(Succeed()) n, err := w.drain(ctx, 2) Expect(err).ToNot(HaveOccurred()) @@ -296,7 +305,7 @@ var _ = Describe("Worker", func() { {ID: "alstale", Name: "Album", FolderIDs: []string{"f1"}}, }) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "alstale"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alstale"})).To(Succeed()) n, err := w.drain(ctx, 2) Expect(err).ToNot(HaveOccurred()) @@ -307,7 +316,7 @@ var _ = Describe("Worker", func() { Expect(it.Attempts).To(Equal(1)) Expect(it.RetryAt).To(BeTemporally(">", time.Now())) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alstale", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alstale", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("folder"), "the fallback art is served meanwhile") @@ -327,7 +336,7 @@ var _ = Describe("Worker", func() { racing := &reenqueueOnDequeue{MockArtworkQueueRepo: queueRepo} ds.MockedArtworkQueue = racing w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: "al7", Priority: model.ArtworkPriorityScan, })).To(Succeed()) @@ -337,7 +346,7 @@ var _ = Describe("Worker", func() { // The concurrent re-enqueue changed retry_at, so the found-path delete was a no-op. Expect(findQueued(queueRepo, "al", "al7")).ToNot(BeNil()) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al7", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al7", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("folder")) }) @@ -349,7 +358,7 @@ var _ = Describe("Worker", func() { racing := &reenqueueOnDequeue{MockArtworkQueueRepo: queueRepo} ds.MockedArtworkQueue = racing w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al8"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al8"})).To(Succeed()) dequeued := findQueued(queueRepo, "al", "al8").RetryAt n, err := w.drain(ctx, 1) @@ -368,7 +377,7 @@ var _ = Describe("Worker", func() { ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al9", Name: "Album"}}) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al9"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al9"})).To(Succeed()) // Age the row past the retry budget. expireQueued(queueRepo, "al9") @@ -377,7 +386,7 @@ var _ = Describe("Worker", func() { Expect(n).To(Equal(1)) Expect(findQueued(queueRepo, "al", "al9")).To(BeNil()) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al9", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al9", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(BeEmpty()) }) @@ -385,13 +394,13 @@ var _ = Describe("Worker", func() { It("keeps already-served art when the retry budget is exhausted", func() { conf.Server.CoverArtPriority = "external" ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al10", Name: "Album"}}) - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al10", ImageType: model.ImageTypePrimary, Hash: "cafebabe", Source: "external:lastfm", })).To(Succeed()) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al10"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al10"})).To(Succeed()) expireQueued(queueRepo, "al10") n, err := w.drain(ctx, 1) @@ -399,7 +408,7 @@ var _ = Describe("Worker", func() { Expect(n).To(Equal(1)) Expect(findQueued(queueRepo, "al", "al10")).To(BeNil()) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al10", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al10", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(Equal("cafebabe"), "a persistent outage must not discard served art") }) @@ -409,7 +418,7 @@ var _ = Describe("Worker", func() { ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al11", Name: "Album"}}) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al11"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al11"})).To(Succeed()) _, err := w.drain(ctx, 1) Expect(err).ToNot(HaveOccurred()) @@ -430,13 +439,13 @@ var _ = Describe("Worker", func() { ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al13", Name: "Album"}}) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al13"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al13"})).To(Succeed()) expireQueued(queueRepo, "al13") _, err := w.drain(ctx, 1) Expect(err).ToNot(HaveOccurred()) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al13", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al13", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred(), "settling absent must create the row the failure is written to") Expect(ia.Hash).To(BeEmpty()) Expect(DecodeTrace(ia.LastFailure, "")).ToNot(BeEmpty()) @@ -445,20 +454,20 @@ var _ = Describe("Worker", func() { It("keeps the failure on the state row after the queue row is deleted", func() { conf.Server.CoverArtPriority = "external" ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al12", Name: "Album"}}) - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al12", ImageType: model.ImageTypePrimary, Hash: "cafebabe", Source: "external:lastfm", })).To(Succeed()) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) w = NewWorker(ds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al12"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al12"})).To(Succeed()) expireQueued(queueRepo, "al12") _, err := w.drain(ctx, 1) Expect(err).ToNot(HaveOccurred()) Expect(findQueued(queueRepo, "al", "al12")).To(BeNil()) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al12", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al12", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(DecodeTrace(ia.LastFailure, "")).ToNot(BeEmpty(), "the queue row is gone, so this is the only remaining record of the failure") @@ -473,7 +482,7 @@ var _ = Describe("Worker", func() { ds.MockedMediaFile.(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: "mfX", LibraryID: 0, Path: "tests/fixtures/artist/an-album/gone.mp3", HasCoverArt: true}, }) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "mf", ItemID: "mfX"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "mf", ItemID: "mfX"})).To(Succeed()) expireQueued(queueRepo, "mfX") n, err := w.drain(ctx, 1) @@ -481,7 +490,7 @@ var _ = Describe("Worker", func() { Expect(n).To(Equal(1)) Expect(findQueued(queueRepo, "mf", "mfX")).To(BeNil(), "the row must stop retrying") - _, err = artRepo.GetItemArtwork(model.KindMediaFileArtwork, "mfX", model.ImageTypePrimary) + _, err = artRepo.GetItemArtwork(ctx, model.KindMediaFileArtwork, "mfX", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "no row leaves the track unresolved, so a later view can still recover it") // Known gap: with no row and no absent settle, there is nowhere to keep the failure. @@ -496,14 +505,14 @@ var _ = Describe("Worker", func() { tracks: &tests.MockPlaylistTrackRepo{}, } w = NewWorker(vds, store, ag, ffm, broker, imgCache) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "pl", ItemID: "plPriv"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "pl", ItemID: "plPriv"})).To(Succeed()) n, err := w.drain(ctx, 1) Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(1)) Expect(findQueued(queueRepo, "pl", "plPriv")).To(BeNil()) - ia, err := artRepo.GetItemArtwork(model.KindPlaylistArtwork, "plPriv", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindPlaylistArtwork, "plPriv", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(BeEmpty()) }) @@ -523,9 +532,9 @@ var _ = Describe("Worker", func() { {ID: "al1", Name: "Album 1", FolderIDs: []string{"f1"}}, {ID: "al2", Name: "Album 2", FolderIDs: []string{"f1"}}, }) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al1", Priority: model.ArtworkPriorityScan})).To(Succeed()) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al2", Priority: model.ArtworkPriorityScan})).To(Succeed()) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar1", Priority: model.ArtworkPriorityScan})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al1", Priority: model.ArtworkPriorityScan})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al2", Priority: model.ArtworkPriorityScan})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "ar", ItemID: "ar1", Priority: model.ArtworkPriorityScan})).To(Succeed()) n, err := w.drain(ctx, 3) Expect(err).ToNot(HaveOccurred()) @@ -572,7 +581,7 @@ var _ = Describe("Worker", func() { conf.Server.CoverArtPriority = "cover.*" // local-only; no folder image → absent ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al3", Name: "Artless"}}) folderRepo.result = nil - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "al3", Priority: model.ArtworkPriorityScan})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "al3", Priority: model.ArtworkPriorityScan})).To(Succeed()) n, err := w.drain(ctx, 2) Expect(err).ToNot(HaveOccurred()) @@ -582,7 +591,7 @@ var _ = Describe("Worker", func() { Expect(evts).To(HaveLen(1), "a removed cover must live-refresh clients so they drop it") Expect(evts[0].(*events.RefreshResource).Data(evts[0])).To(ContainSubstring("al3")) - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al3", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al3", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(BeEmpty(), "the outcome was absent, not found") }) @@ -591,7 +600,7 @@ var _ = Describe("Worker", func() { conf.Server.CoverArtPriority = "external" ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "alx", Name: "Album"}}) imageAgents(&fakeImageAgent{name: "failAgent", err: errors.New("agent timed out")}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ItemKind: "al", ItemID: "alx"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alx"})).To(Succeed()) n, err := w.drain(ctx, 2) Expect(err).ToNot(HaveOccurred()) @@ -725,7 +734,7 @@ var _ = Describe("Worker", func() { {ID: "alpc", Name: "Album", FolderIDs: []string{"f1"}}, }) conf.Server.UICoverArtSize = 300 - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: "alpc", Priority: model.ArtworkPriorityScan, })).To(Succeed()) }) @@ -813,11 +822,11 @@ var _ = Describe("Worker", func() { // Artists first, exactly as Backfill orders them. for _, a := range artists { - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "ar", ItemID: a.ID, Priority: model.ArtworkPriorityBackfill, })).To(Succeed()) } - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: "alx", Priority: model.ArtworkPriorityBackfill, })).To(Succeed()) @@ -832,12 +841,12 @@ var _ = Describe("Worker", func() { }) Eventually(func() bool { - ia, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "alx", model.ImageTypePrimary) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alx", model.ImageTypePrimary) return err == nil && ia.Hash != "" }, 5*time.Second, 50*time.Millisecond).Should(BeTrue(), "a blocked external pool must not hold up local artwork") - _, err := artRepo.GetItemArtwork(model.KindArtistArtwork, "arx0", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "arx0", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound), "artists are still blocked, as intended") }) }) @@ -849,7 +858,7 @@ var _ = Describe("Worker", func() { for i := range 8 { id := fmt.Sprintf("alc%d", i) albums = append(albums, model.Album{ID: id, Name: "Album"}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: id, Priority: model.ArtworkPriorityScan, })).To(Succeed()) } @@ -875,21 +884,21 @@ var _ = Describe("Worker", func() { for i := range 8 { id := fmt.Sprintf("alp%d", i) albums = append(albums, model.Album{ID: id, Name: "Album", FolderIDs: []string{"f1"}}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: id, Priority: model.ArtworkPriorityScan, })).To(Succeed()) } ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(albums) // Pauses as soon as the first item has left the queue. w.PauseWhile(func() bool { - n, _ := queueRepo.Count() + n, _ := queueRepo.Count(ctx) return n < 8 }) _, err := w.drain(ctx, 1) Expect(err).ToNot(HaveOccurred()) - count, err := queueRepo.Count() + count, err := queueRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(7)), "only the item dispatched before the pause may leave the queue") }) @@ -897,7 +906,7 @@ var _ = Describe("Worker", func() { It("dequeues past the worker pool so one drain covers many items", func() { for i := range 16 { ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: fmt.Sprintf("alb%d", i), Name: "Album"}}) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: fmt.Sprintf("alb%d", i), Priority: model.ArtworkPriorityScan, })).To(Succeed()) } @@ -932,7 +941,7 @@ var _ = Describe("Worker", func() { ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{ {ID: "al1", Name: "Album", FolderIDs: []string{"f1"}}, }) - Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ ItemKind: "al", ItemID: "al1", Priority: model.ArtworkPriorityScan, })).To(Succeed()) w.PauseWhile(func() bool { return true }) diff --git a/core/auth/auth.go b/core/auth/auth.go index b36bb2696..1bdc917da 100644 --- a/core/auth/auth.go +++ b/core/auth/auth.go @@ -48,7 +48,7 @@ func Init(ds model.DataStore) { } func loadOrCreateSecret(ctx context.Context, ds model.DataStore, key string) string { - secret, err := ds.Property(ctx).Get(key) + secret, err := ds.Property().Get(ctx, key) if err != nil || secret == "" { log.Info(ctx, "Creating new JWT secret", "key", key) return createNewSecret(ctx, ds, key) @@ -154,9 +154,9 @@ func CheckClaims(c Claims, usr model.User, audience string) error { } func WithAdminUser(ctx context.Context, ds model.DataStore) context.Context { - u, err := ds.User(ctx).FindFirstAdmin() + u, err := ds.User().FindFirstAdmin(ctx) if err != nil { - c, err := ds.User(ctx).CountAll() + c, err := ds.User().CountAll(ctx) if c == 0 && err == nil { log.Debug(ctx, "No admin user yet!", err) } else { @@ -176,7 +176,7 @@ func createNewSecret(ctx context.Context, ds model.DataStore, key string) string log.Error(ctx, "Could not encrypt JWT secret", err) return secret } - if err := ds.Property(ctx).Put(key, encSecret); err != nil { + if err := ds.Property().Put(ctx, key, encSecret); err != nil { log.Error(ctx, "Could not save JWT secret in DB", err) } return secret diff --git a/core/common.go b/core/common.go index 6ff349b1b..db2dbbf1b 100644 --- a/core/common.go +++ b/core/common.go @@ -19,7 +19,7 @@ func userName(ctx context.Context) string { // BFR We should only access files through the `storage.Storage` interface. This will require changing how // TagLib and ffmpeg access files var AbsolutePath = func(ctx context.Context, ds model.DataStore, libId int, path string) string { - libPath, err := ds.Library(ctx).GetPath(libId) + libPath, err := ds.Library().GetPath(ctx, libId) if err != nil { return path } diff --git a/core/external/extdata_helper_test.go b/core/external/extdata_helper_test.go index f7a155cd9..73d88e5b4 100644 --- a/core/external/extdata_helper_test.go +++ b/core/external/extdata_helper_test.go @@ -31,7 +31,7 @@ func (m *mockArtistRepo) SetData(artists model.Artists) { } // Get implements model.ArtistRepository. -func (m *mockArtistRepo) Get(id string) (*model.Artist, error) { +func (m *mockArtistRepo) Get(_ context.Context, id string) (*model.Artist, error) { args := m.Called(id) if args.Get(0) == nil { return nil, args.Error(1) @@ -40,7 +40,7 @@ func (m *mockArtistRepo) Get(id string) (*model.Artist, error) { } // GetAll implements model.ArtistRepository. -func (m *mockArtistRepo) GetAll(options ...model.QueryOptions) (model.Artists, error) { +func (m *mockArtistRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Artists, error) { argsSlice := make([]any, len(options)) for i, v := range options { argsSlice[i] = v @@ -85,7 +85,7 @@ func (m *mockMediaFileRepo) SetData(mediaFiles model.MediaFiles) { } // Get implements model.MediaFileRepository. -func (m *mockMediaFileRepo) Get(id string) (*model.MediaFile, error) { +func (m *mockMediaFileRepo) Get(ctx context.Context, id string) (*model.MediaFile, error) { args := m.Called(id) if args.Get(0) == nil { return nil, args.Error(1) @@ -94,12 +94,12 @@ func (m *mockMediaFileRepo) Get(id string) (*model.MediaFile, error) { } // GetAllByTags implements model.MediaFileRepository. -func (m *mockMediaFileRepo) GetAllByTags(_ model.TagName, _ []string, options ...model.QueryOptions) (model.MediaFiles, error) { - return m.GetAll(options...) +func (m *mockMediaFileRepo) GetAllByTags(ctx context.Context, _ model.TagName, _ []string, options ...model.QueryOptions) (model.MediaFiles, error) { + return m.GetAll(ctx, options...) } // GetAll implements model.MediaFileRepository. -func (m *mockMediaFileRepo) GetAll(options ...model.QueryOptions) (model.MediaFiles, error) { +func (m *mockMediaFileRepo) GetAll(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { argsSlice := make([]any, len(options)) for i, v := range options { argsSlice[i] = v @@ -112,7 +112,7 @@ func (m *mockMediaFileRepo) GetAll(options ...model.QueryOptions) (model.MediaFi } // GetRandom implements model.MediaFileRepository. -func (m *mockMediaFileRepo) GetRandom(options ...model.QueryOptions) (model.MediaFiles, error) { +func (m *mockMediaFileRepo) GetRandom(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { argsSlice := make([]any, len(options)) for i, v := range options { argsSlice[i] = v @@ -156,7 +156,7 @@ func newMockAlbumRepo() *mockAlbumRepo { } // Get implements model.AlbumRepository. -func (m *mockAlbumRepo) Get(id string) (*model.Album, error) { +func (m *mockAlbumRepo) Get(_ context.Context, id string) (*model.Album, error) { args := m.Called(id) if args.Get(0) == nil { return nil, args.Error(1) @@ -165,7 +165,7 @@ func (m *mockAlbumRepo) Get(id string) (*model.Album, error) { } // GetAll implements model.AlbumRepository. -func (m *mockAlbumRepo) GetAll(options ...model.QueryOptions) (model.Albums, error) { +func (m *mockAlbumRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Albums, error) { argsSlice := make([]any, len(options)) for i, v := range options { argsSlice[i] = v diff --git a/core/external/provider.go b/core/external/provider.go index 3a3f4bd46..185725259 100644 --- a/core/external/provider.go +++ b/core/external/provider.go @@ -182,7 +182,7 @@ func (e *provider) populateAlbumInfo(ctx context.Context, album auxAlbum) (auxAl } } - err = e.ds.Album(ctx).UpdateExternalInfo(&album.Album) + err = e.ds.Album().UpdateExternalInfo(ctx, &album.Album) if err != nil { log.Error(ctx, "Error trying to update album external information", "id", album.ID, "name", albumName, "elapsed", time.Since(start), err) @@ -285,7 +285,7 @@ func (e *provider) populateArtistInfo(ctx context.Context, artist auxArtist) (au if !throttled { artist.ExternalInfoUpdatedAt = new(time.Now()) } - err := e.ds.Artist(ctx).UpdateExternalInfo(&artist.Artist) + err := e.ds.Artist().UpdateExternalInfo(ctx, &artist.Artist) if err != nil { log.Error(ctx, "Error trying to update artist external information", "id", artist.ID, "name", artistName, "elapsed", time.Since(start), err) @@ -548,7 +548,7 @@ func (e *provider) loadArtistsByID(ctx context.Context, similar []agents.Artist) if len(ids) == 0 { return matches, nil } - res, err := e.ds.Artist(ctx).GetAll(model.QueryOptions{ + res, err := e.ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"artist.id": ids}, }) if err != nil { @@ -577,7 +577,7 @@ func (e *provider) loadArtistsByMBID(ctx context.Context, similar []agents.Artis if len(mbids) == 0 { return matches, nil } - res, err := e.ds.Artist(ctx).GetAll(model.QueryOptions{ + res, err := e.ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"mbz_artist_id": mbids}, }) if err != nil { @@ -612,7 +612,7 @@ func (e *provider) loadArtistsByName(ctx context.Context, similar []agents.Artis clauses := slice.Map(names, func(name string) squirrel.Sqlizer { return squirrel.Like{"artist.name": name} }) - res, err := e.ds.Artist(ctx).GetAll(model.QueryOptions{ + res, err := e.ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Or(clauses), }) if err != nil { @@ -628,7 +628,7 @@ func (e *provider) loadArtistsByName(ctx context.Context, similar []agents.Artis func (e *provider) findArtist(ctx context.Context, artistName, id string) (*auxArtist, error) { if id != "" { - artist, err := e.ds.Artist(ctx).Get(id) + artist, err := e.ds.Artist().Get(ctx, id) if err == nil { return &auxArtist{Artist: *artist}, nil } @@ -644,7 +644,7 @@ func (e *provider) findArtist(ctx context.Context, artistName, id string) (*auxA return nil, model.ErrNotFound } - artists, err := e.ds.Artist(ctx).GetAll(model.QueryOptions{ + artists, err := e.ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Like{"artist.name": artistName}, Max: 1, }) @@ -666,7 +666,7 @@ func (e *provider) loadSimilar(ctx context.Context, artist *auxArtist, count int ids = append(ids, sa.ID) } - similar, err := e.ds.Artist(ctx).GetAll(model.QueryOptions{ + similar, err := e.ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"artist.id": ids}, }) if err != nil { diff --git a/core/external/provider_refreshinfo_test.go b/core/external/provider_refreshinfo_test.go index e7910a734..a1216dd94 100644 --- a/core/external/provider_refreshinfo_test.go +++ b/core/external/provider_refreshinfo_test.go @@ -71,8 +71,8 @@ var _ = Describe("Provider - RefreshInfo", func() { ag = new(mockAgents) broker = &fakeBroker{} p = external.NewProvider(ds, ag, matcher.New(ds), broker) - mockArtistRepo = ds.Artist(ctx).(*tests.MockArtistRepo) - mockAlbumRepo = ds.Album(ctx).(*tests.MockAlbumRepo) + mockArtistRepo = ds.Artist().(*tests.MockArtistRepo) + mockAlbumRepo = ds.Album().(*tests.MockAlbumRepo) }) It("repopulates an artist even when its info is fresh", func() { @@ -84,7 +84,7 @@ var _ = Describe("Provider - RefreshInfo", func() { Expect(p.RefreshInfo(ctx, model.KindArtistArtwork, "ar-1")).To(Succeed()) - saved, err := mockArtistRepo.Get("ar-1") + saved, err := mockArtistRepo.Get(ctx, "ar-1") Expect(err).ToNot(HaveOccurred()) Expect(saved.Biography).To(Equal("Fresh Bio")) }) @@ -99,7 +99,7 @@ var _ = Describe("Provider - RefreshInfo", func() { Expect(p.RefreshInfo(ctx, model.KindAlbumArtwork, "al-1")).To(Succeed()) - saved, err := mockAlbumRepo.Get("al-1") + saved, err := mockAlbumRepo.Get(ctx, "al-1") Expect(err).ToNot(HaveOccurred()) Expect(saved.Description).To(Equal("Fresh Notes")) }) diff --git a/core/external/provider_similarsongs.go b/core/external/provider_similarsongs.go index 4ab465b03..305b8d6ec 100644 --- a/core/external/provider_similarsongs.go +++ b/core/external/provider_similarsongs.go @@ -37,7 +37,7 @@ func (e *provider) SimilarSongs(ctx context.Context, id string, count int) (mode if !errors.Is(err, model.ErrNotFound) { return nil, err } - genre, err := e.ds.Genre(ctx).Get(id) + genre, err := e.ds.Genre().Get(ctx, id) if err != nil { return nil, err } @@ -178,13 +178,13 @@ func (e *provider) seedMix(ctx context.Context, count int, sample func() (model. func (e *provider) samplePlaylistTracks(ctx context.Context, playlistID string, n int) (model.MediaFiles, error) { // Refresh: a smart playlist materializes no tracks until it is evaluated, so skipping it would // mix an empty seed set. It is a no-op for regular playlists and inside the refresh delay. - repo := e.ds.Playlist(ctx).Tracks(playlistID, true) + repo := e.ds.Playlist().Tracks(ctx, playlistID, true) if repo == nil { return nil, model.ErrNotFound } // A playlist can hold the same file at several positions, so over-fetch and dedup: a repeated // seed wastes an agent call and can reach the mix twice through the seed fallback. - tracks, err := repo.GetAll(model.QueryOptions{ + tracks, err := repo.GetAll(ctx, model.QueryOptions{ Sort: "random", Max: n * 4, Filters: squirrel.Eq{"missing": false}, @@ -225,7 +225,7 @@ func (e *provider) sampleGenreTracks(ctx context.Context, genre *model.Genre, n // sampleTracks returns up to n random present tracks. Seeds can end up in the mix verbatim, so // missing files would surface as unplayable entries. func (e *provider) sampleTracks(ctx context.Context, filter squirrel.Sqlizer, n int) (model.MediaFiles, error) { - return e.ds.MediaFile(ctx).GetRandom(model.QueryOptions{ + return e.ds.MediaFile().GetRandom(ctx, model.QueryOptions{ Filters: squirrel.And{filter, squirrel.Eq{"missing": false}}, Max: n, }) diff --git a/core/external/provider_updatealbuminfo_test.go b/core/external/provider_updatealbuminfo_test.go index e168aa026..d2fd4364e 100644 --- a/core/external/provider_updatealbuminfo_test.go +++ b/core/external/provider_updatealbuminfo_test.go @@ -35,7 +35,7 @@ var _ = Describe("Provider - UpdateAlbumInfo", func() { ds = new(tests.MockDataStore) ag = new(mockAgents) p = external.NewProvider(ds, ag, matcher.New(ds), &fakeBroker{}) - mockAlbumRepo = ds.Album(ctx).(*tests.MockAlbumRepo) + mockAlbumRepo = ds.Album().(*tests.MockAlbumRepo) conf.Server.DevAlbumInfoTimeToLive = 1 * time.Hour }) diff --git a/core/external/provider_updateartistinfo_test.go b/core/external/provider_updateartistinfo_test.go index c722aaee8..5e2087d35 100644 --- a/core/external/provider_updateartistinfo_test.go +++ b/core/external/provider_updateartistinfo_test.go @@ -38,7 +38,7 @@ var _ = Describe("Provider - UpdateArtistInfo", func() { ds = new(tests.MockDataStore) ag = new(mockAgents) p = external.NewProvider(ds, ag, matcher.New(ds), &fakeBroker{}) - mockArtistRepo = ds.Artist(ctx).(*tests.MockArtistRepo) + mockArtistRepo = ds.Artist().(*tests.MockArtistRepo) }) It("returns error when artist is not found", func() { diff --git a/core/library.go b/core/library.go index 6df0e95b5..628ee4b7b 100644 --- a/core/library.go +++ b/core/library.go @@ -32,25 +32,28 @@ type Library interface { SetUserLibraries(ctx context.Context, userID string, libraryIDs []int) error ValidateLibraryAccess(ctx context.Context, userID string, libraryID int) error - NewRepository(ctx context.Context) rest.Repository + Repository() rest.Repository[model.Library] } type libraryService struct { - ds model.DataStore - scanner model.Scanner - watcher Watcher - broker events.Broker - pluginManager PluginUnloader + ds model.DataStore + broker events.Broker + repo *libraryRepositoryWrapper } // NewLibrary creates a new Library service func NewLibrary(ds model.DataStore, scanner model.Scanner, watcher Watcher, broker events.Broker, pluginManager PluginUnloader) Library { return &libraryService{ - ds: ds, - scanner: scanner, - watcher: watcher, - broker: broker, - pluginManager: pluginManager, + ds: ds, + broker: broker, + repo: &libraryRepositoryWrapper{ + LibraryRepository: ds.Library(), + ds: ds, + scanner: scanner, + watcher: watcher, + broker: broker, + pluginManager: pluginManager, + }, } } @@ -58,16 +61,16 @@ func NewLibrary(ds model.DataStore, scanner model.Scanner, watcher Watcher, brok func (s *libraryService) GetUserLibraries(ctx context.Context, userID string) (model.Libraries, error) { // Verify user exists - if _, err := s.ds.User(ctx).Get(userID); err != nil { + if _, err := s.ds.User().Get(ctx, userID); err != nil { return nil, err } - return s.ds.User(ctx).GetUserLibraries(userID) + return s.ds.User().GetUserLibraries(ctx, userID) } func (s *libraryService) SetUserLibraries(ctx context.Context, userID string, libraryIDs []int) error { // Verify user exists - user, err := s.ds.User(ctx).Get(userID) + user, err := s.ds.User().Get(ctx, userID) if err != nil { return err } @@ -90,7 +93,7 @@ func (s *libraryService) SetUserLibraries(ctx context.Context, userID string, li } // Set user libraries - err = s.ds.User(ctx).SetUserLibraries(userID, libraryIDs) + err = s.ds.User().SetUserLibraries(ctx, userID, libraryIDs) if err != nil { return fmt.Errorf("error setting user libraries: %w", err) } @@ -115,7 +118,7 @@ func (s *libraryService) ValidateLibraryAccess(ctx context.Context, userID strin } // Check if user has explicit access to this library - libraries, err := s.ds.User(ctx).GetUserLibraries(userID) + libraries, err := s.ds.User().GetUserLibraries(ctx, userID) if err != nil { log.Error(ctx, "Error checking library access", "userID", userID, "libraryID", libraryID, err) return fmt.Errorf("error checking library access: %w", err) @@ -132,25 +135,14 @@ func (s *libraryService) ValidateLibraryAccess(ctx context.Context, userID strin // REST repository wrapper -func (s *libraryService) NewRepository(ctx context.Context) rest.Repository { - repo := s.ds.Library(ctx) - wrapper := &libraryRepositoryWrapper{ - ctx: ctx, - LibraryRepository: repo, - Repository: repo.(rest.Repository), - ds: s.ds, - scanner: s.scanner, - watcher: s.watcher, - broker: s.broker, - pluginManager: s.pluginManager, - } - return wrapper +func (s *libraryService) Repository() rest.Repository[model.Library] { + return s.repo } +var _ rest.Persistable[model.Library] = (*libraryRepositoryWrapper)(nil) + type libraryRepositoryWrapper struct { - rest.Repository model.LibraryRepository - ctx context.Context ds model.DataStore scanner model.Scanner watcher Watcher @@ -158,59 +150,58 @@ type libraryRepositoryWrapper struct { pluginManager PluginUnloader } -func (r *libraryRepositoryWrapper) Save(entity any) (string, error) { - lib := entity.(*model.Library) - if err := r.validateLibrary(lib); err != nil { +func (r *libraryRepositoryWrapper) Save(ctx context.Context, lib *model.Library) (string, error) { + if err := r.validateLibrary(ctx, lib); err != nil { return "", err } - err := r.LibraryRepository.Put(lib) + err := r.LibraryRepository.Put(ctx, lib) if err != nil { return "", r.mapError(err) } // Start watcher and trigger scan after successful library creation if r.watcher != nil { - if err := r.watcher.Watch(r.ctx, lib); err != nil { - log.Warn(r.ctx, "Failed to start watcher for new library", "libraryID", lib.ID, "name", lib.Name, "path", lib.Path, err) + if err := r.watcher.Watch(ctx, lib); err != nil { + log.Warn(ctx, "Failed to start watcher for new library", "libraryID", lib.ID, "name", lib.Name, "path", lib.Path, err) } } if r.scanner != nil { - go r.triggerScan(lib, "new") + go r.triggerScan(ctx, lib, "new") } // Send library refresh event to all clients if r.broker != nil { event := &events.RefreshResource{} - r.broker.SendBroadcastMessage(r.ctx, event.With("library", strconv.Itoa(lib.ID))) - log.Debug(r.ctx, "Library created - sent refresh event", "libraryID", lib.ID, "name", lib.Name) + r.broker.SendBroadcastMessage(ctx, event.With("library", strconv.Itoa(lib.ID))) + log.Debug(ctx, "Library created - sent refresh event", "libraryID", lib.ID, "name", lib.Name) } return strconv.Itoa(lib.ID), nil } -func (r *libraryRepositoryWrapper) Update(id string, entity any, cols ...string) error { - lib := entity.(*model.Library) +func (r *libraryRepositoryWrapper) Update(ctx context.Context, id string, entity model.Library, cols ...string) error { + lib := &entity libID, err := strconv.Atoi(id) if err != nil { return fmt.Errorf("invalid library ID: %s", id) } lib.ID = libID - if err := r.validateLibrary(lib); err != nil { + if err := r.validateLibrary(ctx, lib); err != nil { return err } // Get the original library to check if path changed - originalLib, err := r.Get(libID) + originalLib, err := r.Get(ctx, libID) if err != nil { return r.mapError(err) } pathChanged := originalLib.Path != lib.Path - err = r.LibraryRepository.Put(lib, cols...) + err = r.LibraryRepository.Put(ctx, lib, cols...) if err != nil { return r.mapError(err) } @@ -218,27 +209,36 @@ func (r *libraryRepositoryWrapper) Update(id string, entity any, cols ...string) // Restart watcher and trigger scan if path was updated if pathChanged { if r.watcher != nil { - if err := r.watcher.Watch(r.ctx, lib); err != nil { - log.Warn(r.ctx, "Failed to restart watcher for updated library", "libraryID", lib.ID, "name", lib.Name, "path", lib.Path, err) + if err := r.watcher.Watch(ctx, lib); err != nil { + log.Warn(ctx, "Failed to restart watcher for updated library", "libraryID", lib.ID, "name", lib.Name, "path", lib.Path, err) } } if r.scanner != nil { - go r.triggerScan(lib, "updated") + go r.triggerScan(ctx, lib, "updated") } } // Send library refresh event to all clients if r.broker != nil { event := &events.RefreshResource{} - r.broker.SendBroadcastMessage(r.ctx, event.With("library", id)) - log.Debug(r.ctx, "Library updated - sent refresh event", "libraryID", libID, "name", lib.Name) + r.broker.SendBroadcastMessage(ctx, event.With("library", id)) + log.Debug(ctx, "Library updated - sent refresh event", "libraryID", libID, "name", lib.Name) } return nil } -func (r *libraryRepositoryWrapper) Delete(id string) error { +func (r *libraryRepositoryWrapper) Delete(ctx context.Context, ids ...string) error { + for _, id := range ids { + if err := r.deleteOne(ctx, id); err != nil { + return err + } + } + return nil +} + +func (r *libraryRepositoryWrapper) deleteOne(ctx context.Context, id string) error { libID, err := strconv.Atoi(id) if err != nil { return &rest.ValidationError{Errors: map[string]string{ @@ -247,7 +247,7 @@ func (r *libraryRepositoryWrapper) Delete(id string) error { } // Get library info before deletion for logging - lib, err := r.Get(libID) + lib, err := r.Get(ctx, libID) if err != nil { return r.mapError(err) } @@ -255,7 +255,7 @@ func (r *libraryRepositoryWrapper) Delete(id string) error { // Run the deletion in a transaction so the cascade delete and the orphaned-artist // reconciliation it triggers (see libraryRepository.Delete) commit atomically. err = r.ds.WithTx(func(tx model.DataStore) error { - return tx.Library(r.ctx).Delete(libID) + return tx.Library().Delete(ctx, libID) }, "delete library") if err != nil { return r.mapError(err) @@ -263,25 +263,25 @@ func (r *libraryRepositoryWrapper) Delete(id string) error { // Stop watcher and trigger scan after successful library deletion to clean up orphaned data if r.watcher != nil { - if err := r.watcher.StopWatching(r.ctx, libID); err != nil { - log.Warn(r.ctx, "Failed to stop watcher for deleted library", "libraryID", libID, "name", lib.Name, "path", lib.Path, err) + if err := r.watcher.StopWatching(ctx, libID); err != nil { + log.Warn(ctx, "Failed to stop watcher for deleted library", "libraryID", libID, "name", lib.Name, "path", lib.Path, err) } } if r.scanner != nil { - go r.triggerScan(lib, "deleted") + go r.triggerScan(ctx, lib, "deleted") } // Send library refresh event to all clients if r.broker != nil { event := &events.RefreshResource{} - r.broker.SendBroadcastMessage(r.ctx, event.With("library", id)) - log.Debug(r.ctx, "Library deleted - sent refresh event", "libraryID", libID, "name", lib.Name) + r.broker.SendBroadcastMessage(ctx, event.With("library", id)) + log.Debug(ctx, "Library deleted - sent refresh event", "libraryID", libID, "name", lib.Name) } // After successful deletion, check if any plugins were auto-disabled // and need to be unloaded from memory - r.pluginManager.UnloadDisabledPlugins(r.ctx) + r.pluginManager.UnloadDisabledPlugins(ctx) return nil } @@ -309,7 +309,7 @@ func (r *libraryRepositoryWrapper) mapError(err error) error { return err } -func (r *libraryRepositoryWrapper) validateLibrary(library *model.Library) error { +func (r *libraryRepositoryWrapper) validateLibrary(ctx context.Context, library *model.Library) error { validationErrors := make(map[string]string) if library.Name == "" { @@ -320,7 +320,7 @@ func (r *libraryRepositoryWrapper) validateLibrary(library *model.Library) error validationErrors["path"] = "ra.validation.required" } else { // Validate path format and accessibility - if err := r.validateLibraryPath(library); err != nil { + if err := r.validateLibraryPath(ctx, library); err != nil { validationErrors["path"] = err.Error() } } @@ -332,7 +332,7 @@ func (r *libraryRepositoryWrapper) validateLibrary(library *model.Library) error return nil } -func (r *libraryRepositoryWrapper) validateLibraryPath(library *model.Library) error { +func (r *libraryRepositoryWrapper) validateLibraryPath(ctx context.Context, library *model.Library) error { // Validate path format if !filepath.IsAbs(library.Path) { return fmt.Errorf("library path must be absolute") @@ -350,7 +350,7 @@ func (r *libraryRepositoryWrapper) validateLibraryPath(library *model.Library) e fsys, err := fileStore.FS() if err != nil { - log.Warn(r.ctx, "Error validating library.path", "path", library.Path, err) + log.Warn(ctx, "Error validating library.path", "path", library.Path, err) return fmt.Errorf("resources.library.validation.pathInvalid") } @@ -358,7 +358,7 @@ func (r *libraryRepositoryWrapper) validateLibraryPath(library *model.Library) e info, err := fs.Stat(fsys, ".") if err != nil { // Parse the error message to check for "not a directory" - log.Warn(r.ctx, "Error stating library.path", "path", library.Path, err) + log.Warn(ctx, "Error stating library.path", "path", library.Path, err) errStr := err.Error() if strings.Contains(errStr, "not a directory") || strings.Contains(errStr, "The directory name is invalid.") { @@ -385,7 +385,7 @@ func (s *libraryService) validateLibraryIDs(ctx context.Context, libraryIDs []in } // Use CountAll to efficiently validate library IDs exist - count, err := s.ds.Library(ctx).CountAll(model.QueryOptions{ + count, err := s.ds.Library().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"id": libraryIDs}, }) if err != nil { @@ -399,13 +399,13 @@ func (s *libraryService) validateLibraryIDs(ctx context.Context, libraryIDs []in return nil } -func (r *libraryRepositoryWrapper) triggerScan(lib *model.Library, action string) { - log.Info(r.ctx, fmt.Sprintf("Triggering scan for %s library", action), "libraryID", lib.ID, "name", lib.Name, "path", lib.Path) +func (r *libraryRepositoryWrapper) triggerScan(ctx context.Context, lib *model.Library, action string) { + log.Info(ctx, fmt.Sprintf("Triggering scan for %s library", action), "libraryID", lib.ID, "name", lib.Name, "path", lib.Path) start := time.Now() - warnings, err := r.scanner.ScanAll(r.ctx, false) // Quick scan for new library + warnings, err := r.scanner.ScanAll(ctx, false) // Quick scan for new library if err != nil { - log.Error(r.ctx, fmt.Sprintf("Error scanning %s library", action), "libraryID", lib.ID, "name", lib.Name, err) + log.Error(ctx, fmt.Sprintf("Error scanning %s library", action), "libraryID", lib.ID, "name", lib.Name, err) } else { - log.Info(r.ctx, fmt.Sprintf("Scan completed for %s library", action), "libraryID", lib.ID, "name", lib.Name, "warnings", len(warnings), "elapsed", time.Since(start)) + log.Info(ctx, fmt.Sprintf("Scan completed for %s library", action), "libraryID", lib.ID, "name", lib.Name, "warnings", len(warnings), "elapsed", time.Since(start)) } } diff --git a/core/library_test.go b/core/library_test.go index 43097414d..5402eac22 100644 --- a/core/library_test.go +++ b/core/library_test.go @@ -66,18 +66,18 @@ var _ = Describe("Library Service", func() { }) Describe("Library CRUD Operations", func() { - var repo rest.Persistable + var repo rest.Persistable[model.Library] BeforeEach(func() { - r := service.NewRepository(ctx) - repo = r.(rest.Persistable) + r := service.Repository() + repo = r.(rest.Persistable[model.Library]) }) Describe("Create", func() { It("creates a new library successfully", func() { library := &model.Library{ID: 1, Name: "New Library", Path: tempDir} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).NotTo(HaveOccurred()) Expect(libraryRepo.Data[1].Name).To(Equal("New Library")) @@ -87,7 +87,7 @@ var _ = Describe("Library Service", func() { It("fails when library name is empty", func() { library := &model.Library{Path: tempDir} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("ra.validation.required")) @@ -96,7 +96,7 @@ var _ = Describe("Library Service", func() { It("fails when library path is empty", func() { library := &model.Library{Name: "Test"} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("ra.validation.required")) @@ -105,7 +105,7 @@ var _ = Describe("Library Service", func() { It("fails when library path is not absolute", func() { library := &model.Library{Name: "Test", Path: "relative/path"} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -140,7 +140,7 @@ var _ = Describe("Library Service", func() { return errors.New("UNIQUE constraint failed: library.name") } - _, err = repo.Save(library) + _, err = repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -157,7 +157,7 @@ var _ = Describe("Library Service", func() { return errors.New("UNIQUE constraint failed: library.path") } - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -181,7 +181,7 @@ var _ = Describe("Library Service", func() { library := &model.Library{ID: 1, Name: "Updated Library", Path: newTempDir} - err = repo.Update("1", library) + err = repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) Expect(libraryRepo.Data[1].Name).To(Equal("Updated Library")) @@ -191,7 +191,7 @@ var _ = Describe("Library Service", func() { It("forwards the columns sent by the client to the repository", func() { library := &model.Library{ID: 1, Name: "Updated Library", Path: tempDir} - err := repo.Update("1", library, "name", "path") + err := repo.Update(ctx, "1", *library, "name", "path") Expect(err).NotTo(HaveOccurred()) Expect(libraryRepo.PutCols).To(Equal([]string{"name", "path"})) @@ -205,7 +205,7 @@ var _ = Describe("Library Service", func() { library := &model.Library{ID: 999, Name: "Non-existent", Path: uniqueTempDir} - err = repo.Update("999", library) + err = repo.Update(ctx, "999", *library) Expect(err).To(HaveOccurred()) Expect(err).To(Equal(model.ErrNotFound)) @@ -214,7 +214,7 @@ var _ = Describe("Library Service", func() { It("fails when library name is empty", func() { library := &model.Library{ID: 1, Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("ra.validation.required")) @@ -224,7 +224,7 @@ var _ = Describe("Library Service", func() { unnormalizedPath := tempDir + "//../" + filepath.Base(tempDir) library := &model.Library{ID: 1, Name: "Updated Library", Path: unnormalizedPath} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) Expect(libraryRepo.Data[1].Path).To(Equal(filepath.Clean(unnormalizedPath))) @@ -239,7 +239,7 @@ var _ = Describe("Library Service", func() { // Update the library keeping the same name (should be allowed) library := &model.Library{ID: 1, Name: "Test Library", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) }) @@ -253,7 +253,7 @@ var _ = Describe("Library Service", func() { // Update the library keeping the same path (should be allowed) library := &model.Library{ID: 1, Name: "Test Library", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) }) @@ -284,7 +284,7 @@ var _ = Describe("Library Service", func() { // Try to update library 2 to have the same name as library 1 library := &model.Library{ID: 2, Name: "Library One", Path: otherTempDir} - err = repo.Update("2", library) + err = repo.Update(ctx, "2", *library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -312,7 +312,7 @@ var _ = Describe("Library Service", func() { // Try to update library 2 to have the same path as library 1 library := &model.Library{ID: 2, Name: "Library Two", Path: tempDir} - err = repo.Update("2", library) + err = repo.Update(ctx, "2", *library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -327,7 +327,7 @@ var _ = Describe("Library Service", func() { It("fails when path is not absolute", func() { library := &model.Library{Name: "Test", Path: "relative/path"} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -339,7 +339,7 @@ var _ = Describe("Library Service", func() { nonExistentPath := filepath.Join(tempDir, "nonexistent") library := &model.Library{Name: "Test", Path: nonExistentPath} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -354,7 +354,7 @@ var _ = Describe("Library Service", func() { library := &model.Library{Name: "Test", Path: testFile} - _, err = repo.Save(library) + _, err = repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -371,7 +371,7 @@ var _ = Describe("Library Service", func() { It("handles multiple validation errors", func() { library := &model.Library{Name: "", Path: "relative/path"} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -393,7 +393,7 @@ var _ = Describe("Library Service", func() { It("fails when updated path is not absolute", func() { library := &model.Library{ID: 1, Name: "Test", Path: "relative/path"} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -410,7 +410,7 @@ var _ = Describe("Library Service", func() { // Update the library keeping the same name (should be allowed) library := &model.Library{ID: 1, Name: "Test Library", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) }) @@ -419,7 +419,7 @@ var _ = Describe("Library Service", func() { nonExistentPath := filepath.Join(tempDir, "nonexistent") library := &model.Library{ID: 1, Name: "Test", Path: nonExistentPath} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -434,7 +434,7 @@ var _ = Describe("Library Service", func() { library := &model.Library{ID: 1, Name: "Test", Path: testFile} - err = repo.Update("1", library) + err = repo.Update(ctx, "1", *library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -446,7 +446,7 @@ var _ = Describe("Library Service", func() { // Try to update with empty name and invalid path library := &model.Library{ID: 1, Name: "", Path: "relative/path"} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).To(HaveOccurred()) var validationErr *rest.ValidationError @@ -467,14 +467,14 @@ var _ = Describe("Library Service", func() { }) It("deletes an existing library successfully", func() { - err := repo.Delete("1") + err := repo.Delete(ctx, "1") Expect(err).NotTo(HaveOccurred()) Expect(libraryRepo.Data).To(HaveLen(0)) }) It("fails when library doesn't exist", func() { - err := repo.Delete("999") + err := repo.Delete(ctx, "999") Expect(err).To(HaveOccurred()) Expect(err).To(Equal(model.ErrNotFound)) @@ -613,17 +613,17 @@ var _ = Describe("Library Service", func() { }) Describe("Scan Triggering", func() { - var repo rest.Persistable + var repo rest.Persistable[model.Library] BeforeEach(func() { - r := service.NewRepository(ctx) - repo = r.(rest.Persistable) + r := service.Repository() + repo = r.(rest.Persistable[model.Library]) }) It("triggers scan when creating a new library", func() { library := &model.Library{ID: 1, Name: "New Library", Path: tempDir} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).NotTo(HaveOccurred()) // Wait briefly for the goroutine to complete @@ -649,7 +649,7 @@ var _ = Describe("Library Service", func() { // Update the library with a new path library := &model.Library{ID: 1, Name: "Updated Library", Path: newTempDir} - err = repo.Update("1", library) + err = repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) // Wait briefly for the goroutine to complete @@ -670,7 +670,7 @@ var _ = Describe("Library Service", func() { // Update the library name only (same path) library := &model.Library{ID: 1, Name: "Updated Name", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) // Wait a bit to ensure no scan was triggered @@ -683,7 +683,7 @@ var _ = Describe("Library Service", func() { // Try to create library with invalid data (empty name) library := &model.Library{Path: tempDir} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).To(HaveOccurred()) // Ensure no scan was triggered since creation failed @@ -700,7 +700,7 @@ var _ = Describe("Library Service", func() { // Try to update with invalid data (empty name) library := &model.Library{ID: 1, Name: "", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).To(HaveOccurred()) // Ensure no scan was triggered since update failed @@ -716,7 +716,7 @@ var _ = Describe("Library Service", func() { }) // Delete the library - err := repo.Delete("1") + err := repo.Delete(ctx, "1") Expect(err).NotTo(HaveOccurred()) // Wait briefly for the goroutine to complete @@ -731,7 +731,7 @@ var _ = Describe("Library Service", func() { It("does not trigger scan when library deletion fails", func() { // Try to delete a non-existent library - err := repo.Delete("999") + err := repo.Delete(ctx, "999") Expect(err).To(HaveOccurred()) // Ensure no scan was triggered since deletion failed @@ -744,7 +744,7 @@ var _ = Describe("Library Service", func() { It("starts watcher when creating a new library", func() { library := &model.Library{ID: 1, Name: "New Library", Path: tempDir} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).NotTo(HaveOccurred()) // Verify watcher was started @@ -773,7 +773,7 @@ var _ = Describe("Library Service", func() { // Update library with new path library := &model.Library{ID: 1, Name: "Updated Library", Path: newTempDir} - err = repo.Update("1", library) + err = repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) // Verify watcher was restarted @@ -793,7 +793,7 @@ var _ = Describe("Library Service", func() { // Update library with same path but different name library := &model.Library{ID: 1, Name: "Updated Name", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) // Verify watcher was NOT restarted (since path didn't change) @@ -808,7 +808,7 @@ var _ = Describe("Library Service", func() { {ID: 1, Name: "Test Library", Path: tempDir}, }) - err := repo.Delete("1") + err := repo.Delete(ctx, "1") Expect(err).NotTo(HaveOccurred()) // Verify watcher was stopped @@ -826,7 +826,7 @@ var _ = Describe("Library Service", func() { }) // Mock deletion to fail by trying to delete non-existent library - err := repo.Delete("999") + err := repo.Delete(ctx, "999") Expect(err).To(HaveOccurred()) // Verify watcher was NOT stopped since deletion failed @@ -838,11 +838,11 @@ var _ = Describe("Library Service", func() { }) Describe("Event Broadcasting", func() { - var repo rest.Persistable + var repo rest.Persistable[model.Library] BeforeEach(func() { - r := service.NewRepository(ctx) - repo = r.(rest.Persistable) + r := service.Repository() + repo = r.(rest.Persistable[model.Library]) // Clear any events from broker broker.Events = []events.Event{} }) @@ -850,7 +850,7 @@ var _ = Describe("Library Service", func() { It("sends refresh event when creating a library", func() { library := &model.Library{ID: 1, Name: "New Library", Path: tempDir} - _, err := repo.Save(library) + _, err := repo.Save(ctx, library) Expect(err).NotTo(HaveOccurred()) Expect(broker.Events).To(HaveLen(1)) @@ -863,7 +863,7 @@ var _ = Describe("Library Service", func() { }) library := &model.Library{ID: 1, Name: "Updated Library", Path: tempDir} - err := repo.Update("1", library) + err := repo.Update(ctx, "1", *library) Expect(err).NotTo(HaveOccurred()) Expect(broker.Events).To(HaveLen(1)) @@ -875,7 +875,7 @@ var _ = Describe("Library Service", func() { {ID: 2, Name: "Library to Delete", Path: tempDir}, }) - err := repo.Delete("2") + err := repo.Delete(ctx, "2") Expect(err).NotTo(HaveOccurred()) Expect(broker.Events).To(HaveLen(1)) @@ -883,13 +883,13 @@ var _ = Describe("Library Service", func() { }) Describe("Plugin Manager Integration", func() { - var repo rest.Persistable + var repo rest.Persistable[model.Library] BeforeEach(func() { // Reset the call count for each test pluginManager.unloadCalls = 0 - r := service.NewRepository(ctx) - repo = r.(rest.Persistable) + r := service.Repository() + repo = r.(rest.Persistable[model.Library]) }) It("calls UnloadDisabledPlugins after successful library deletion", func() { @@ -897,14 +897,14 @@ var _ = Describe("Library Service", func() { {ID: 2, Name: "Library to Delete", Path: tempDir}, }) - err := repo.Delete("2") + err := repo.Delete(ctx, "2") Expect(err).NotTo(HaveOccurred()) Expect(pluginManager.unloadCalls).To(Equal(1)) }) It("does not call UnloadDisabledPlugins when library deletion fails", func() { // Try to delete non-existent library - err := repo.Delete("999") + err := repo.Delete(ctx, "999") Expect(err).To(HaveOccurred()) Expect(pluginManager.unloadCalls).To(Equal(0)) }) diff --git a/core/lyrics/lyrics.go b/core/lyrics/lyrics.go index b9fb8cb74..d3e8d72b8 100644 --- a/core/lyrics/lyrics.go +++ b/core/lyrics/lyrics.go @@ -57,7 +57,7 @@ func (l *lyricsService) GetLyrics(ctx context.Context, mf *model.MediaFile) (mod func (l *lyricsService) GetLyricsByArtistTitle(ctx context.Context, artist, title string) (model.LyricList, error) { opts := songsByArtistTitleWithLyricsFirst(artist, title) opts.Max = maxLegacyLyricsCandidates - mediaFiles, err := l.ds.MediaFile(ctx).GetAll(opts) + mediaFiles, err := l.ds.MediaFile().GetAll(ctx, opts) if err != nil { return nil, err } diff --git a/core/maintenance.go b/core/maintenance.go index 56c0ac18d..58265ea2c 100644 --- a/core/maintenance.go +++ b/core/maintenance.go @@ -58,7 +58,7 @@ func (s *maintenanceService) RemapMissingFile(ctx context.Context, missingID, ta return fmt.Errorf("%w: %q", ErrSameFile, missingID) } - missing, err := s.ds.MediaFile(ctx).Get(missingID) + missing, err := s.ds.MediaFile().Get(ctx, missingID) if err != nil { return fmt.Errorf("loading missing file %q: %w", missingID, err) } @@ -66,7 +66,7 @@ func (s *maintenanceService) RemapMissingFile(ctx context.Context, missingID, ta return fmt.Errorf("%w: %q", ErrNotMissing, missingID) } - target, err := s.ds.MediaFile(ctx).GetWithParticipants(targetID) + target, err := s.ds.MediaFile().GetWithParticipants(ctx, targetID) if err != nil { return fmt.Errorf("loading target file %q: %w", targetID, err) } @@ -82,27 +82,27 @@ func (s *maintenanceService) RemapMissingFile(ctx context.Context, missingID, ta // Preserve the original created_at so the remapped track doesn't resurface in "Recently Added" target.CreatedAt = missing.CreatedAt target.ID = missing.ID - if err := tx.MediaFile(ctx).Put(target); err != nil { + if err := tx.MediaFile().Put(ctx, target); err != nil { return fmt.Errorf("update matched track: %w", err) } // Unlike the scanner's freshly-imported target, this one may carry history of its own - if err := tx.MediaFile(ctx).ReassignReferences(discardedID, missing.ID); err != nil { + if err := tx.MediaFile().ReassignReferences(ctx, discardedID, missing.ID); err != nil { return fmt.Errorf("reassign target references: %w", err) } - if err := tx.MediaFile(ctx).Delete(discardedID); err != nil { + if err := tx.MediaFile().Delete(ctx, discardedID); err != nil { return fmt.Errorf("delete discarded track: %w", err) } if oldAlbumID != newAlbumID { - oldAlbumTracks, err := tx.MediaFile(ctx).CountAll(model.QueryOptions{Filters: squirrel.Eq{"album_id": oldAlbumID}}) + oldAlbumTracks, err := tx.MediaFile().CountAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album_id": oldAlbumID}}) if err != nil { return fmt.Errorf("get old album tracks: %w", err) } if oldAlbumTracks == 0 { - if err := tx.Album(ctx).ReassignAnnotation(oldAlbumID, newAlbumID); err != nil { + if err := tx.Album().ReassignAnnotation(ctx, oldAlbumID, newAlbumID); err != nil { return fmt.Errorf("reassign album annotations: %w", err) } - if err := tx.Album(ctx).CopyAttributes(oldAlbumID, newAlbumID, "created_at"); err != nil && !errors.Is(err, model.ErrNotFound) { + if err := tx.Album().CopyAttributes(ctx, oldAlbumID, newAlbumID, "created_at"); err != nil && !errors.Is(err, model.ErrNotFound) { return fmt.Errorf("copy album attributes: %w", err) } } @@ -121,7 +121,7 @@ func (s *maintenanceService) RemapMissingFile(ctx context.Context, missingID, ta // Stats are refreshed synchronously, unlike deleteMissing, so the CLI sees them before it exits. // album/artist play count aggregates are not recalculated here; they are refreshed by the next scan. - if _, err := s.ds.Artist(ctx).RefreshStats(true); err != nil { + if _, err := s.ds.Artist().RefreshStats(ctx, true); err != nil { log.Error(ctx, "Error refreshing artist stats after remapping missing file", err) } affectedAlbumIDs := []string{newAlbumID} @@ -146,10 +146,10 @@ func (s *maintenanceService) deleteMissing(ctx context.Context, ids []string) er // Delete missing files within a transaction err = s.ds.WithTx(func(tx model.DataStore) error { if len(ids) == 0 { - _, err := tx.MediaFile(ctx).DeleteAllMissing() + _, err := tx.MediaFile().DeleteAllMissing(ctx) return err } - return tx.MediaFile(ctx).DeleteMissing(ids) + return tx.MediaFile().DeleteMissing(ctx, ids) }) if err != nil { log.Error(ctx, "Error deleting missing tracks from DB", "ids", ids, err) @@ -192,11 +192,11 @@ func (s *maintenanceService) refreshAlbums(ctx context.Context, albumIDs []strin // refreshAlbumChunk processes a single chunk of album IDs func (s *maintenanceService) refreshAlbumChunk(ctx context.Context, albumIDs []string) error { - albumRepo := s.ds.Album(ctx) - mfRepo := s.ds.MediaFile(ctx) + albumRepo := s.ds.Album() + mfRepo := s.ds.MediaFile() // Batch load existing albums - albums, err := albumRepo.GetAll(model.QueryOptions{ + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.id": albumIDs}, }) if err != nil { @@ -210,7 +210,7 @@ func (s *maintenanceService) refreshAlbumChunk(ctx context.Context, albumIDs []s } // Batch load all media files for these albums - mediaFiles, err := mfRepo.GetAll(model.QueryOptions{ + mediaFiles, err := mfRepo.GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album_id": albumIDs}, Sort: "album_id, path", }) @@ -243,7 +243,7 @@ func (s *maintenanceService) refreshAlbumChunk(ctx context.Context, albumIDs []s newAlbum.UpdatedAt = time.Now() newAlbum.CreatedAt = oldAlbum.CreatedAt - if err := albumRepo.Put(&newAlbum); err != nil { + if err := albumRepo.Put(ctx, &newAlbum); err != nil { log.Error(ctx, "Error updating album during refresh", "albumID", albumID, err) // Continue with other albums instead of failing entirely continue @@ -265,7 +265,7 @@ func (s *maintenanceService) getAffectedAlbumIDs(ctx context.Context, ids []stri } } - mfs, err := s.ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := s.ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: filters, }) if err != nil { @@ -293,7 +293,7 @@ func (s *maintenanceService) refreshStatsAsync(ctx context.Context, affectedAlbu // Refresh artist stats in background s.wg.Go(func() { bgCtx := request.AddValues(context.Background(), ctx) - if _, err := s.ds.Artist(bgCtx).RefreshStats(true); err != nil { + if _, err := s.ds.Artist().RefreshStats(bgCtx, true); err != nil { log.Error(bgCtx, "Error refreshing artist stats after deleting missing files", err) } else { log.Debug(bgCtx, "Successfully refreshed artist stats after deleting missing files") diff --git a/core/maintenance_test.go b/core/maintenance_test.go index 56e6d13b5..4ffc098d9 100644 --- a/core/maintenance_test.go +++ b/core/maintenance_test.go @@ -262,12 +262,12 @@ var _ = Describe("Maintenance", func() { Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) - got, err := mfRepo.Get("m1") + got, err := mfRepo.Get(ctx, "m1") Expect(err).ToNot(HaveOccurred()) Expect(got.Path).To(Equal("new/song.mp3")) // moved to target's location Expect(got.Missing).To(BeFalse()) Expect(got.CreatedAt).To(BeTemporally("==", created)) // created_at preserved - exists, _ := mfRepo.Exists("t1") + exists, _ := mfRepo.Exists(ctx, "t1") Expect(exists).To(BeFalse()) // discarded row removed Expect(ds.GCCalled).To(BeTrue()) }) @@ -369,14 +369,14 @@ var _ = Describe("Maintenance", func() { Expect(artistRepo.IsRefreshStatsCalled()).To(BeTrue(), "Artist stats should be refreshed") // The old album lost the remapped track, so its stats are recalculated from the remaining one - oldAlbum, err := albumRepo.Get("album1") + oldAlbum, err := albumRepo.Get(ctx, "album1") Expect(err).ToNot(HaveOccurred()) Expect(oldAlbum.SongCount).To(Equal(1)) Expect(oldAlbum.Size).To(Equal(int64(1000))) Expect(oldAlbum.Duration).To(BeNumerically("==", 100)) // The target album keeps the track, now under the missing file's ID - newAlbum, err := albumRepo.Get("album2") + newAlbum, err := albumRepo.Get(ctx, "album2") Expect(err).ToNot(HaveOccurred()) Expect(newAlbum.SongCount).To(Equal(1)) Expect(newAlbum.Size).To(Equal(int64(2000))) @@ -407,7 +407,7 @@ var _ = Describe("Maintenance", func() { Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) // The surviving row is the missing file's ID, holding the target's data - got, err := mfRepo.GetWithParticipants("m1") + got, err := mfRepo.GetWithParticipants(ctx, "m1") Expect(err).ToNot(HaveOccurred()) Expect(got.Participants).To(HaveKeyWithValue(model.RoleArtist, model.ParticipantList{participant})) }) @@ -447,7 +447,7 @@ type extendedMediaFileRepo struct { deleteMissingError error } -func (m *extendedMediaFileRepo) DeleteMissing(ids []string) error { +func (m *extendedMediaFileRepo) DeleteMissing(ctx context.Context, ids []string) error { m.deleteMissingCalled = true m.deletedIDs = ids if m.deleteMissingError != nil { @@ -470,7 +470,7 @@ type extendedAlbumRepo struct { failOnce bool } -func (m *extendedAlbumRepo) Put(album *model.Album) error { +func (m *extendedAlbumRepo) Put(ctx context.Context, album *model.Album) error { m.mu.Lock() m.putCallCount++ m.lastPutData = album @@ -490,7 +490,7 @@ func (m *extendedAlbumRepo) Put(album *model.Album) error { } m.mu.Unlock() - return m.MockAlbumRepo.Put(album) + return m.MockAlbumRepo.Put(ctx, album) } func (m *extendedAlbumRepo) GetPutCallCount() int { @@ -507,7 +507,7 @@ type extendedArtistRepo struct { refreshStatsError error } -func (m *extendedArtistRepo) RefreshStats(allArtists bool) (int64, error) { +func (m *extendedArtistRepo) RefreshStats(ctx context.Context, allArtists bool) (int64, error) { m.mu.Lock() m.refreshStatsCalled = true err := m.refreshStatsError @@ -516,7 +516,7 @@ func (m *extendedArtistRepo) RefreshStats(allArtists bool) (int64, error) { if err != nil { return 0, err } - return m.MockArtistRepo.RefreshStats(allArtists) + return m.MockArtistRepo.RefreshStats(ctx, allArtists) } func (m *extendedArtistRepo) IsRefreshStatsCalled() bool { diff --git a/core/matcher/matcher.go b/core/matcher/matcher.go index 25b8fda5f..3eb4db152 100644 --- a/core/matcher/matcher.go +++ b/core/matcher/matcher.go @@ -95,7 +95,7 @@ func (m *Matcher) matchByID(ctx context.Context, songs []agents.Song, result map if len(ids) == 0 { return nil } - res, err := m.ds.MediaFile(ctx).GetAll(model.QueryOptions{ + res, err := m.ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Eq{"media_file.id": ids}, squirrel.Eq{"missing": false}, @@ -134,7 +134,7 @@ func (m *Matcher) matchByMBID(ctx context.Context, songs []agents.Song, result m if len(mbids) == 0 { return nil } - res, err := m.ds.MediaFile(ctx).GetAll(model.QueryOptions{ + res, err := m.ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Eq{"mbz_recording_id": mbids}, squirrel.Eq{"missing": false}, @@ -180,7 +180,7 @@ func (m *Matcher) matchByISRC(ctx context.Context, songs []agents.Song, result m if len(isrcs) == 0 { return nil } - res, err := m.ds.MediaFile(ctx).GetAllByTags(model.TagISRC, isrcs, model.QueryOptions{ + res, err := m.ds.MediaFile().GetAllByTags(ctx, model.TagISRC, isrcs, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, Sort: "starred desc, rating desc, year asc, compilation asc", }) @@ -442,7 +442,7 @@ func (m *Matcher) resolveArtists(ctx context.Context, queries []indexedQuery) (r filter = append(filter, squirrel.Eq{"id": slices.Collect(maps.Keys(allIDs))}) } if len(filter) > 0 { - artists, err := m.ds.Artist(ctx).GetAll(model.QueryOptions{Filters: filter}) + artists, err := m.ds.Artist().GetAll(ctx, model.QueryOptions{Filters: filter}) if err != nil { return resolvedArtists{}, err } @@ -543,7 +543,7 @@ func (m *Matcher) fetchTracksCreditedTo(ctx context.Context, artistIDs []string) return nil, nil } args := slice.Map(artistIDs, func(id string) any { return id }) - return m.ds.MediaFile(ctx).GetAll(model.QueryOptions{ + return m.ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Expr( "media_file.id IN (SELECT media_file_id FROM media_file_artists "+ diff --git a/core/matcher/matcher_test.go b/core/matcher/matcher_test.go index a46db8a09..0f5389ca5 100644 --- a/core/matcher/matcher_test.go +++ b/core/matcher/matcher_test.go @@ -1267,7 +1267,7 @@ func newMockMediaFileRepo() *mockMediaFileRepo { return &mockMediaFileRepo{} } -func (m *mockMediaFileRepo) GetAll(options ...model.QueryOptions) (model.MediaFiles, error) { +func (m *mockMediaFileRepo) GetAll(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { argsSlice := make([]any, len(options)) for i, v := range options { argsSlice[i] = v @@ -1279,8 +1279,8 @@ func (m *mockMediaFileRepo) GetAll(options ...model.QueryOptions) (model.MediaFi return args.Get(0).(model.MediaFiles), args.Error(1) } -func (m *mockMediaFileRepo) GetAllByTags(_ model.TagName, _ []string, options ...model.QueryOptions) (model.MediaFiles, error) { - return m.GetAll(options...) +func (m *mockMediaFileRepo) GetAllByTags(ctx context.Context, _ model.TagName, _ []string, options ...model.QueryOptions) (model.MediaFiles, error) { + return m.GetAll(ctx, options...) } func (m *mockMediaFileRepo) SetError(hasError bool) { @@ -1298,7 +1298,7 @@ func newMockArtistRepo() *mockArtistRepo { return &mockArtistRepo{} } -func (m *mockArtistRepo) GetAll(options ...model.QueryOptions) (model.Artists, error) { +func (m *mockArtistRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Artists, error) { argsSlice := make([]any, len(options)) for i, v := range options { argsSlice[i] = v diff --git a/core/metrics/insights.go b/core/metrics/insights.go index d952f517a..4a78a7f3f 100644 --- a/core/metrics/insights.go +++ b/core/metrics/insights.go @@ -47,11 +47,11 @@ type insightsCollector struct { func GetInstance(ds model.DataStore) Insights { return singleton.GetInstance(func() *insightsCollector { - id, err := ds.Property(context.TODO()).Get(consts.InsightsIDKey) + id, err := ds.Property().Get(context.TODO(), consts.InsightsIDKey) if err != nil { log.Trace("Could not get Insights ID from DB. Creating one", err) id = uuid.NewString() - err = ds.Property(context.TODO()).Put(consts.InsightsIDKey, id) + err = ds.Property().Put(context.TODO(), consts.InsightsIDKey, id) if err != nil { log.Trace("Could not save Insights ID to DB", err) } @@ -87,7 +87,7 @@ func (c *insightsCollector) LastRun(context.Context) (timestamp time.Time, succe } func (c *insightsCollector) sendInsights(ctx context.Context) { - count, err := c.ds.User(ctx).CountAll(model.QueryOptions{}) + count, err := c.ds.User().CountAll(ctx, model.QueryOptions{}) if err != nil { log.Trace(ctx, "Could not check user count", err) return @@ -245,41 +245,41 @@ func (c *insightsCollector) collect(ctx context.Context) []byte { // Library info var err error - data.Library.Tracks, err = c.ds.MediaFile(ctx).CountAll() + data.Library.Tracks, err = c.ds.MediaFile().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading tracks count", err) } - data.Library.Albums, err = c.ds.Album(ctx).CountAll() + data.Library.Albums, err = c.ds.Album().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading albums count", err) } - data.Library.Artists, err = c.ds.Artist(ctx).CountAll() + data.Library.Artists, err = c.ds.Artist().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading artists count", err) } - data.Library.Playlists, err = c.ds.Playlist(ctx).CountAll() + data.Library.Playlists, err = c.ds.Playlist().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading playlists count", err) } - data.Library.Shares, err = c.ds.Share(ctx).CountAll() + data.Library.Shares, err = c.ds.Share().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading shares count", err) } - data.Library.Radios, err = c.ds.Radio(ctx).Count() + data.Library.Radios, err = c.ds.Radio().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading radios count", err) } - data.Library.Libraries, err = c.ds.Library(ctx).CountAll() + data.Library.Libraries, err = c.ds.Library().CountAll(ctx) if err != nil { log.Trace(ctx, "Error reading libraries count", err) } - data.Library.ActiveUsers, err = c.ds.User(ctx).CountAll(model.QueryOptions{ + data.Library.ActiveUsers, err = c.ds.User().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Gt{"last_access_at": time.Now().Add(-7 * 24 * time.Hour)}, }) if err != nil { log.Trace(ctx, "Error reading active users count", err) } - data.Library.FileSuffixes, err = c.ds.MediaFile(ctx).CountBySuffix() + data.Library.FileSuffixes, err = c.ds.MediaFile().CountBySuffix(ctx) if err != nil { log.Trace(ctx, "Error reading file suffixes count", err) } @@ -297,7 +297,7 @@ func (c *insightsCollector) collect(ctx context.Context) []byte { // Collect active players if permitted if conf.Server.DevEnablePlayerInsights { - data.Library.ActivePlayers, err = c.ds.Player(ctx).CountByClient(model.QueryOptions{ + data.Library.ActivePlayers, err = c.ds.Player().CountByClient(ctx, model.QueryOptions{ Filters: squirrel.Gt{"last_seen": time.Now().Add(-7 * 24 * time.Hour)}, }) if err != nil { @@ -324,7 +324,7 @@ func (c *insightsCollector) collect(ctx context.Context) []byte { // hasSmartPlaylists checks if there are any smart playlists (playlists with rules) func (c *insightsCollector) hasSmartPlaylists(ctx context.Context) (bool, error) { - count, err := c.ds.Playlist(ctx).CountAll(model.QueryOptions{ + count, err := c.ds.Playlist().CountAll(ctx, model.QueryOptions{ Filters: squirrel.And{squirrel.NotEq{"rules": ""}, squirrel.NotEq{"rules": nil}}, }) return count > 0, err diff --git a/core/metrics/prometheus.go b/core/metrics/prometheus.go index 412483156..c85706560 100644 --- a/core/metrics/prometheus.go +++ b/core/metrics/prometheus.go @@ -197,28 +197,28 @@ var getPrometheusMetrics = sync.OnceValue(func() *prometheusMetrics { }) func processSqlAggregateMetrics(ctx context.Context, ds model.DataStore, targetGauge *prometheus.GaugeVec) { - albumsCount, err := ds.Album(ctx).CountAll() + albumsCount, err := ds.Album().CountAll(ctx) if err != nil { log.Warn("album CountAll error", err) return } targetGauge.With(prometheus.Labels{"model": "album"}).Set(float64(albumsCount)) - artistCount, err := ds.Artist(ctx).CountAll() + artistCount, err := ds.Artist().CountAll(ctx) if err != nil { log.Warn("artist CountAll error", err) return } targetGauge.With(prometheus.Labels{"model": "artist"}).Set(float64(artistCount)) - songsCount, err := ds.MediaFile(ctx).CountAll() + songsCount, err := ds.MediaFile().CountAll(ctx) if err != nil { log.Warn("media CountAll error", err) return } targetGauge.With(prometheus.Labels{"model": "media"}).Set(float64(songsCount)) - usersCount, err := ds.User(ctx).CountAll() + usersCount, err := ds.User().CountAll(ctx) if err != nil { log.Warn("user CountAll error", err) return diff --git a/core/playback/device.go b/core/playback/device.go index fd08b340e..8e4e8880e 100644 --- a/core/playback/device.go +++ b/core/playback/device.go @@ -23,7 +23,7 @@ type Track interface { } type playbackDevice struct { - serviceCtx context.Context + serviceCtx context.Context //nolint:containedctx // playback service lifecycle ctx ParentPlaybackServer PlaybackServer Default bool User string diff --git a/core/playback/playbackserver.go b/core/playback/playbackserver.go index 7dd02dcb1..48e3bcfaf 100644 --- a/core/playback/playbackserver.go +++ b/core/playback/playbackserver.go @@ -111,7 +111,7 @@ func (ps *playbackServer) getDefaultDevice() (*playbackDevice, error) { // GetMediaFile retrieves the MediaFile given by the id parameter func (ps *playbackServer) GetMediaFile(id string) (*model.MediaFile, error) { - return ps.datastore.MediaFile(*ps.ctx).Get(id) + return ps.datastore.MediaFile().Get(*ps.ctx, id) } // GetDeviceForUser returns the audio playback device for the given user. As of now this is but only the default device. diff --git a/core/players.go b/core/players.go index b757f8460..e03d8caa2 100644 --- a/core/players.go +++ b/core/players.go @@ -37,14 +37,14 @@ func (p *players) Register(ctx context.Context, playerID, client, userAgent, ip var err error user, _ := request.UserFrom(ctx) if playerID != "" { - plr, err = p.ds.Player(ctx).Get(playerID) + plr, err = p.ds.Player().Get(ctx, playerID) if err == nil && (plr.Client != client || plr.UserId != user.ID) { playerID = "" } } username := userName(ctx) if err != nil || playerID == "" { - plr, err = p.ds.Player(ctx).FindMatch(user.ID, client, userAgent) + plr, err = p.ds.Player().FindMatch(ctx, user.ID, client, userAgent) if err == nil { log.Debug(ctx, "Found matching player", "id", plr.ID, "client", client, "username", username, "type", userAgent) } else { @@ -66,17 +66,17 @@ func (p *players) Register(ctx context.Context, playerID, client, userAgent, ip ctx, cancel := context.WithTimeout(ctx, time.Second) defer cancel() - err = p.ds.Player(ctx).Put(plr) + err = p.ds.Player().Put(ctx, plr) if err != nil { log.Warn(ctx, "Could not save player", "id", plr.ID, "client", client, "username", username, "type", userAgent, err) } }) if plr.TranscodingId != "" { - trc, err = p.ds.Transcoding(ctx).Get(plr.TranscodingId) + trc, err = p.ds.Transcoding().Get(ctx, plr.TranscodingId) } return plr, trc, err } func (p *players) Get(ctx context.Context, playerId string) (*model.Player, error) { - return p.ds.Player(ctx).Get(playerId) + return p.ds.Player().Get(ctx, playerId) } diff --git a/core/players_test.go b/core/players_test.go index 55ec16833..302d63157 100644 --- a/core/players_test.go +++ b/core/players_test.go @@ -145,14 +145,14 @@ func (m *mockPlayerRepository) add(p *model.Player) { m.data[p.ID] = *p } -func (m *mockPlayerRepository) Get(id string) (*model.Player, error) { +func (m *mockPlayerRepository) Get(_ context.Context, id string) (*model.Player, error) { if p, ok := m.data[id]; ok { return &p, nil } return nil, model.ErrNotFound } -func (m *mockPlayerRepository) FindMatch(userId, client, userAgent string) (*model.Player, error) { +func (m *mockPlayerRepository) FindMatch(_ context.Context, userId, client, userAgent string) (*model.Player, error) { for _, p := range m.data { if p.Client == client && p.UserId == userId && p.UserAgent == userAgent { return &p, nil @@ -161,7 +161,7 @@ func (m *mockPlayerRepository) FindMatch(userId, client, userAgent string) (*mod return nil, model.ErrNotFound } -func (m *mockPlayerRepository) Put(p *model.Player) error { +func (m *mockPlayerRepository) Put(_ context.Context, p *model.Player) error { m.lastSaved = p return nil } diff --git a/core/playlists/import.go b/core/playlists/import.go index e41f61bd1..658bd92dc 100644 --- a/core/playlists/import.go +++ b/core/playlists/import.go @@ -39,7 +39,7 @@ func (s *playlists) ImportFile(ctx context.Context, absolutePath string, sync bo } if pls.ID != "" && pls.Sync != sync { pls.Sync = sync - if putErr := s.ds.Playlist(ctx).Put(pls); putErr != nil { + if putErr := s.ds.Playlist().Put(ctx, pls); putErr != nil { return nil, putErr } } @@ -74,7 +74,7 @@ func (s *playlists) ImportFile(ctx context.Context, absolutePath string, sync bo var errNotInLibrary = fmt.Errorf("path not in any library") func (s *playlists) resolveFolder(ctx context.Context, dir string) (*model.Folder, error) { - libs, err := s.ds.Library(ctx).GetAll() + libs, err := s.ds.Library().GetAll(ctx) if err != nil { return nil, err } @@ -84,7 +84,7 @@ func (s *playlists) resolveFolder(ctx context.Context, dir string) (*model.Folde return nil, fmt.Errorf("%w: %s", errNotInLibrary, dir) } - folder, err := s.ds.Folder(ctx).GetByPath(lib, dir) + folder, err := s.ds.Folder().GetByPath(ctx, lib, dir) if err != nil { return nil, fmt.Errorf("resolving folder for path %s: %w", dir, err) } @@ -122,7 +122,7 @@ func (s *playlists) ImportM3U(ctx context.Context, reader io.Reader) (*model.Pla log.Error(ctx, "Error parsing playlist", err) return nil, err } - err = s.ds.Playlist(ctx).Put(pls) + err = s.ds.Playlist().Put(ctx, pls) if err != nil { log.Error(ctx, "Error saving playlist", err) return nil, err @@ -166,14 +166,14 @@ func fingerprint(h *xxh3.Hasher) string { // findByPathNormalized looks up a playlist by path, trying both NFC and NFD Unicode // normalization forms to handle cross-platform filesystem differences. func (s *playlists) findByPathNormalized(ctx context.Context, path string) (*model.Playlist, error) { - pls, err := s.ds.Playlist(ctx).FindByPath(path) + pls, err := s.ds.Playlist().FindByPath(ctx, path) if errors.Is(err, model.ErrNotFound) { altPath := norm.NFD.String(path) if altPath == path { altPath = norm.NFC.String(path) } if altPath != path { - pls, err = s.ds.Playlist(ctx).FindByPath(altPath) + pls, err = s.ds.Playlist().FindByPath(ctx, altPath) } } return pls, err @@ -221,5 +221,5 @@ func (s *playlists) updatePlaylist(ctx context.Context, newPls *model.Playlist, newPls.Public = conf.Server.DefaultPlaylistPublicVisibility } } - return s.ds.Playlist(ctx).Put(newPls) + return s.ds.Playlist().Put(ctx, newPls) } diff --git a/core/playlists/import_test.go b/core/playlists/import_test.go index a90a703d9..25960f0fe 100644 --- a/core/playlists/import_test.go +++ b/core/playlists/import_test.go @@ -1172,7 +1172,7 @@ type mockedMediaFileRepo struct { data map[string]model.MediaFile } -func (r *mockedMediaFileRepo) FindByPaths(paths []string) (model.MediaFiles, error) { +func (r *mockedMediaFileRepo) FindByPaths(ctx context.Context, paths []string) (model.MediaFiles, error) { var mfs model.MediaFiles // If data map provided, look up files @@ -1212,7 +1212,7 @@ type mockedMediaFileFromListRepo struct { data []string } -func (r *mockedMediaFileFromListRepo) FindByPaths(paths []string) (model.MediaFiles, error) { +func (r *mockedMediaFileFromListRepo) FindByPaths(ctx context.Context, paths []string) (model.MediaFiles, error) { var mfs model.MediaFiles for idx, dataPath := range r.data { @@ -1247,7 +1247,7 @@ type mockFolderRepoForImport struct { folder *model.Folder } -func (m *mockFolderRepoForImport) GetByPath(_ model.Library, _ string) (*model.Folder, error) { +func (m *mockFolderRepoForImport) GetByPath(_ context.Context, _ model.Library, _ string) (*model.Folder, error) { if m.folder != nil { return m.folder, nil } diff --git a/core/playlists/parse_m3u.go b/core/playlists/parse_m3u.go index b9cb154cb..9610e9dbb 100644 --- a/core/playlists/parse_m3u.go +++ b/core/playlists/parse_m3u.go @@ -20,7 +20,7 @@ import ( ) func (s *playlists) parseM3U(ctx context.Context, pls *model.Playlist, folder *model.Folder, reader io.Reader) error { - mediaFileRepository := s.ds.MediaFile(ctx) + mediaFileRepository := s.ds.MediaFile() resolver, err := newPathResolver(ctx, s.ds) if err != nil { return err @@ -96,7 +96,7 @@ func (s *playlists) parseM3U(ctx context.Context, pls *model.Playlist, folder *m } } - found, err := mediaFileRepository.FindByPaths(lookupCandidates) + found, err := mediaFileRepository.FindByPaths(ctx, lookupCandidates) if err != nil { log.Warn(ctx, "Error reading files from DB", "playlist", pls.Name, err) continue @@ -215,7 +215,7 @@ type pathResolver struct { // newPathResolver creates a pathResolver with libraries loaded from the datastore. func newPathResolver(ctx context.Context, ds model.DataStore) (*pathResolver, error) { - libs, err := ds.Library(ctx).GetAll() + libs, err := ds.Library().GetAll(ctx) if err != nil { return nil, err } diff --git a/core/playlists/parse_m3u_test.go b/core/playlists/parse_m3u_test.go index d7fd5e001..b6a3a96f9 100644 --- a/core/playlists/parse_m3u_test.go +++ b/core/playlists/parse_m3u_test.go @@ -24,7 +24,7 @@ var _ = Describe("libraryMatcher", func() { // Helper function to create a libraryMatcher from the mock datastore createMatcher := func(ds model.DataStore) *libraryMatcher { - libs, err := ds.Library(ctx).GetAll() + libs, err := ds.Library().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) return newLibraryMatcher(libs) } diff --git a/core/playlists/playlists.go b/core/playlists/playlists.go index 1f8cc9581..c9bc03b97 100644 --- a/core/playlists/playlists.go +++ b/core/playlists/playlists.go @@ -49,8 +49,8 @@ type Playlists interface { ImportM3U(ctx context.Context, reader io.Reader) (*model.Playlist, error) // REST adapters - NewRepository(ctx context.Context) rest.Repository - TracksRepository(ctx context.Context, playlistId string, refreshSmartPlaylist bool) rest.Repository + Repository() rest.Repository[model.Playlist] + TracksRepository(ctx context.Context, playlistId string, refreshSmartPlaylist bool) rest.Repository[model.PlaylistTrack] } // ImageUploadService is a local interface satisfied by artwork.Uploader. @@ -64,10 +64,13 @@ type ImageUploadService interface { type playlists struct { ds model.DataStore imgUpload ImageUploadService + repo *playlistRepositoryWrapper } func NewPlaylists(ds model.DataStore, imgUpload ImageUploadService) Playlists { - return &playlists{ds: ds, imgUpload: imgUpload} + s := &playlists{ds: ds, imgUpload: imgUpload} + s.repo = &playlistRepositoryWrapper{PlaylistRepository: ds.Playlist(), service: s} + return s } func InPath(folder model.Folder) bool { @@ -86,30 +89,30 @@ func InPath(folder model.Folder) bool { // --- Read operations --- func (s *playlists) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Playlists, error) { - return s.ds.Playlist(ctx).GetAll(options...) + return s.ds.Playlist().GetAll(ctx, options...) } func (s *playlists) Get(ctx context.Context, id string) (*model.Playlist, error) { - return s.ds.Playlist(ctx).Get(id) + return s.ds.Playlist().Get(ctx, id) } func (s *playlists) GetWithTracks(ctx context.Context, id string) (*model.Playlist, error) { - return s.ds.Playlist(ctx).GetWithTracks(id, true, false) + return s.ds.Playlist().GetWithTracks(ctx, id, true, false) } func (s *playlists) GetPlaylists(ctx context.Context, mediaFileId string) (model.Playlists, error) { - return s.ds.Playlist(ctx).GetPlaylists(mediaFileId) + return s.ds.Playlist().GetPlaylists(ctx, mediaFileId) } // Tracks scopes a repository to one playlist's tracks, for callers that page or stream them rather // than loading every one like GetWithTracks. Gets first because PlaylistRepository.Tracks discards // its error behind a nil (and warns), and this is probed with ids that are usually not playlists. func (s *playlists) Tracks(ctx context.Context, id string) (model.PlaylistTrackRepository, error) { - repo := s.ds.Playlist(ctx) - if _, err := repo.Get(id); err != nil { + repo := s.ds.Playlist() + if _, err := repo.Get(ctx, id); err != nil { return nil, err } - tracks := repo.Tracks(id, true) + tracks := repo.Tracks(ctx, id, true) if tracks == nil { return nil, model.ErrNotFound } @@ -127,7 +130,7 @@ func (s *playlists) Create(ctx context.Context, playlistId string, name string, var err error if playlistId != "" { - pls, err = tx.Playlist(ctx).Get(playlistId) + pls, err = tx.Playlist().Get(ctx, playlistId) if err != nil { return err } @@ -145,7 +148,7 @@ func (s *playlists) Create(ctx context.Context, playlistId string, name string, pls.Tracks = nil pls.AddMediaFilesByID(ids) - err = tx.Playlist(ctx).Put(pls) + err = tx.Playlist().Put(ctx, pls) playlistId = pls.ID return err }) @@ -165,7 +168,7 @@ func (s *playlists) Delete(ctx context.Context, id string) error { } } - return s.ds.Playlist(ctx).Delete(id) + return s.ds.Playlist().Delete(ctx, id) } func (s *playlists) Update(ctx context.Context, playlistID string, @@ -183,21 +186,21 @@ func (s *playlists) Update(ctx context.Context, playlistID string, return err } return s.ds.WithTxImmediate(func(tx model.DataStore) error { - repo := tx.Playlist(ctx) + repo := tx.Playlist() if len(idxToRemove) > 0 { - tracksRepo := repo.Tracks(playlistID, false) + tracksRepo := repo.Tracks(ctx, playlistID, false) // Convert 0-based indices to 1-based position IDs and delete them directly, // avoiding the need to load all tracks into memory. positions := make([]string, len(idxToRemove)) for i, idx := range idxToRemove { positions[i] = strconv.Itoa(idx + 1) } - if err := tracksRepo.Delete(positions...); err != nil { + if err := tracksRepo.Delete(ctx, positions...); err != nil { return err } if len(idsToAdd) > 0 { - if _, err := tracksRepo.Add(idsToAdd); err != nil { + if _, err := tracksRepo.Add(ctx, idsToAdd); err != nil { return err } } @@ -205,7 +208,7 @@ func (s *playlists) Update(ctx context.Context, playlistID string, } if len(idsToAdd) > 0 { - if _, err := repo.Tracks(playlistID, false).Add(idsToAdd); err != nil { + if _, err := repo.Tracks(ctx, playlistID, false).Add(ctx, idsToAdd); err != nil { return err } } @@ -221,7 +224,7 @@ func (s *playlists) Update(ctx context.Context, playlistID string, // checkWritable fetches the playlist and verifies the current user can modify it. func (s *playlists) checkWritable(ctx context.Context, id string) (*model.Playlist, error) { - pls, err := s.ds.Playlist(ctx).Get(id) + pls, err := s.ds.Playlist().Get(ctx, id) if err != nil { return nil, err } @@ -257,7 +260,7 @@ func (s *playlists) updateMetadata(ctx context.Context, ds model.DataStore, pls if public != nil { pls.Public = *public } - return ds.Playlist(ctx).Put(pls) + return ds.Playlist().Put(ctx, pls) } // --- Track management operations --- @@ -266,7 +269,7 @@ func (s *playlists) AddTracks(ctx context.Context, playlistID string, ids []stri if _, err := s.checkTracksEditable(ctx, playlistID); err != nil { return 0, err } - return s.ds.Playlist(ctx).Tracks(playlistID, false).Add(ids) + return s.ds.Playlist().Tracks(ctx, playlistID, false).Add(ctx, ids) } // InsertTracks adds tracks before the 1-based position pos; a position past the end appends. @@ -279,7 +282,7 @@ func (s *playlists) InsertTracks(ctx context.Context, playlistID string, ids []s // concurrent writers instead of waiting for the lock. err := s.ds.WithTxImmediate(func(tx model.DataStore) error { var err error - count, err = tx.Playlist(ctx).Tracks(playlistID, false).Insert(ids, pos) + count, err = tx.Playlist().Tracks(ctx, playlistID, false).Insert(ctx, ids, pos) return err }) return count, err @@ -289,21 +292,21 @@ func (s *playlists) AddAlbums(ctx context.Context, playlistID string, albumIds [ if _, err := s.checkTracksEditable(ctx, playlistID); err != nil { return 0, err } - return s.ds.Playlist(ctx).Tracks(playlistID, false).AddAlbums(albumIds) + return s.ds.Playlist().Tracks(ctx, playlistID, false).AddAlbums(ctx, albumIds) } func (s *playlists) AddArtists(ctx context.Context, playlistID string, artistIds []string) (int, error) { if _, err := s.checkTracksEditable(ctx, playlistID); err != nil { return 0, err } - return s.ds.Playlist(ctx).Tracks(playlistID, false).AddArtists(artistIds) + return s.ds.Playlist().Tracks(ctx, playlistID, false).AddArtists(ctx, artistIds) } func (s *playlists) AddDiscs(ctx context.Context, playlistID string, discs []model.DiscID) (int, error) { if _, err := s.checkTracksEditable(ctx, playlistID); err != nil { return 0, err } - return s.ds.Playlist(ctx).Tracks(playlistID, false).AddDiscs(discs) + return s.ds.Playlist().Tracks(ctx, playlistID, false).AddDiscs(ctx, discs) } func (s *playlists) RemoveTracks(ctx context.Context, playlistID string, trackIds []string) error { @@ -311,7 +314,7 @@ func (s *playlists) RemoveTracks(ctx context.Context, playlistID string, trackId return err } return s.ds.WithTx(func(tx model.DataStore) error { - return tx.Playlist(ctx).Tracks(playlistID, false).Delete(trackIds...) + return tx.Playlist().Tracks(ctx, playlistID, false).Delete(ctx, trackIds...) }) } @@ -320,7 +323,7 @@ func (s *playlists) ReorderTrack(ctx context.Context, playlistID string, pos int return err } return s.ds.WithTxImmediate(func(tx model.DataStore) error { - return tx.Playlist(ctx).Tracks(playlistID, false).Reorder(pos, newPos) + return tx.Playlist().Tracks(ctx, playlistID, false).Reorder(ctx, pos, newPos) }) } @@ -339,7 +342,7 @@ func (s *playlists) SetImage(ctx context.Context, playlistID string, reader io.R } pls.UploadedImage = filename - if err := s.ds.Playlist(ctx).Put(pls); err != nil { + if err := s.ds.Playlist().Put(ctx, pls); err != nil { return err } s.imgUpload.EnqueueArtwork(ctx, consts.EntityPlaylist, pls.ID) @@ -357,7 +360,7 @@ func (s *playlists) RemoveImage(ctx context.Context, playlistID string) error { } pls.UploadedImage = "" - if err := s.ds.Playlist(ctx).Put(pls); err != nil { + if err := s.ds.Playlist().Put(ctx, pls); err != nil { return err } s.imgUpload.EnqueueArtwork(ctx, consts.EntityPlaylist, pls.ID) diff --git a/core/playlists/playlists_test.go b/core/playlists/playlists_test.go index dd6213555..ec44de3bc 100644 --- a/core/playlists/playlists_test.go +++ b/core/playlists/playlists_test.go @@ -493,15 +493,15 @@ var _ = Describe("Playlists", func() { It("clears the resolved artwork state and re-queues after removing an upload", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - Expect(ds.Artwork(ctx).PutItemArtwork(&model.ItemArtwork{ + Expect(ds.Artwork().PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "pl", ItemID: "pls-1", Hash: "oldhash", Source: "upload", })).To(Succeed()) Expect(ps.RemoveImage(ctx, "pls-1")).To(Succeed()) - _, err := ds.Artwork(ctx).GetItemArtwork(model.KindPlaylistArtwork, "pls-1", model.ImageTypePrimary) + _, err := ds.Artwork().GetItemArtwork(ctx, model.KindPlaylistArtwork, "pls-1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) - queued, _ := ds.ArtworkQueue(ctx).DequeueBatch(100) + queued, _ := ds.ArtworkQueue().DequeueBatch(ctx, 100) Expect(queued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "pl"), HaveField("ItemID", "pls-1"), diff --git a/core/playlists/rest_adapter.go b/core/playlists/rest_adapter.go index 45642c819..f8e8a9d41 100644 --- a/core/playlists/rest_adapter.go +++ b/core/playlists/rest_adapter.go @@ -14,42 +14,43 @@ import ( // --- REST adapter (follows Share/Library pattern) --- -func (s *playlists) NewRepository(ctx context.Context) rest.Repository { - return &playlistRepositoryWrapper{ - ctx: ctx, - PlaylistRepository: s.ds.Playlist(ctx), - service: s, - } +func (s *playlists) Repository() rest.Repository[model.Playlist] { + return s.repo } -// playlistRepositoryWrapper wraps the playlist repository as a thin REST-to-service adapter. -// It satisfies rest.Repository through the embedded PlaylistRepository (via ResourceRepository), -// and rest.Persistable by delegating to service methods for all mutations. +// playlistRepositoryWrapper wraps the playlist repository as a thin REST-to-service adapter, +// delegating to service methods for all mutations. type playlistRepositoryWrapper struct { model.PlaylistRepository - ctx context.Context service *playlists } -func (r *playlistRepositoryWrapper) Save(entity any) (string, error) { - return r.service.savePlaylist(r.ctx, entity.(*model.Playlist)) +var _ rest.Persistable[model.Playlist] = (*playlistRepositoryWrapper)(nil) + +func (r *playlistRepositoryWrapper) Save(ctx context.Context, entity *model.Playlist) (string, error) { + return r.service.savePlaylist(ctx, entity) } -func (r *playlistRepositoryWrapper) Update(id string, entity any, cols ...string) error { - return r.service.updatePlaylistEntity(r.ctx, id, entity.(*model.Playlist), cols...) +func (r *playlistRepositoryWrapper) Update(ctx context.Context, id string, entity model.Playlist, cols ...string) error { + return r.service.updatePlaylistEntity(ctx, id, &entity, cols...) } -func (r *playlistRepositoryWrapper) Delete(id string) error { - return r.service.Delete(r.ctx, id) +func (r *playlistRepositoryWrapper) Delete(ctx context.Context, ids ...string) error { + for _, id := range ids { + if err := r.service.Delete(ctx, id); err != nil { + return err + } + } + return nil } -func (s *playlists) TracksRepository(ctx context.Context, playlistId string, refreshSmartPlaylist bool) rest.Repository { - repo := s.ds.Playlist(ctx) - tracks := repo.Tracks(playlistId, refreshSmartPlaylist) +func (s *playlists) TracksRepository(ctx context.Context, playlistId string, refreshSmartPlaylist bool) rest.Repository[model.PlaylistTrack] { + repo := s.ds.Playlist() + tracks := repo.Tracks(ctx, playlistId, refreshSmartPlaylist) if tracks == nil { return nil } - return tracks.(rest.Repository) + return tracks } // savePlaylist creates a new playlist, assigning the owner from context. @@ -63,7 +64,7 @@ func (s *playlists) savePlaylist(ctx context.Context, pls *model.Playlist) (stri pls.UploadedImage = "" // Managed by image upload endpoint pls.ExternalImageURL = "" // Managed by M3U import / plugins only pls.EvaluatedAt = nil // Server-managed - err := s.ds.Playlist(ctx).Put(pls) + err := s.ds.Playlist().Put(ctx, pls) if err != nil { return "", err } @@ -156,7 +157,7 @@ func (s *playlists) applyFlagsOnly(ctx context.Context, current, entity *model.P if len(updateCols) == 0 { return nil } - return s.ds.Playlist(ctx).Put(current, updateCols...) + return s.ds.Playlist().Put(ctx, current, updateCols...) } // sentFields returns a predicate that reports whether a JSON field was present diff --git a/core/playlists/rest_adapter_test.go b/core/playlists/rest_adapter_test.go index fbfab350c..e82dd9df5 100644 --- a/core/playlists/rest_adapter_test.go +++ b/core/playlists/rest_adapter_test.go @@ -31,7 +31,7 @@ var _ = Describe("REST Adapter", func() { }) Describe("NewRepository", func() { - var repo rest.Persistable + var repo rest.Persistable[model.Playlist] BeforeEach(func() { mockPlsRepo.Data = map[string]*model.Playlist{ @@ -43,9 +43,9 @@ var _ = Describe("REST Adapter", func() { Describe("Save", func() { It("sets the owner from the context user", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "New Playlist"} - id, err := repo.Save(pls) + id, err := repo.Save(ctx, pls) Expect(err).ToNot(HaveOccurred()) Expect(id).ToNot(BeEmpty()) Expect(pls.OwnerID).To(Equal("user-1")) @@ -53,16 +53,16 @@ var _ = Describe("REST Adapter", func() { It("forces a new creation by clearing ID", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{ID: "should-be-cleared", Name: "New"} - _, err := repo.Save(pls) + _, err := repo.Save(ctx, pls) Expect(err).ToNot(HaveOccurred()) Expect(pls.ID).ToNot(Equal("should-be-cleared")) }) It("clears server-managed fields to prevent injection via REST API", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{ Name: "Legit Playlist", Comment: "A comment", @@ -74,7 +74,7 @@ var _ = Describe("REST Adapter", func() { ExternalImageURL: "http://evil.example.com/ssrf", EvaluatedAt: new(time.Now()), } - _, err := repo.Save(pls) + _, err := repo.Save(ctx, pls) Expect(err).ToNot(HaveOccurred()) saved := mockPlsRepo.Last @@ -95,33 +95,33 @@ var _ = Describe("REST Adapter", func() { Describe("Update", func() { It("allows owner to update their playlist", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "Updated"} - err := repo.Update("pls-1", pls) + err := repo.Update(ctx, "pls-1", *pls) Expect(err).ToNot(HaveOccurred()) }) It("allows admin to update any playlist", func() { ctx = request.WithUser(ctx, model.User{ID: "admin-1", IsAdmin: true}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "Updated"} - err := repo.Update("pls-1", pls) + err := repo.Update(ctx, "pls-1", *pls) Expect(err).ToNot(HaveOccurred()) }) It("denies non-owner, non-admin", func() { ctx = request.WithUser(ctx, model.User{ID: "other-user", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "Updated"} - err := repo.Update("pls-1", pls) + err := repo.Update(ctx, "pls-1", *pls) Expect(err).To(Equal(rest.ErrPermissionDenied)) }) It("denies regular user from changing ownership", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "Updated", OwnerID: "other-user"} - err := repo.Update("pls-1", pls) + err := repo.Update(ctx, "pls-1", *pls) Expect(err).To(Equal(rest.ErrPermissionDenied)) }) @@ -133,9 +133,9 @@ var _ = Describe("REST Adapter", func() { // entity.OwnerID. sentFields normalizes both sides so the // permission gate fires regardless of casing. ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{OwnerID: "other-user"} - err := repo.Update("pls-1", pls, colName) + err := repo.Update(ctx, "pls-1", *pls, colName) Expect(err).To(Equal(rest.ErrPermissionDenied)) }, Entry("canonical camelCase", "ownerId"), @@ -152,10 +152,10 @@ var _ = Describe("REST Adapter", func() { Rules: &criteria.Criteria{Expression: criteria.Contains{"title": "old"}}, } ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) newRules := &criteria.Criteria{Expression: criteria.Contains{"title": "new"}} pls := &model.Playlist{Name: "Smart Playlist", Rules: newRules} - err := repo.Update("smart-1", pls) + err := repo.Update(ctx, "smart-1", *pls) Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Rules).To(Equal(newRules)) }) @@ -171,10 +171,10 @@ var _ = Describe("REST Adapter", func() { Rules: &criteria.Criteria{Expression: criteria.Contains{"title": "old"}}, } ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) newRules := &criteria.Criteria{Expression: criteria.Contains{"title": "new"}} pls := &model.Playlist{Rules: newRules} - err := repo.Update("smart-1", pls, "rules") + err := repo.Update(ctx, "smart-1", *pls, "rules") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.ImportedHash).To(BeEmpty()) }) @@ -190,9 +190,9 @@ var _ = Describe("REST Adapter", func() { UpdatedAt: originalTime, } ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "File Playlist", Sync: false} - err := repo.Update("file-pls", pls) + err := repo.Update(ctx, "file-pls", *pls) Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Sync).To(BeFalse()) Expect(mockPlsRepo.Last.UpdatedAt).To(Equal(originalTime)) @@ -207,9 +207,9 @@ var _ = Describe("REST Adapter", func() { Sync: false, } ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "Manual Playlist", Sync: true} - err := repo.Update("manual-pls", pls) + err := repo.Update(ctx, "manual-pls", *pls) Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last).To(BeNil()) }) @@ -224,9 +224,9 @@ var _ = Describe("REST Adapter", func() { UpdatedAt: originalTime, } ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "My Playlist", Public: true} - err := repo.Update("pls-pub", pls) + err := repo.Update(ctx, "pls-pub", *pls) Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Public).To(BeTrue()) Expect(mockPlsRepo.Last.UpdatedAt).To(Equal(originalTime)) @@ -241,9 +241,9 @@ var _ = Describe("REST Adapter", func() { Sync: true, } ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "New Name", Sync: false} - err := repo.Update("file-pls2", pls) + err := repo.Update(ctx, "file-pls2", *pls) Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Name).To(Equal("New Name")) Expect(mockPlsRepo.Last.Sync).To(BeFalse()) @@ -251,9 +251,9 @@ var _ = Describe("REST Adapter", func() { It("returns rest.ErrNotFound when playlist doesn't exist", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) pls := &model.Playlist{Name: "Updated"} - err := repo.Update("nonexistent", pls) + err := repo.Update(ctx, "nonexistent", *pls) Expect(err).To(Equal(rest.ErrNotFound)) }) @@ -274,8 +274,8 @@ var _ = Describe("REST Adapter", func() { }) It("preserves name and comment when only public is sent (bulk Make Public)", func() { - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("partial", &model.Playlist{Public: true}, "public") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "partial", model.Playlist{Public: true}, "public") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Name).To(Equal("Original Name")) Expect(mockPlsRepo.Last.Comment).To(Equal("Original comment")) @@ -290,16 +290,16 @@ var _ = Describe("REST Adapter", func() { Path: "/music/p.m3u", Sync: true, } - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("file-partial", &model.Playlist{Sync: false}, "sync") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "file-partial", model.Playlist{Sync: false}, "sync") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Name).To(Equal("Keep Me")) Expect(mockPlsRepo.Last.Sync).To(BeFalse()) }) It("renames the playlist when only name is sent", func() { - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("partial", &model.Playlist{Name: "Renamed"}, "name") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "partial", model.Playlist{Name: "Renamed"}, "name") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Name).To(Equal("Renamed")) Expect(mockPlsRepo.Last.Comment).To(Equal("Original comment")) @@ -307,8 +307,8 @@ var _ = Describe("REST Adapter", func() { }) It("clears the comment when an empty comment is sent explicitly", func() { - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("partial", &model.Playlist{Comment: ""}, "comment") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "partial", model.Playlist{Comment: ""}, "comment") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Comment).To(BeEmpty()) Expect(mockPlsRepo.Last.Name).To(Equal("Original Name")) @@ -323,9 +323,9 @@ var _ = Describe("REST Adapter", func() { Public: true, Rules: &criteria.Criteria{Expression: criteria.Is{"genre": "Rock"}}, } - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) newRules := &criteria.Criteria{Expression: criteria.Is{"genre": "Jazz"}, Sort: "year DESC"} - err := repo.Update("smart-partial", &model.Playlist{Rules: newRules}, "rules") + err := repo.Update(ctx, "smart-partial", model.Playlist{Rules: newRules}, "rules") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Rules).To(Equal(newRules)) Expect(mockPlsRepo.Last.Name).To(Equal("Smart Original")) @@ -342,9 +342,9 @@ var _ = Describe("REST Adapter", func() { Rules: &criteria.Criteria{Expression: criteria.Is{"genre": "Rock"}}, EvaluatedAt: &evaluatedAt, } - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) newRules := &criteria.Criteria{Expression: criteria.Is{"genre": "Jazz"}} - err := repo.Update("smart-reset", &model.Playlist{Rules: newRules}, "rules") + err := repo.Update(ctx, "smart-reset", model.Playlist{Rules: newRules}, "rules") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.EvaluatedAt).To(BeNil()) }) @@ -358,8 +358,8 @@ var _ = Describe("REST Adapter", func() { Rules: &criteria.Criteria{Expression: criteria.Is{"genre": "Rock"}}, EvaluatedAt: &evaluatedAt, } - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("smart-keep", &model.Playlist{Name: "Renamed Smart"}, "name") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "smart-keep", model.Playlist{Name: "Renamed Smart"}, "name") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.EvaluatedAt).ToNot(BeNil()) Expect(*mockPlsRepo.Last.EvaluatedAt).To(BeTemporally("~", evaluatedAt, time.Second)) @@ -373,10 +373,10 @@ var _ = Describe("REST Adapter", func() { OwnerID: "user-1", Rules: &criteria.Criteria{Expression: criteria.Is{"genre": "Rock"}}, } - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) newRules := &criteria.Criteria{Expression: criteria.Is{"artist": "Miles Davis"}, Sort: "album"} - err := repo.Update("smart-edit", - &model.Playlist{Name: "Smart Renamed", Rules: newRules}, + err := repo.Update(ctx, "smart-edit", + model.Playlist{Name: "Smart Renamed", Rules: newRules}, "name", "rules") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Name).To(Equal("Smart Renamed")) @@ -392,11 +392,11 @@ var _ = Describe("REST Adapter", func() { OwnerID: "user-1", Rules: rules, } - repo = ps.NewRepository(ctx).(rest.Persistable) + repo = ps.Repository().(rest.Persistable[model.Playlist]) // Same rules sent back — rulesEqual should report no change and // the request should no-op (no Put call). sameRules := &criteria.Criteria{Expression: criteria.Is{"genre": "Rock"}} - err := repo.Update("smart-idempotent", &model.Playlist{Rules: sameRules}, "rules") + err := repo.Update(ctx, "smart-idempotent", model.Playlist{Rules: sameRules}, "rules") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last).To(BeNil()) // no Put happened }) @@ -410,8 +410,8 @@ var _ = Describe("REST Adapter", func() { Public: false, Rules: rules, } - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("smart-public", &model.Playlist{Public: true}, "public") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "smart-public", model.Playlist{Public: true}, "public") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Public).To(BeTrue()) Expect(mockPlsRepo.Last.Rules).To(Equal(rules)) @@ -421,8 +421,8 @@ var _ = Describe("REST Adapter", func() { It("does not treat a missing ownerId as an ownership transfer attempt", func() { // A non-admin user sending only {public:true} should not be blocked // just because OwnerID is the zero value in the deserialized entity. - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("partial", &model.Playlist{Public: true}, "public") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "partial", model.Playlist{Public: true}, "public") Expect(err).ToNot(HaveOccurred()) }) @@ -431,8 +431,8 @@ var _ = Describe("REST Adapter", func() { // like {"Name":"x"}, but rest.Put's field-name extraction is // case-sensitive. sentFields normalizes both sides so a request // with {"Name":"Renamed"} is honored, not silently ignored. - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Update("partial", &model.Playlist{Name: "Renamed"}, "Name") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Update(ctx, "partial", model.Playlist{Name: "Renamed"}, "Name") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Last.Name).To(Equal("Renamed")) Expect(mockPlsRepo.Last.Comment).To(Equal("Original comment")) @@ -443,16 +443,16 @@ var _ = Describe("REST Adapter", func() { Describe("Delete", func() { It("delegates to service Delete with permission checks", func() { ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Delete("pls-1") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Delete(ctx, "pls-1") Expect(err).ToNot(HaveOccurred()) Expect(mockPlsRepo.Deleted).To(ContainElement("pls-1")) }) It("denies non-owner", func() { ctx = request.WithUser(ctx, model.User{ID: "other-user", IsAdmin: false}) - repo = ps.NewRepository(ctx).(rest.Persistable) - err := repo.Delete("pls-1") + repo = ps.Repository().(rest.Persistable[model.Playlist]) + err := repo.Delete(ctx, "pls-1") Expect(err).To(Equal(rest.ErrPermissionDenied)) }) }) diff --git a/core/scrobbler/buffered_scrobbler.go b/core/scrobbler/buffered_scrobbler.go index c5c096a0b..30bced623 100644 --- a/core/scrobbler/buffered_scrobbler.go +++ b/core/scrobbler/buffered_scrobbler.go @@ -70,7 +70,7 @@ type bufferedScrobbler struct { loader Loader service string wakeSignal chan struct{} - ctx context.Context + ctx context.Context //nolint:containedctx // scrobbler lifecycle ctx, cancelled by Stop cancel context.CancelFunc } @@ -97,7 +97,7 @@ func (b *bufferedScrobbler) NowPlaying(ctx context.Context, userId string, track } func (b *bufferedScrobbler) Scrobble(ctx context.Context, userId string, s Scrobble) error { - err := b.ds.ScrobbleBuffer(ctx).Enqueue(b.service, userId, s.ID, s.TimeStamp) + err := b.ds.ScrobbleBuffer().Enqueue(ctx, b.service, userId, s.ID, s.TimeStamp) if err != nil { return err } @@ -154,8 +154,8 @@ func (b *bufferedScrobbler) run(ctx context.Context) { } func (b *bufferedScrobbler) processQueue(ctx context.Context) (bool, time.Duration) { - buffer := b.ds.ScrobbleBuffer(ctx) - userIds, err := buffer.UserIDs(b.service) + buffer := b.ds.ScrobbleBuffer() + userIds, err := buffer.UserIDs(ctx, b.service) if err != nil { log.Error(ctx, "Error retrieving userIds from scrobble buffer", "scrobbler", b.service, err) return false, 0 @@ -176,14 +176,14 @@ func (b *bufferedScrobbler) processUserQueue(ctx context.Context, userId string) // Scrobbles are drained on a background context that no longer carries the // request's authenticated user. Restore it from the buffered userId so that // scrobblers relying on the user in the context (e.g. plugins) still get it. - if user, err := b.ds.User(ctx).Get(userId); err != nil { + if user, err := b.ds.User().Get(ctx, userId); err != nil { log.Warn(ctx, "Could not load user for buffered scrobble", "userId", userId, "scrobbler", b.service, err) } else { ctx = request.WithUser(ctx, *user) } - buffer := b.ds.ScrobbleBuffer(ctx) + buffer := b.ds.ScrobbleBuffer() for { - entry, err := buffer.Next(b.service, userId) + entry, err := buffer.Next(ctx, b.service, userId) if err != nil { log.Error(ctx, "Error reading from scrobble buffer", "scrobbler", b.service, err) return false, 0 @@ -210,7 +210,7 @@ func (b *bufferedScrobbler) processUserQueue(ctx context.Context, userId string) log.Error(ctx, "Error sending scrobble to service. Discarding", "scrobbler", b.service, "userId", entry.UserID, "artist", entry.Artist, "track", entry.Title, err) } - err = buffer.Dequeue(entry) + err = buffer.Dequeue(ctx, entry) if err != nil { log.Error(ctx, "Error removing entry from scrobble buffer", "userId", entry.UserID, "track", entry.Title, "artist", entry.Artist, "scrobbler", b.service, err) diff --git a/core/scrobbler/buffered_scrobbler_test.go b/core/scrobbler/buffered_scrobbler_test.go index 16172194b..6ecd92cab 100644 --- a/core/scrobbler/buffered_scrobbler_test.go +++ b/core/scrobbler/buffered_scrobbler_test.go @@ -26,7 +26,7 @@ var _ = Describe("BufferedScrobbler", func() { ctx = context.Background() buffer = tests.CreateMockedScrobbleBufferRepo() userRepo := tests.CreateMockUserRepo() - Expect(userRepo.Put(&model.User{ID: "user1", UserName: "alice"})).To(Succeed()) + Expect(userRepo.Put(ctx, &model.User{ID: "user1", UserName: "alice"})).To(Succeed()) ds = &tests.MockDataStore{ MockedScrobbleBuffer: buffer, MockedUser: userRepo, @@ -55,7 +55,7 @@ var _ = Describe("BufferedScrobbler", func() { track := model.MediaFile{ID: "123", Title: "Test Track"} now := time.Now() scrobble := Scrobble{MediaFile: track, TimeStamp: now} - Expect(buffer.Length()).To(Equal(int64(0))) + Expect(buffer.Length(ctx)).To(Equal(int64(0))) Expect(scr.ScrobbleCalled.Load()).To(BeFalse()) Expect(bs.Scrobble(ctx, "user1", scrobble)).To(Succeed()) @@ -131,7 +131,7 @@ func TestBufferedScrobblerBackoffSchedule(t *testing.T) { g := NewWithT(t) buffer := tests.CreateMockedScrobbleBufferRepo() userRepo := tests.CreateMockUserRepo() - g.Expect(userRepo.Put(&model.User{ID: "user1", UserName: "alice"})).To(Succeed()) + g.Expect(userRepo.Put(t.Context(), &model.User{ID: "user1", UserName: "alice"})).To(Succeed()) ds := &tests.MockDataStore{MockedScrobbleBuffer: buffer, MockedUser: userRepo} flaky := &recoveringScrobbler{} @@ -147,7 +147,7 @@ func TestBufferedScrobblerBackoffSchedule(t *testing.T) { // First attempt fires immediately on the enqueue wake and is left buffered. synctest.Wait() g.Expect(flaky.count.Load()).To(Equal(int32(1))) - g.Expect(buffer.Length()).To(Equal(int64(1))) + g.Expect(buffer.Length(t.Context())).To(Equal(int64(1))) // Each subsequent retry waits exactly double the previous: 5s, 10s, 20s, 40s. for i, gap := range []time.Duration{5 * time.Second, 10 * time.Second, 20 * time.Second, 40 * time.Second} { @@ -165,10 +165,10 @@ func TestBufferedScrobblerBackoffSchedule(t *testing.T) { flaky.succeed() bs.sendWakeSignal() synctest.Wait() - g.Expect(buffer.Length()).To(Equal(int64(1)), "wake during backoff drained early") + g.Expect(buffer.Length(t.Context())).To(Equal(int64(1)), "wake during backoff drained early") time.Sleep(80 * time.Second) synctest.Wait() - g.Expect(buffer.Length()).To(Equal(int64(0))) + g.Expect(buffer.Length(t.Context())).To(Equal(int64(0))) }) } @@ -176,7 +176,7 @@ func TestBufferedScrobblerBackoffWindow(t *testing.T) { synctest.Test(t, func(t *testing.T) { buffer := tests.CreateMockedScrobbleBufferRepo() userRepo := tests.CreateMockUserRepo() - _ = userRepo.Put(&model.User{ID: "user1", UserName: "alice"}) + _ = userRepo.Put(t.Context(), &model.User{ID: "user1", UserName: "alice"}) ds := &tests.MockDataStore{MockedScrobbleBuffer: buffer, MockedUser: userRepo} scr := &fakeScrobbler{Authorized: true} scr.SetError(errors.Join(errors.New("boom"), ErrRetryLater)) @@ -211,7 +211,7 @@ func TestBufferedScrobblerHonorsServerDelay(t *testing.T) { synctest.Test(t, func(t *testing.T) { buffer := tests.CreateMockedScrobbleBufferRepo() userRepo := tests.CreateMockUserRepo() - _ = userRepo.Put(&model.User{ID: "user1", UserName: "alice"}) + _ = userRepo.Put(t.Context(), &model.User{ID: "user1", UserName: "alice"}) ds := &tests.MockDataStore{MockedScrobbleBuffer: buffer, MockedUser: userRepo} scr := &fakeScrobbler{Authorized: true} scr.SetError(errors.Join(errors.New("429"), &agents.RetryLaterError{RetryIn: 30 * time.Second})) @@ -244,8 +244,8 @@ func TestBufferedScrobblerTakesTheLongestServerDelayAcrossUsers(t *testing.T) { synctest.Test(t, func(t *testing.T) { buffer := tests.CreateMockedScrobbleBufferRepo() userRepo := tests.CreateMockUserRepo() - _ = userRepo.Put(&model.User{ID: "user1", UserName: "alice"}) - _ = userRepo.Put(&model.User{ID: "user2", UserName: "bob"}) + _ = userRepo.Put(t.Context(), &model.User{ID: "user1", UserName: "alice"}) + _ = userRepo.Put(t.Context(), &model.User{ID: "user2", UserName: "bob"}) ds := &tests.MockDataStore{MockedScrobbleBuffer: buffer, MockedUser: userRepo} scr := &recoveringScrobbler{delays: map[string]time.Duration{ "user1": 10 * time.Second, @@ -253,8 +253,8 @@ func TestBufferedScrobblerTakesTheLongestServerDelayAcrossUsers(t *testing.T) { }} // Both are buffered before the drain goroutine exists: it drains once on startup, and // seeing only one user there would park it on that user's delay, ignoring the other. - _ = buffer.Enqueue("test", "user1", "1", time.Now()) - _ = buffer.Enqueue("test", "user2", "2", time.Now()) + _ = buffer.Enqueue(t.Context(), "test", "user1", "1", time.Now()) + _ = buffer.Enqueue(t.Context(), "test", "user2", "2", time.Now()) bs := newBufferedScrobbler(ds, scr, "test") defer bs.Stop() diff --git a/core/scrobbler/play_tracker.go b/core/scrobbler/play_tracker.go index e68fcd942..541a8e92b 100644 --- a/core/scrobbler/play_tracker.go +++ b/core/scrobbler/play_tracker.go @@ -68,14 +68,14 @@ type ReportPlaybackParams struct { } type nowPlayingEntry struct { - ctx context.Context + ctx context.Context //nolint:containedctx // queued work item carries the request ctx to the worker userId string track *model.MediaFile position int } type playbackReportEntry struct { - ctx context.Context + ctx context.Context //nolint:containedctx // queued work item carries the request ctx to the worker info PlaybackSession filtered bool } @@ -293,7 +293,7 @@ func (p *playTracker) ReportPlayback(ctx context.Context, params ReportPlaybackP log.Trace(ctx, "Ignoring out-of-order starting report for playing session", "clientId", clientId, "mediaId", params.MediaId) return nil } - mf, err := p.ds.MediaFile(ctx).GetWithParticipants(params.MediaId) + mf, err := p.ds.MediaFile().GetWithParticipants(ctx, params.MediaId) if err != nil { return err } @@ -328,7 +328,7 @@ func (p *playTracker) ReportPlayback(ctx context.Context, params ReportPlaybackP case StatePlaying, StatePaused: info, getErr := p.playMap.Get(clientId) if getErr != nil || info.MediaFile.ID != params.MediaId { - mf, err := p.ds.MediaFile(ctx).GetWithParticipants(params.MediaId) + mf, err := p.ds.MediaFile().GetWithParticipants(ctx, params.MediaId) if err != nil { return err } @@ -364,7 +364,7 @@ func (p *playTracker) ReportPlayback(ctx context.Context, params ReportPlaybackP var loadedMF *model.MediaFile haveVerdict := false if !params.IgnoreScrobble && player.ScrobbleEnabled { - mf, err := p.ds.MediaFile(ctx).GetWithParticipants(params.MediaId) + mf, err := p.ds.MediaFile().GetWithParticipants(ctx, params.MediaId) if err != nil { return err } @@ -409,7 +409,7 @@ func (p *playTracker) ReportPlayback(ctx context.Context, params ReportPlaybackP mf := loadedMF if mf == nil { var mfErr error - mf, mfErr = p.ds.MediaFile(ctx).GetWithParticipants(params.MediaId) + mf, mfErr = p.ds.MediaFile().GetWithParticipants(ctx, params.MediaId) if mfErr != nil { return mfErr } @@ -477,7 +477,7 @@ func (p *playTracker) Submit(ctx context.Context, submissions []Submission) erro success := 0 for _, s := range submissions { - mf, err := p.ds.MediaFile(ctx).GetWithParticipants(s.TrackID) + mf, err := p.ds.MediaFile().GetWithParticipants(ctx, s.TrackID) if err != nil { log.Error(ctx, "Cannot find track for scrobbling", "id", s.TrackID, "user", username, err) continue @@ -504,22 +504,22 @@ func (p *playTracker) Submit(ctx context.Context, submissions []Submission) erro func (p *playTracker) incPlay(ctx context.Context, track *model.MediaFile, timestamp time.Time) error { return p.ds.WithTx(func(tx model.DataStore) error { - err := tx.MediaFile(ctx).IncPlayCount(track.ID, timestamp) + err := tx.MediaFile().IncPlayCount(ctx, track.ID, timestamp) if err != nil { return err } - err = tx.Album(ctx).IncPlayCount(track.AlbumID, timestamp) + err = tx.Album().IncPlayCount(ctx, track.AlbumID, timestamp) if err != nil { return err } for _, artist := range track.Participants[model.RoleArtist] { - err = tx.Artist(ctx).IncPlayCount(artist.ID, timestamp) + err = tx.Artist().IncPlayCount(ctx, artist.ID, timestamp) if err != nil { return err } } if conf.Server.EnableScrobbleHistory { - return tx.Scrobble(ctx).RecordScrobble(track.ID, timestamp) + return tx.Scrobble().RecordScrobble(ctx, track.ID, timestamp) } return nil }) @@ -538,7 +538,7 @@ func (p *playTracker) isFilteredOut(ctx context.Context, t *model.MediaFile) boo log.Warn(ctx, "Invalid scrobble filter, ignoring", "user", u.UserName, err) return false } - match, err := p.ds.MediaFile(ctx).MatchesCriteria(t.ID, c) + match, err := p.ds.MediaFile().MatchesCriteria(ctx, t.ID, c) if err != nil { log.Warn(ctx, "Error evaluating scrobble filter, ignoring", "user", u.UserName, "track", t.Title, err) return false diff --git a/core/scrobbler/play_tracker_test.go b/core/scrobbler/play_tracker_test.go index d79233e15..045827811 100644 --- a/core/scrobbler/play_tracker_test.go +++ b/core/scrobbler/play_tracker_test.go @@ -55,12 +55,12 @@ type flipOnPlayRepo struct { played atomic.Bool } -func (r *flipOnPlayRepo) IncPlayCount(id string, ts time.Time) error { +func (r *flipOnPlayRepo) IncPlayCount(ctx context.Context, id string, ts time.Time) error { r.played.Store(true) - return r.MediaFileRepository.IncPlayCount(id, ts) + return r.MediaFileRepository.IncPlayCount(ctx, id, ts) } -func (r *flipOnPlayRepo) MatchesCriteria(string, criteria.Criteria) (bool, error) { +func (r *flipOnPlayRepo) MatchesCriteria(context.Context, string, criteria.Criteria) (bool, error) { if r.played.Load() { return r.after, nil } @@ -73,9 +73,9 @@ type slowMediaFileRepo struct { model.MediaFileRepository } -func (s *slowMediaFileRepo) GetWithParticipants(id string) (*model.MediaFile, error) { +func (s *slowMediaFileRepo) GetWithParticipants(ctx context.Context, id string) (*model.MediaFile, error) { time.Sleep(5 * time.Millisecond) - return s.MediaFileRepository.GetWithParticipants(id) + return s.MediaFileRepository.GetWithParticipants(ctx, id) } var _ = Describe("PlayTracker", func() { @@ -119,13 +119,13 @@ var _ = Describe("PlayTracker", func() { model.RoleArtist: []model.Participant{_p("ar-1", "Artist 1"), _p("ar-2", "Artist 2")}, }, } - _ = ds.MediaFile(ctx).Put(&track) + _ = ds.MediaFile().Put(ctx, &track) artist1 = model.Artist{ID: "ar-1"} - _ = ds.Artist(ctx).Put(&artist1) + _ = ds.Artist().Put(ctx, &artist1) artist2 = model.Artist{ID: "ar-2"} - _ = ds.Artist(ctx).Put(&artist2) + _ = ds.Artist().Put(ctx, &artist2) album = model.Album{ID: "al-1"} - _ = ds.Album(ctx).(*tests.MockAlbumRepo).Put(&album) + _ = ds.Album().(*tests.MockAlbumRepo).Put(ctx, &album) }) AfterEach(func() { @@ -149,7 +149,7 @@ var _ = Describe("PlayTracker", func() { It("returns current playing music", func() { track2 := track track2.ID = "456" - _ = ds.MediaFile(ctx).Put(&track2) + _ = ds.MediaFile().Put(ctx, &track2) ctx1 := request.WithUser(GinkgoT().Context(), model.User{UserName: "user-1"}) ctx1 = request.WithPlayer(ctx1, model.Player{ScrobbleEnabled: true}) _ = tracker.ReportPlayback(ctx1, ReportPlaybackParams{ @@ -180,7 +180,7 @@ var _ = Describe("PlayTracker", func() { hidden := track hidden.ID = "789" hidden.LibraryID = 2 - _ = ds.MediaFile(ctx).Put(&hidden) + _ = ds.MediaFile().Put(ctx, &hidden) reporter := request.WithPlayer( request.WithUser(GinkgoT().Context(), model.User{ID: "u-2", UserName: "user-2"}), model.Player{ScrobbleEnabled: true}, @@ -199,7 +199,7 @@ var _ = Describe("PlayTracker", func() { hidden := track hidden.ID = "789" hidden.LibraryID = 2 - _ = ds.MediaFile(ctx).Put(&hidden) + _ = ds.MediaFile().Put(ctx, &hidden) reporter := request.WithPlayer( request.WithUser(GinkgoT().Context(), model.User{ID: "u-2", UserName: "user-2"}), model.Player{ScrobbleEnabled: true}, @@ -360,7 +360,7 @@ var _ = Describe("PlayTracker", func() { Expect(err).ToNot(HaveOccurred()) mockDS := ds.(*tests.MockDataStore) - mockScrobble := mockDS.Scrobble(ctx).(*tests.MockScrobbleRepo) + mockScrobble := mockDS.Scrobble().(*tests.MockScrobbleRepo) Expect(mockScrobble.RecordedScrobbles).To(HaveLen(1)) Expect(mockScrobble.RecordedScrobbles[0].MediaFileID).To(Equal("123")) Expect(mockScrobble.RecordedScrobbles[0].UserID).To(Equal("u-1")) @@ -376,7 +376,7 @@ var _ = Describe("PlayTracker", func() { Expect(err).ToNot(HaveOccurred()) mockDS := ds.(*tests.MockDataStore) - mockScrobble := mockDS.Scrobble(ctx).(*tests.MockScrobbleRepo) + mockScrobble := mockDS.Scrobble().(*tests.MockScrobbleRepo) Expect(mockScrobble.RecordedScrobbles).To(HaveLen(0)) }) }) @@ -388,7 +388,7 @@ var _ = Describe("PlayTracker", func() { BeforeEach(func() { ctx = request.WithUser(ctx, model.User{ID: "u-1", UserName: "user-1", ScrobbleFilter: `{"all":[{"contains":{"title":"Track"}}]}`}) - repo = ds.MediaFile(ctx).(*tests.MockMediaFileRepo) + repo = ds.MediaFile().(*tests.MockMediaFileRepo) }) It("does not send a matching track to the agent", func() { @@ -489,7 +489,7 @@ var _ = Describe("PlayTracker", func() { var flip *flipOnPlayRepo install := func(before, after bool) { - flip = &flipOnPlayRepo{MediaFileRepository: ds.MediaFile(ctx), before: before, after: after} + flip = &flipOnPlayRepo{MediaFileRepository: ds.MediaFile(), before: before, after: after} ds.(*tests.MockDataStore).MockedMediaFile = flip } @@ -632,7 +632,7 @@ var _ = Describe("PlayTracker", func() { It("starting replaces existing entry when switching tracks on same player", func() { track2 := track track2.ID = "456" - _ = ds.MediaFile(ctx).Put(&track2) + _ = ds.MediaFile().Put(ctx, &track2) err := tracker.ReportPlayback(ctx, ReportPlaybackParams{ MediaId: "123", PositionMs: 50000, State: "playing", PlaybackRate: 1.0, ClientId: defaultClientId, @@ -659,7 +659,7 @@ var _ = Describe("PlayTracker", func() { track2 := track track2.ID = "456" - _ = ds.MediaFile(ctx).Put(&track2) + _ = ds.MediaFile().Put(ctx, &track2) err := tracker.ReportPlayback(ctx1, ReportPlaybackParams{ MediaId: "123", PositionMs: 0, State: "playing", PlaybackRate: 1.0, ClientId: "client-1", @@ -760,7 +760,7 @@ var _ = Describe("PlayTracker", func() { model.RoleArtist: []model.Participant{_p("ar-1", "Artist 1")}, }, } - _ = ds.MediaFile(ctx).Put(&longTrack) + _ = ds.MediaFile().Put(ctx, &longTrack) err := tracker.ReportPlayback(ctx, ReportPlaybackParams{ MediaId: "long", PositionMs: 0, State: "starting", PlaybackRate: 1.0, ClientId: defaultClientId, @@ -958,7 +958,7 @@ var _ = Describe("PlayTracker", func() { BeforeEach(func() { track2 := track track2.ID = "456" - _ = ds.MediaFile(ctx).Put(&track2) + _ = ds.MediaFile().Put(ctx, &track2) }) It("does not downgrade an actively playing session when a late starting report arrives for the same track", func() { @@ -1023,7 +1023,7 @@ var _ = Describe("PlayTracker", func() { }) It("never lets a concurrent starting report downgrade the playing session", func() { - ds.(*tests.MockDataStore).MockedMediaFile = &slowMediaFileRepo{MediaFileRepository: ds.MediaFile(ctx)} + ds.(*tests.MockDataStore).MockedMediaFile = &slowMediaFileRepo{MediaFileRepository: ds.MediaFile()} for i := range 20 { raceClientId := fmt.Sprintf("race-client-%d", i) var wg sync.WaitGroup diff --git a/core/share.go b/core/share.go index a291a060c..33c8996ef 100644 --- a/core/share.go +++ b/core/share.go @@ -8,7 +8,6 @@ import ( "time" "github.com/Masterminds/squirrel" - "github.com/deluan/rest" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" @@ -21,22 +20,24 @@ import ( type Share interface { Load(ctx context.Context, id string) (*model.Share, error) - NewRepository(ctx context.Context) rest.Repository + Repository() model.ShareRepository } func NewShare(ds model.DataStore) Share { return &shareService{ - ds: ds, + ds: ds, + repo: &shareRepositoryWrapper{ShareRepository: ds.Share(), ds: ds}, } } type shareService struct { - ds model.DataStore + ds model.DataStore + repo *shareRepositoryWrapper } func (s *shareService) Load(ctx context.Context, id string) (*model.Share, error) { - repo := s.ds.Share(ctx) - share, err := repo.Get(id) + repo := s.ds.Share() + share, err := repo.Get(ctx, id) if err != nil { return nil, err } @@ -47,40 +48,29 @@ func (s *shareService) Load(ctx context.Context, id string) (*model.Share, error share.LastVisitedAt = new(time.Now()) share.VisitCount++ - err = repo.(rest.Persistable).Update(id, share, "last_visited_at", "visit_count") + err = repo.Update(ctx, id, *share, "last_visited_at", "visit_count") if err != nil { log.Warn(ctx, "Could not increment visit count for share", "share", share.ID) } return share, nil } -func (s *shareService) NewRepository(ctx context.Context) rest.Repository { - repo := s.ds.Share(ctx) - wrapper := &shareRepositoryWrapper{ - ctx: ctx, - ShareRepository: repo, - Repository: repo.(rest.Repository), - Persistable: repo.(rest.Persistable), - ds: s.ds, - } - return wrapper +func (s *shareService) Repository() model.ShareRepository { + return s.repo } type shareRepositoryWrapper struct { model.ShareRepository - rest.Repository - rest.Persistable - ctx context.Context - ds model.DataStore + ds model.DataStore } -func (r *shareRepositoryWrapper) newId() (string, error) { +func (r *shareRepositoryWrapper) newId(ctx context.Context) (string, error) { for { id, err := nanoid.Generate("0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz", 10) if err != nil { return "", err } - exists, err := r.Exists(id) + exists, err := r.Exists(ctx, id) if err != nil { return "", err } @@ -90,14 +80,13 @@ func (r *shareRepositoryWrapper) newId() (string, error) { } } -func (r *shareRepositoryWrapper) Save(entity any) (string, error) { - s := entity.(*model.Share) +func (r *shareRepositoryWrapper) Save(ctx context.Context, s *model.Share) (string, error) { // Owner is always the caller; never trust a client-supplied UserID, as it // determines the library-access context used to resolve the share contents. - if user, ok := request.UserFrom(r.ctx); ok { + if user, ok := request.UserFrom(ctx); ok { s.UserID = user.ID } - id, err := r.newId() + id, err := r.newId(ctx) if err != nil { return "", err } @@ -106,39 +95,39 @@ func (r *shareRepositoryWrapper) Save(entity any) (string, error) { s.ExpiresAt = new(time.Now().Add(conf.Server.DefaultShareExpiration)) } - s.ResourceType, err = r.resourceType(s.ResourceIDs) + s.ResourceType, err = r.resourceType(ctx, s.ResourceIDs) if err != nil { return "", err } switch s.ResourceType { case "artist": - s.Contents = r.contentsLabelFromArtist(s.ID, s.ResourceIDs) + s.Contents = r.contentsLabelFromArtist(ctx, s.ID, s.ResourceIDs) case "album": - s.Contents = r.contentsLabelFromAlbums(s.ID, s.ResourceIDs) + s.Contents = r.contentsLabelFromAlbums(ctx, s.ID, s.ResourceIDs) case "playlist": - s.Contents = r.contentsLabelFromPlaylist(s.ID, s.ResourceIDs) + s.Contents = r.contentsLabelFromPlaylist(ctx, s.ID, s.ResourceIDs) case "media_file": - s.Contents = r.contentsLabelFromMediaFiles(s.ID, s.ResourceIDs) + s.Contents = r.contentsLabelFromMediaFiles(ctx, s.ID, s.ResourceIDs) } s.Contents = str.TruncateRunes(s.Contents, 30, "...") - return r.Persistable.Save(s) + return r.ShareRepository.Save(ctx, s) } var shareableKinds = []model.Kind{model.KindArtistArtwork, model.KindAlbumArtwork, model.KindPlaylistArtwork, model.KindMediaFileArtwork} // resourceType resolves every ID as the current user, so an entity they cannot see cannot // ride along behind a valid first one, and requires all IDs to be of the same kind. -func (r *shareRepositoryWrapper) resourceType(resourceIDs string) (string, error) { +func (r *shareRepositoryWrapper) resourceType(ctx context.Context, resourceIDs string) (string, error) { resourceType := "" for _, id := range strings.Split(resourceIDs, ",") { - kind, err := model.GetEntityKindByID(r.ctx, r.ds, id) + kind, err := model.GetEntityKindByID(ctx, r.ds, id) if err != nil { return "", err } if !slices.Contains(shareableKinds, kind) { - log.Error(r.ctx, "Invalid Resource ID", "id", id) + log.Error(ctx, "Invalid Resource ID", "id", id) return "", model.ErrNotFound } if resourceType != "" && kind.String() != resourceType { @@ -149,53 +138,53 @@ func (r *shareRepositoryWrapper) resourceType(resourceIDs string) (string, error return resourceType, nil } -func (r *shareRepositoryWrapper) Update(id string, entity any, _ ...string) error { +func (r *shareRepositoryWrapper) Update(ctx context.Context, id string, entity model.Share, _ ...string) error { cols := []string{"description", "downloadable"} // TODO Better handling of Share expiration - if !V(entity.(*model.Share).ExpiresAt).IsZero() { + if !V(entity.ExpiresAt).IsZero() { cols = append(cols, "expires_at") } - return r.Persistable.Update(id, entity, cols...) + return r.ShareRepository.Update(ctx, id, entity, cols...) } -func (r *shareRepositoryWrapper) contentsLabelFromArtist(shareID string, ids string) string { +func (r *shareRepositoryWrapper) contentsLabelFromArtist(ctx context.Context, shareID string, ids string) string { idList := strings.SplitN(ids, ",", 2) - a, err := r.ds.Artist(r.ctx).Get(idList[0]) + a, err := r.ds.Artist().Get(ctx, idList[0]) if err != nil { - log.Error(r.ctx, "Error retrieving artist name for share", "share", shareID, err) + log.Error(ctx, "Error retrieving artist name for share", "share", shareID, err) return "" } return a.Name } -func (r *shareRepositoryWrapper) contentsLabelFromAlbums(shareID string, ids string) string { +func (r *shareRepositoryWrapper) contentsLabelFromAlbums(ctx context.Context, shareID string, ids string) string { idList := strings.Split(ids, ",") - all, err := r.ds.Album(r.ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album.id": idList}}) + all, err := r.ds.Album().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album.id": idList}}) if err != nil { - log.Error(r.ctx, "Error retrieving album names for share", "share", shareID, err) + log.Error(ctx, "Error retrieving album names for share", "share", shareID, err) return "" } names := slice.Map(all, func(a model.Album) string { return a.Name }) return strings.Join(names, ", ") } -func (r *shareRepositoryWrapper) contentsLabelFromPlaylist(shareID string, id string) string { - pls, err := r.ds.Playlist(r.ctx).Get(id) +func (r *shareRepositoryWrapper) contentsLabelFromPlaylist(ctx context.Context, shareID string, id string) string { + pls, err := r.ds.Playlist().Get(ctx, id) if err != nil { - log.Error(r.ctx, "Error retrieving album names for share", "share", shareID, err) + log.Error(ctx, "Error retrieving album names for share", "share", shareID, err) return "" } return pls.Name } -func (r *shareRepositoryWrapper) contentsLabelFromMediaFiles(shareID string, ids string) string { +func (r *shareRepositoryWrapper) contentsLabelFromMediaFiles(ctx context.Context, shareID string, ids string) string { idList := strings.Split(ids, ",") - mfs, err := r.ds.MediaFile(r.ctx).GetAll(model.QueryOptions{Filters: squirrel.And{ + mfs, err := r.ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: squirrel.And{ squirrel.Eq{"media_file.id": idList}, squirrel.Eq{"missing": false}, }}) if err != nil { - log.Error(r.ctx, "Error retrieving media files for share", "share", shareID, err) + log.Error(ctx, "Error retrieving media files for share", "share", shareID, err) return "" } diff --git a/core/share_test.go b/core/share_test.go index e6dec1b92..8ce2bf270 100644 --- a/core/share_test.go +++ b/core/share_test.go @@ -14,27 +14,27 @@ import ( var _ = Describe("Share", func() { var ds model.DataStore var share Share - var mockedRepo rest.Persistable + var mockedRepo rest.Persistable[model.Share] ctx := context.Background() BeforeEach(func() { ds = &tests.MockDataStore{} - mockedRepo = ds.Share(ctx).(rest.Persistable) + mockedRepo = ds.Share().(rest.Persistable[model.Share]) share = NewShare(ds) }) Describe("NewRepository", func() { - var repo rest.Persistable + var repo rest.Persistable[model.Share] BeforeEach(func() { - repo = share.NewRepository(ctx).(rest.Persistable) - _ = ds.Album(ctx).Put(&model.Album{ID: "123", Name: "Album"}) + repo = share.Repository().(rest.Persistable[model.Share]) + _ = ds.Album().Put(ctx, &model.Album{ID: "123", Name: "Album"}) }) Describe("Save", func() { It("it sets a random ID", func() { entity := &model.Share{Description: "test", ResourceIDs: "123"} - id, err := repo.Save(entity) + id, err := repo.Save(ctx, entity) Expect(err).ToNot(HaveOccurred()) Expect(id).ToNot(BeEmpty()) Expect(entity.ID).To(Equal(id)) @@ -42,63 +42,62 @@ var _ = Describe("Share", func() { It("assigns the logged-in user as owner, ignoring a client-supplied UserID", func() { loggedInCtx := request.WithUser(context.Background(), model.User{ID: "logged-in-user"}) - repo := share.NewRepository(loggedInCtx).(rest.Persistable) + repo := share.Repository().(rest.Persistable[model.Share]) entity := &model.Share{Description: "test", ResourceIDs: "123", UserID: "victim-user"} - _, err := repo.Save(entity) + _, err := repo.Save(loggedInCtx, entity) Expect(err).ToNot(HaveOccurred()) Expect(entity.UserID).To(Equal("logged-in-user")) }) It("does not truncate ASCII labels shorter than 30 characters", func() { - _ = ds.MediaFile(ctx).Put(&model.MediaFile{ID: "456", Title: "Example Media File"}) + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "456", Title: "Example Media File"}) entity := &model.Share{Description: "test", ResourceIDs: "456"} - _, err := repo.Save(entity) + _, err := repo.Save(ctx, entity) Expect(err).ToNot(HaveOccurred()) Expect(entity.Contents).To(Equal("Example Media File")) }) It("truncates ASCII labels longer than 30 characters", func() { - _ = ds.MediaFile(ctx).Put(&model.MediaFile{ID: "789", Title: "Example Media File But The Title Is Really Long For Testing Purposes"}) + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "789", Title: "Example Media File But The Title Is Really Long For Testing Purposes"}) entity := &model.Share{Description: "test", ResourceIDs: "789"} - _, err := repo.Save(entity) + _, err := repo.Save(ctx, entity) Expect(err).ToNot(HaveOccurred()) Expect(entity.Contents).To(Equal("Example Media File But The ...")) }) It("does not truncate CJK labels shorter than 30 runes", func() { - _ = ds.MediaFile(ctx).Put(&model.MediaFile{ID: "456", Title: "青春コンプレックス"}) + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "456", Title: "青春コンプレックス"}) entity := &model.Share{Description: "test", ResourceIDs: "456"} - _, err := repo.Save(entity) + _, err := repo.Save(ctx, entity) Expect(err).ToNot(HaveOccurred()) Expect(entity.Contents).To(Equal("青春コンプレックス")) }) It("truncates CJK labels longer than 30 runes", func() { - _ = ds.MediaFile(ctx).Put(&model.MediaFile{ID: "789", Title: "私の中の幻想的世界観及びその顕現を想起させたある現実での出来事に関する一考察"}) + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "789", Title: "私の中の幻想的世界観及びその顕現を想起させたある現実での出来事に関する一考察"}) entity := &model.Share{Description: "test", ResourceIDs: "789"} - _, err := repo.Save(entity) + _, err := repo.Save(ctx, entity) Expect(err).ToNot(HaveOccurred()) Expect(entity.Contents).To(Equal("私の中の幻想的世界観及びその顕現を想起させたある現実で...")) }) It("fails when any of the resource IDs does not exist", func() { entity := &model.Share{Description: "test", ResourceIDs: "123,missing"} - _, err := repo.Save(entity) + _, err := repo.Save(ctx, entity) Expect(err).To(MatchError(model.ErrNotFound)) }) It("fails when the resource IDs are of mixed types", func() { - _ = ds.MediaFile(ctx).Put(&model.MediaFile{ID: "456", Title: "Example Media File"}) + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "456", Title: "Example Media File"}) entity := &model.Share{Description: "test", ResourceIDs: "123,456"} - _, err := repo.Save(entity) + _, err := repo.Save(ctx, entity) Expect(err).To(HaveOccurred()) }) }) Describe("Update", func() { It("filters out read-only fields", func() { - entity := &model.Share{} - err := repo.Update("id", entity) + err := repo.Update(ctx, "id", model.Share{}) Expect(err).ToNot(HaveOccurred()) Expect(mockedRepo.(*tests.MockShareRepo).Cols).To(ConsistOf("description", "downloadable")) }) diff --git a/core/sonic/sonic.go b/core/sonic/sonic.go index 67f5cc7da..8645d3d28 100644 --- a/core/sonic/sonic.go +++ b/core/sonic/sonic.go @@ -100,7 +100,7 @@ func (s *Sonic) GetSonicSimilarTracks(ctx context.Context, id string, count int) return nil, err } - mf, err := s.ds.MediaFile(ctx).Get(id) + mf, err := s.ds.MediaFile().Get(ctx, id) if err != nil { return nil, fmt.Errorf("getting media file %s: %w", id, err) } @@ -120,11 +120,11 @@ func (s *Sonic) FindSonicPath(ctx context.Context, startID, endID string, count return nil, err } - startMF, err := s.ds.MediaFile(ctx).Get(startID) + startMF, err := s.ds.MediaFile().Get(ctx, startID) if err != nil { return nil, fmt.Errorf("getting start media file %s: %w", startID, err) } - endMF, err := s.ds.MediaFile(ctx).Get(endID) + endMF, err := s.ds.MediaFile().Get(ctx, endID) if err != nil { return nil, fmt.Errorf("getting end media file %s: %w", endID, err) } diff --git a/core/stream/decider.go b/core/stream/decider.go index 886898206..38839cdf0 100644 --- a/core/stream/decider.go +++ b/core/stream/decider.go @@ -312,7 +312,7 @@ func (s *deciderService) computeTranscodedStream(ctx context.Context, src *Detai // It checks the DB first (for user-customized values), then falls back to // the built-in defaults, and finally to fallbackBitrate. func lookupDefaultBitrate(ctx context.Context, ds model.DataStore, format string) int { - if t, err := ds.Transcoding(ctx).FindByFormat(format); err == nil && t.DefaultBitRate > 0 { + if t, err := ds.Transcoding().FindByFormat(ctx, format); err == nil && t.DefaultBitRate > 0 { return t.DefaultBitRate } for _, dt := range consts.DefaultTranscodings { @@ -327,7 +327,7 @@ func lookupDefaultBitrate(ctx context.Context, ds model.DataStore, format string // It checks the DB first (for user-customized commands), then falls back to // the built-in default command. Returns "" if the format is unknown. func LookupTranscodeCommand(ctx context.Context, ds model.DataStore, format string) string { - t, err := ds.Transcoding(ctx).FindByFormat(format) + t, err := ds.Transcoding().FindByFormat(ctx, format) if err == nil && t.Command != "" { return t.Command } @@ -447,7 +447,7 @@ func (s *deciderService) ensureProbed(ctx context.Context, mf *model.MediaFile) } mf.ProbeData = string(data) - if err := s.ds.MediaFile(ctx).UpdateProbeData(mf.ID, mf.ProbeData); err != nil { + if err := s.ds.MediaFile().UpdateProbeData(ctx, mf.ID, mf.ProbeData); err != nil { log.Error(ctx, "Failed to persist probe data", "mediaID", mf.ID, err) // Don't fail the decision — we have the data in memory } diff --git a/core/stream/media_streamer.go b/core/stream/media_streamer.go index 6db2f6338..c250b0c4e 100644 --- a/core/stream/media_streamer.go +++ b/core/stream/media_streamer.go @@ -133,7 +133,7 @@ func (ms *mediaStreamer) NewStream(ctx context.Context, mf *model.MediaFile, req } type Stream struct { - ctx context.Context + ctx context.Context //nolint:containedctx // stream outlives the call that built it; Read has no ctx mf *model.MediaFile bitRate int format string diff --git a/core/stream/media_streamer_test.go b/core/stream/media_streamer_test.go index e06599208..5a4bcd480 100644 --- a/core/stream/media_streamer_test.go +++ b/core/stream/media_streamer_test.go @@ -36,7 +36,7 @@ var _ = Describe("MediaStreamer", func() { conf.Server.CacheFolder = conf.NewDir(cacheDir) conf.Server.TranscodingCacheSize = "100MB" ds = &tests.MockDataStore{MockedTranscoding: &tests.MockTranscodingRepo{}} - ds.MediaFile(ctx).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: "123", Path: "tests/fixtures/test.mp3", Suffix: "mp3", BitRate: 128, Duration: 257.0}, }) testCache := stream.NewTranscodingCache() @@ -51,7 +51,7 @@ var _ = Describe("MediaStreamer", func() { var mf *model.MediaFile BeforeEach(func() { var err error - mf, err = ds.MediaFile(ctx).Get("123") + mf, err = ds.MediaFile().Get(ctx, "123") Expect(err).ToNot(HaveOccurred()) }) It("returns a seekable stream if format is 'raw'", func() { @@ -151,7 +151,7 @@ var _ = Describe("MediaStreamer", func() { var mf *model.MediaFile BeforeEach(func() { var err error - mf, err = ds.MediaFile(ctx).Get("123") + mf, err = ds.MediaFile().Get(ctx, "123") Expect(err).ToNot(HaveOccurred()) }) diff --git a/core/user.go b/core/user.go index f13e90167..d67d69dcb 100644 --- a/core/user.go +++ b/core/user.go @@ -15,62 +15,40 @@ type PluginUnloader interface { // User provides business logic for user management with plugin coordination. type User interface { - NewRepository(ctx context.Context) rest.Repository + Repository() rest.Repository[model.User] } type userService struct { - ds model.DataStore - pluginManager PluginUnloader + repo *userRepositoryWrapper } // NewUser creates a new User service func NewUser(ds model.DataStore, pluginManager PluginUnloader) User { return &userService{ - ds: ds, - pluginManager: pluginManager, + repo: &userRepositoryWrapper{ + UserRepository: ds.User(), + pluginManager: pluginManager, + }, } } -// NewRepository returns a REST repository wrapper for user operations. +// Repository returns a REST repository wrapper for user operations. // The wrapper intercepts Delete operations to coordinate plugin unloading. -func (s *userService) NewRepository(ctx context.Context) rest.Repository { - repo := s.ds.User(ctx) - wrapper := &userRepositoryWrapper{ - ctx: ctx, - UserRepository: repo, - pluginManager: s.pluginManager, - } - return wrapper +func (s *userService) Repository() rest.Repository[model.User] { + return s.repo } type userRepositoryWrapper struct { model.UserRepository - ctx context.Context pluginManager PluginUnloader } -// Save implements rest.Persistable by delegating to the underlying repository. -func (r *userRepositoryWrapper) Save(entity any) (string, error) { - return r.UserRepository.(rest.Persistable).Save(entity) -} - -// Update implements rest.Persistable by delegating to the underlying repository. -func (r *userRepositoryWrapper) Update(id string, entity any, cols ...string) error { - return r.UserRepository.(rest.Persistable).Update(id, entity, cols...) -} - -// Delete implements rest.Persistable and coordinates plugin unloading. -func (r *userRepositoryWrapper) Delete(id string) error { - // The underlying repository Delete handles the database cleanup - // including calling cleanupPluginUserReferences - err := r.UserRepository.(rest.Persistable).Delete(id) - if err != nil { - return err - } - - // After successful deletion, check if any plugins were auto-disabled - // and need to be unloaded from memory - r.pluginManager.UnloadDisabledPlugins(r.ctx) - - return nil +var _ rest.Persistable[model.User] = (*userRepositoryWrapper)(nil) + +// Delete unloads plugins even on error: a bulk delete can fail after earlier users were removed +// and their plugins auto-disabled. +func (r *userRepositoryWrapper) Delete(ctx context.Context, ids ...string) error { + err := r.UserRepository.Delete(ctx, ids...) + r.pluginManager.UnloadDisabledPlugins(ctx) + return err } diff --git a/core/user_test.go b/core/user_test.go index b2d3117f8..880637770 100644 --- a/core/user_test.go +++ b/core/user_test.go @@ -29,19 +29,19 @@ var _ = Describe("User Service", func() { }) Describe("NewRepository", func() { - It("returns a rest.Persistable", func() { - repo := service.NewRepository(ctx) - _, ok := repo.(rest.Persistable) + It("returns a rest.Persistable[model.User]", func() { + repo := service.Repository() + _, ok := repo.(rest.Persistable[model.User]) Expect(ok).To(BeTrue()) }) }) Describe("Delete", func() { - var repo rest.Persistable + var repo rest.Persistable[model.User] BeforeEach(func() { - r := service.NewRepository(ctx) - repo = r.(rest.Persistable) + r := service.Repository() + repo = r.(rest.Persistable[model.User]) // Add a test user user := &model.User{ @@ -50,37 +50,45 @@ var _ = Describe("User Service", func() { IsAdmin: false, } user.NewPassword = "password" - Expect(userRepo.Put(user)).To(Succeed()) + Expect(userRepo.Put(ctx, user)).To(Succeed()) }) It("deletes the user successfully", func() { - err := repo.Delete("user-123") + err := repo.Delete(ctx, "user-123") Expect(err).NotTo(HaveOccurred()) // Verify user is deleted - _, err = userRepo.Get("user-123") + _, err = userRepo.Get(ctx, "user-123") Expect(err).To(Equal(model.ErrNotFound)) }) It("calls UnloadDisabledPlugins after successful deletion", func() { - err := repo.Delete("user-123") + err := repo.Delete(ctx, "user-123") Expect(err).NotTo(HaveOccurred()) Expect(pluginManager.unloadCalls).To(Equal(1)) }) - It("does not call UnloadDisabledPlugins when deletion fails", func() { - // Try to delete non-existent user - err := repo.Delete("non-existent") - Expect(err).To(HaveOccurred()) - Expect(pluginManager.unloadCalls).To(Equal(0)) + It("still calls UnloadDisabledPlugins when deletion fails", func() { + err := repo.Delete(ctx, "non-existent") + Expect(err).To(MatchError(model.ErrNotFound)) + Expect(pluginManager.unloadCalls).To(Equal(1)) + }) + + It("unloads plugins when a bulk delete fails after removing earlier users", func() { + err := repo.Delete(ctx, "user-123", "non-existent") + Expect(err).To(MatchError(model.ErrNotFound)) + + _, err = userRepo.Get(ctx, "user-123") + Expect(err).To(Equal(model.ErrNotFound)) + Expect(pluginManager.unloadCalls).To(Equal(1)) }) It("returns error when repository fails", func() { userRepo.Error = errors.New("database error") - err := repo.Delete("user-123") + err := repo.Delete(ctx, "user-123") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("database error")) - Expect(pluginManager.unloadCalls).To(Equal(0)) + Expect(pluginManager.unloadCalls).To(Equal(1)) }) }) }) diff --git a/db/db.go b/db/db.go index a9c6c4a15..66e48cede 100644 --- a/db/db.go +++ b/db/db.go @@ -184,7 +184,7 @@ func isSchemaEmpty(ctx context.Context, db *sql.DB) bool { } type logAdapter struct { - ctx context.Context + ctx context.Context //nolint:containedctx // goose logger interface has no ctx silent bool } diff --git a/go.mod b/go.mod index 9a693d651..01fb0f73d 100644 --- a/go.mod +++ b/go.mod @@ -9,7 +9,7 @@ require ( github.com/Masterminds/squirrel v1.5.4 github.com/andybalholm/cascadia v1.3.5 github.com/bmatcuk/doublestar/v4 v4.10.0 - github.com/deluan/rest v0.0.0-20260913134927-47b21f30cc12 + github.com/deluan/rest v1.0.1 github.com/deluan/sanitize v0.0.0-20241120162836-fdfd8fdfaa55 github.com/dexterlb/mpvipc v0.0.0-20260722094525-0cf47d745b36 github.com/djherbis/atime v1.1.0 diff --git a/go.sum b/go.sum index 4cda9a2f3..bcc1395ba 100644 --- a/go.sum +++ b/go.sum @@ -31,8 +31,8 @@ github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.1 h1:5RVFMOWjMyRy8cARdy79nAmgYw3h github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.1/go.mod h1:ZXNYxsqcloTdSy/rNShjYzMhyjf0LaoftYK0p+A3h40= github.com/deluan/go-taglib v0.0.0-20260913142955-d55e0c9353cb h1:CGVY6RtDsqaleUFogGP03m3a/9OKi3ZMDr5nhm51Emk= github.com/deluan/go-taglib v0.0.0-20260913142955-d55e0c9353cb/go.mod h1:QGxQ4Z1IWyY9w56xNEFjYAaWE8uSxA/gneQ7RPcFJrY= -github.com/deluan/rest v0.0.0-20260913134927-47b21f30cc12 h1:x4N/tx0XC9zcnipk5rpdI4qaG9JmovHW5HL97XYnO1I= -github.com/deluan/rest v0.0.0-20260913134927-47b21f30cc12/go.mod h1:tSgDythFsl0QgS/PFWfIZqcJKnkADWneY80jaVRlqK8= +github.com/deluan/rest v1.0.1 h1:Enuzzfd88C1/lG6Jqr2NgRreJoClaFuWcOuspeKWHZ8= +github.com/deluan/rest v1.0.1/go.mod h1:r0yO0VgBWOb5Xb7aCPIedtcOuwmA2Oe5lGailC5pvc4= github.com/deluan/sanitize v0.0.0-20241120162836-fdfd8fdfaa55 h1:wSCnggTs2f2ji6nFwQmfwgINcmSMj0xF0oHnoyRSPe4= github.com/deluan/sanitize v0.0.0-20241120162836-fdfd8fdfaa55/go.mod h1:ZNCLJfehvEf34B7BbLKjgpsL9lyW7q938w/GY1XgV4E= github.com/dexterlb/mpvipc v0.0.0-20260722094525-0cf47d745b36 h1:KtPfdSST6e0vJbMzMmVqPa5mO1u8vMBlybRCW2ieXpA= @@ -109,8 +109,6 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4= github.com/google/wire v0.7.0/go.mod h1:n6YbUQD9cPKTnHXEBN2DXlOp/mVADhVErcMFb0v3J18= -github.com/gopherjs/gopherjs v0.0.0-20181017120253-0766667cb4d1 h1:EGx4pi6eqNxGaHF6qqu48+N2wcFQ5qg5FXgOdqsJ5d8= -github.com/gopherjs/gopherjs v0.0.0-20181017120253-0766667cb4d1/go.mod h1:wJfORRmW1u3UXTncJ5qlYoELFm8eSnnEO6hX4iZ3EWY= github.com/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= @@ -128,8 +126,6 @@ github.com/jellydator/ttlcache/v3 v3.4.1 h1:bOdXmXiycyK6E6Qjyuj5vl+/vU3SCOoDs8a8 github.com/jellydator/ttlcache/v3 v3.4.1/go.mod h1:j7LO12PNghFg5+0v9budMAT4rDK4JY969jb9vOdOBBk= github.com/joshdk/go-junit v1.0.0 h1:S86cUKIdwBHWwA6xCmFlf3RTLfVXYQfvanM5Uh+K6GE= github.com/joshdk/go-junit v1.0.0/go.mod h1:TiiV0PqkaNfFXjEiyjWM3XXrhVyCa1K4Zfga6W52ung= -github.com/jtolds/gls v4.20.0+incompatible h1:xdiiI2gbIgH/gLH7ADydsJ1uDOEzR8yvV7C0MuV77Wo= -github.com/jtolds/gls v4.20.0+incompatible/go.mod h1:QJZ7F/aHp+rZTRtaJ1ow/lLfFfVYBRgL+9YlvaHOwJU= github.com/kardianos/service v1.3.0 h1:/LGy+xPP2TM+GLTiCZ2di7cy0Jd/qrawlTUfqKYFdTI= github.com/kardianos/service v1.3.0/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc= github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51 h1:Z9n2FFNUXsshfwJMBgNA0RU6/i7WVaAegv3PtuIHPMs= @@ -138,7 +134,6 @@ github.com/klauspost/compress v1.19.2 h1:hMRETovs/pu/dVWN7zIT1PGG8t509MwT6bO7XSi github.com/klauspost/compress v1.19.2/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/cpuid/v2 v2.4.0 h1:S6Hrbc7+ywsr0r+RLapfGBHfyefhCTwEh3A0tV913Dw= github.com/klauspost/cpuid/v2 v2.4.0/go.mod h1:19jmZ9mjzoF//ddRSUsv0zfBTJWh3QJh9FNxZTMrGxU= -github.com/konsorten/go-windows-terminal-sequences v1.0.1/go.mod h1:T0+1ngSBFLxvqU3pZ+m/2kptfBszLMUkC4ZK/EgS/cQ= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= @@ -231,13 +226,8 @@ github.com/segmentio/asm v1.2.1 h1:DTNbBqs57ioxAD4PrArqftgypG4/qNpXoJx8TVXxPR0= github.com/segmentio/asm v1.2.1/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= github.com/sethvargo/go-retry v0.4.0 h1:9qy1OoIAxBL+gBYnkTnTnWle5wlfsXQlwRzIbbpdqPw= github.com/sethvargo/go-retry v0.4.0/go.mod h1:tvsjdKG6xfiCx4LSiUZ06kcv38xvdVQwv8R6/VnnVWg= -github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE= github.com/sirupsen/logrus v1.10.2 h1:G2SED73/qrAu6YwbdxOD6peLkCBI3z7L+ykJFTXJBBo= github.com/sirupsen/logrus v1.10.2/go.mod h1:SLEg8TqYulVKKfIGHldVp2K2aYz2DKSVBq4g/H5bR7Q= -github.com/smartystreets/assertions v0.0.0-20180927180507-b2de0cb4f26d h1:zE9ykElWQ6/NYmHa3jpm/yHnI4xSofP+UP6SpjHcSeM= -github.com/smartystreets/assertions v0.0.0-20180927180507-b2de0cb4f26d/go.mod h1:OnSkiWE9lh6wB0YB77sQom3nweQdgAjqCqsofrRNTgc= -github.com/smartystreets/goconvey v1.6.4 h1:fv0U8FUIMPNf1L9lnHLvLhgicrIVChEkdzIKYqbNC9s= -github.com/smartystreets/goconvey v1.6.4/go.mod h1:syvi0/a8iFYH4r/RixwvyeAJjdLS9QV7WQ/tjFTllLA= github.com/sosodev/duration v1.3.1 h1:qtHBDMQ6lvMQsL15g4aopM4HEfOaYuhWBw3NPTtlqq4= github.com/sosodev/duration v1.3.1/go.mod h1:RQIBBX0+fMLc/D9+Jb/fwvVmo0eZvDDEERAikUR6SDg= github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I= @@ -252,7 +242,6 @@ github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3A github.com/spf13/viper v1.21.0 h1:x5S+0EU27Lbphp4UKm1C+1oQO+rKx36vfCoaVebLFSU= github.com/spf13/viper v1.21.0/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= -github.com/stretchr/objx v0.1.1/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= @@ -310,7 +299,6 @@ golang.org/x/image v0.46.0 h1:b1+oYj0Jbp6K5MDT4i4/eZpYlk3V8SJhhDKh6LBHAyQ= golang.org/x/image v0.46.0/go.mod h1:3B3W05VGVQyuXucLINLjXKrqISASfi4Xj+iCVkLMwew= golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c= golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o= -golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190603091049-60506f45cf65/go.mod h1:HSz+uSET+XFnRR8LxR5pz3Of3rY3CfYBVs4xY44aLks= golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues= golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg= @@ -318,7 +306,6 @@ golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0= golang.org/x/sys v0.0.0-20180926160741-c2ed4eda69e7/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= -golang.org/x/sys v0.0.0-20190422165155-953cdadca894/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20220615213510-4f61da869c0c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= @@ -333,7 +320,6 @@ golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E= golang.org/x/time v0.16.0 h1:vMb6ptszcQMkcwiRTAuNNU50gom6++Q/6gY2hDM6VDE= golang.org/x/time v0.16.0/go.mod h1:rVKOqvZeKvrDKTQiAHJ7wmwP0RzleSphoEA9RcdLA0s= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= -golang.org/x/tools v0.0.0-20190328211700-ab21143f2384/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI= golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo= google.golang.org/appengine v1.6.5/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc= diff --git a/model/album.go b/model/album.go index 114d19e1e..a43195419 100644 --- a/model/album.go +++ b/model/album.go @@ -1,14 +1,15 @@ package model import ( + "context" "iter" "math" "sync" "time" - "github.com/navidrome/navidrome/conf" - + "github.com/deluan/rest" "github.com/gohugoio/hashstructure" + "github.com/navidrome/navidrome/conf" ) type Album struct { @@ -137,24 +138,25 @@ type Albums []Album type AlbumCursor iter.Seq2[Album, error] type AlbumRepository interface { - CountAll(...QueryOptions) (int64, error) - Exists(id string) (bool, error) - Put(*Album) error - UpdateExternalInfo(*Album) error - Get(id string) (*Album, error) - GetAll(...QueryOptions) (Albums, error) + rest.Repository[Album] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Exists(ctx context.Context, id string) (bool, error) + Put(ctx context.Context, m *Album) error + UpdateExternalInfo(ctx context.Context, m *Album) error + Get(ctx context.Context, id string) (*Album, error) + GetAll(ctx context.Context, options ...QueryOptions) (Albums, error) // GetSoleAlbumArtistIDsInSubtrees returns the sole album artists of the albums with folders in // any of the given library-relative subtrees. - GetSoleAlbumArtistIDsInSubtrees(lib Library, paths ...string) ([]string, error) - GetCursor(...QueryOptions) (AlbumCursor, error) - GetYears(libraryIDs ...int) ([]int, error) + GetSoleAlbumArtistIDsInSubtrees(ctx context.Context, lib Library, paths ...string) ([]string, error) + GetCursor(ctx context.Context, options ...QueryOptions) (AlbumCursor, error) + GetYears(ctx context.Context, libraryIDs ...int) ([]int, error) // The following methods are used exclusively by the scanner: - Touch(ids ...string) error - TouchByMissingFolder() (int64, error) - GetTouchedAlbums(libID int) (AlbumCursor, error) - RefreshPlayCounts() (int64, error) - CopyAttributes(fromID, toID string, columns ...string) error + Touch(ctx context.Context, ids ...string) error + TouchByMissingFolder(ctx context.Context) (int64, error) + GetTouchedAlbums(ctx context.Context, libID int) (AlbumCursor, error) + RefreshPlayCounts(ctx context.Context) (int64, error) + CopyAttributes(ctx context.Context, fromID, toID string, columns ...string) error AnnotatedRepository SearchableRepository[Albums] diff --git a/model/annotation.go b/model/annotation.go index 5228028a6..64b8ddc17 100644 --- a/model/annotation.go +++ b/model/annotation.go @@ -1,6 +1,9 @@ package model -import "time" +import ( + "context" + "time" +) type Annotations struct { PlayCount int64 `structs:"play_count" json:"playCount,omitempty"` @@ -13,8 +16,8 @@ type Annotations struct { } type AnnotatedRepository interface { - IncPlayCount(itemID string, ts time.Time) error - SetStar(starred bool, itemIDs ...string) error - SetRating(rating int, itemID string) error - ReassignAnnotation(prevID string, newID string) error + IncPlayCount(ctx context.Context, itemID string, ts time.Time) error + SetStar(ctx context.Context, starred bool, itemIDs ...string) error + SetRating(ctx context.Context, rating int, itemID string) error + ReassignAnnotation(ctx context.Context, prevID string, newID string) error } diff --git a/model/artist.go b/model/artist.go index f88b3a974..985cc6c7b 100644 --- a/model/artist.go +++ b/model/artist.go @@ -1,11 +1,13 @@ package model import ( + "context" "iter" "maps" "slices" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/consts" ) @@ -84,18 +86,19 @@ type ArtistIndexes []ArtistIndex type ArtistCursor iter.Seq2[Artist, error] type ArtistRepository interface { - CountAll(options ...QueryOptions) (int64, error) - Exists(id string) (bool, error) - Put(m *Artist, colsToUpdate ...string) error - UpdateExternalInfo(a *Artist) error - Get(id string) (*Artist, error) - GetAll(options ...QueryOptions) (Artists, error) - GetCursor(options ...QueryOptions) (ArtistCursor, error) - GetIndex(includeMissing bool, libraryIds []int, roles ...Role) (ArtistIndexes, error) + rest.Repository[Artist] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Exists(ctx context.Context, id string) (bool, error) + Put(ctx context.Context, m *Artist, colsToUpdate ...string) error + UpdateExternalInfo(ctx context.Context, a *Artist) error + Get(ctx context.Context, id string) (*Artist, error) + GetAll(ctx context.Context, options ...QueryOptions) (Artists, error) + GetCursor(ctx context.Context, options ...QueryOptions) (ArtistCursor, error) + GetIndex(ctx context.Context, includeMissing bool, libraryIds []int, roles ...Role) (ArtistIndexes, error) // The following methods are used exclusively by the scanner: - RefreshPlayCounts() (int64, error) - RefreshStats(allArtists bool) (int64, error) + RefreshPlayCounts(ctx context.Context) (int64, error) + RefreshStats(ctx context.Context, allArtists bool) (int64, error) AnnotatedRepository SearchableRepository[Artists] diff --git a/model/artwork.go b/model/artwork.go index 8856a35fa..d5f62ab13 100644 --- a/model/artwork.go +++ b/model/artwork.go @@ -1,6 +1,9 @@ package model -import "time" +import ( + "context" + "time" +) // Artwork is one unique image, identified by the XXH3-64 hash of its bytes. type Artwork struct { @@ -116,60 +119,60 @@ const ( // Delete* takes the rows to remove; Purge* finds them itself and reports how many went. type ArtworkRepository interface { - GetImage(hash string) (*Artwork, error) - PutImage(a *Artwork) error + GetImage(ctx context.Context, hash string) (*Artwork, error) + PutImage(ctx context.Context, a *Artwork) error // PurgeOrphans deletes rows referenced by no item_artwork row and older than cutoff. - PurgeOrphans(createdBefore time.Time) (int64, error) - GetItemArtwork(kind Kind, id, imageType string) (*ItemArtwork, error) - PutItemArtwork(ia *ItemArtwork) error + PurgeOrphans(ctx context.Context, createdBefore time.Time) (int64, error) + GetItemArtwork(ctx context.Context, kind Kind, id, imageType string) (*ItemArtwork, error) + PutItemArtwork(ctx context.Context, ia *ItemArtwork) error // PutLastFailure records the trace of the attempt that exhausted the retry budget. - PutLastFailure(kind Kind, id, imageType, trace string) error - DeleteForItems(kind Kind, ids []string) error + PutLastFailure(ctx context.Context, kind Kind, id, imageType, trace string) error + DeleteForItems(ctx context.Context, kind Kind, ids []string) error // GetInfoForItems hydrates a page in one batched query. - GetInfoForItems(kind Kind, ids []string) (map[string]ItemArtworkInfo, error) + GetInfoForItems(ctx context.Context, kind Kind, ids []string) (map[string]ItemArtworkInfo, error) // GetMimeByHash returns hash -> current mime for every stored artwork. - GetMimeByHash() (map[string]string, error) + GetMimeByHash(ctx context.Context) (map[string]string, error) // PurgeDanglingItems removes state rows whose entity no longer exists. - PurgeDanglingItems() (int64, error) + PurgeDanglingItems(ctx context.Context) (int64, error) } type ArtworkQueueRepository interface { // Get returns the pending row for an item, or ErrNotFound when it is not queued. - Get(kind Kind, id, imageType string) (*ArtworkQueueItem, error) + Get(ctx context.Context, kind Kind, id, imageType string) (*ArtworkQueueItem, error) // Enqueue upserts; an existing row keeps the higher priority and has its retry_at reset. - Enqueue(items ...ArtworkQueueItem) error + Enqueue(ctx context.Context, items ...ArtworkQueueItem) error // EnqueuePreservingBackoff upserts like Enqueue but preserves an existing row's retry_at, so a // request-triggered read-through never resets a failed resolution's backoff. - EnqueuePreservingBackoff(items ...ArtworkQueueItem) error + EnqueuePreservingBackoff(ctx context.Context, items ...ArtworkQueueItem) error // EnqueueAllMissing inserts queue rows for all entities with no item_artwork row, at the given priority. - EnqueueAllMissing(kind Kind, priority int) (int64, error) + EnqueueAllMissing(ctx context.Context, kind Kind, priority int) (int64, error) // EnqueueIfMissing inserts only for items with no item_artwork row yet. - EnqueueIfMissing(items ...ArtworkQueueItem) error + EnqueueIfMissing(ctx context.Context, items ...ArtworkQueueItem) error // CountBySource reports how many items of a kind currently resolve from the given sources. // An empty sources slice means every source; "" matches absent state, and the pseudo-source // ArtworkSourceFailed matches the absent states that gave up. - CountBySource(kind Kind, sources []string) (int64, error) + CountBySource(ctx context.Context, kind Kind, sources []string) (int64, error) // SourcesInUse lists the distinct sources items of a kind currently resolve from, "" included. - SourcesInUse(kind Kind) ([]string, error) + SourcesInUse(ctx context.Context, kind Kind) ([]string, error) // EnqueueBySource inserts queue rows for items of a kind whose current source matches. // It does not clear existing artwork state: the current image stays until it is replaced. - EnqueueBySource(kind Kind, sources []string, priority int) (int64, error) + EnqueueBySource(ctx context.Context, kind Kind, sources []string, priority int) (int64, error) // DequeueBatch returns up to n items with retry_at <= now, priority desc, enqueued_at asc. // Restricted to the given kinds when any are passed, so one kind cannot block another's drain. - DequeueBatch(n int, kinds ...string) ([]ArtworkQueueItem, error) + DequeueBatch(ctx context.Context, n int, kinds ...string) ([]ArtworkQueueItem, error) // MarkFailedIfUnchanged applies the failure backoff only while retry_at still matches // seenRetryAt, so a concurrent re-enqueue keeps its fresh eligibility. - MarkFailedIfUnchanged(kind, id, imageType string, seenRetryAt, retryAt time.Time, trace string) error + MarkFailedIfUnchanged(ctx context.Context, kind, id, imageType string, seenRetryAt, retryAt time.Time, trace string) error // DeleteIfUnchanged deletes only while retry_at still matches, sparing a concurrent re-enqueue. - DeleteIfUnchanged(kind, id, imageType string, retryAt time.Time) error - Count() (int64, error) + DeleteIfUnchanged(ctx context.Context, kind, id, imageType string, retryAt time.Time) error + Count(ctx context.Context) (int64, error) // CountQueued reports the pending rows matching the kinds and priorities, grouped by both; // an empty filter means every one. - CountQueued(kinds []Kind, priorities []int) ([]ArtworkQueueStat, error) + CountQueued(ctx context.Context, kinds []Kind, priorities []int) ([]ArtworkQueueStat, error) // PurgeDangling removes queue rows whose entity no longer exists. - PurgeDangling() (int64, error) + PurgeDangling(ctx context.Context) (int64, error) // PurgeQueued removes pending rows matching the kinds and priorities; an empty filter means every one. - PurgeQueued(kinds []Kind, priorities []int) (int64, error) + PurgeQueued(ctx context.Context, kinds []Kind, priorities []int) (int64, error) } type ArtworkQueueStat struct { diff --git a/model/bookmark.go b/model/bookmark.go index 7c6637ce9..768d4a462 100644 --- a/model/bookmark.go +++ b/model/bookmark.go @@ -1,15 +1,18 @@ package model -import "time" +import ( + "context" + "time" +) type Bookmarkable struct { BookmarkPosition int64 `structs:"-" json:"bookmarkPosition"` } type BookmarkableRepository interface { - AddBookmark(id, comment string, position int64) error - DeleteBookmark(id string) error - GetBookmarks() (Bookmarks, error) + AddBookmark(ctx context.Context, id, comment string, position int64) error + DeleteBookmark(ctx context.Context, id string) error + GetBookmarks(ctx context.Context) (Bookmarks, error) } type Bookmark struct { diff --git a/model/datastore.go b/model/datastore.go index 26687d5d4..6ded8c575 100644 --- a/model/datastore.go +++ b/model/datastore.go @@ -4,7 +4,6 @@ import ( "context" "github.com/Masterminds/squirrel" - "github.com/deluan/rest" ) type QueryOptions struct { @@ -16,34 +15,28 @@ type QueryOptions struct { Seed string // for random sorting } -type ResourceRepository interface { - rest.Repository -} - type DataStore interface { - Library(ctx context.Context) LibraryRepository - Folder(ctx context.Context) FolderRepository - Album(ctx context.Context) AlbumRepository - Artist(ctx context.Context) ArtistRepository - MediaFile(ctx context.Context) MediaFileRepository - Genre(ctx context.Context) GenreRepository - Tag(ctx context.Context) TagRepository - Playlist(ctx context.Context) PlaylistRepository - PlayQueue(ctx context.Context) PlayQueueRepository - Transcoding(ctx context.Context) TranscodingRepository - Player(ctx context.Context) PlayerRepository - Radio(ctx context.Context) RadioRepository - Share(ctx context.Context) ShareRepository - Property(ctx context.Context) PropertyRepository - User(ctx context.Context) UserRepository - UserProps(ctx context.Context) UserPropsRepository - ScrobbleBuffer(ctx context.Context) ScrobbleBufferRepository - Scrobble(ctx context.Context) ScrobbleRepository - Plugin(ctx context.Context) PluginRepository - Artwork(ctx context.Context) ArtworkRepository - ArtworkQueue(ctx context.Context) ArtworkQueueRepository - - Resource(ctx context.Context, model any) ResourceRepository + Library() LibraryRepository + Folder() FolderRepository + Album() AlbumRepository + Artist() ArtistRepository + MediaFile() MediaFileRepository + Genre() GenreRepository + Tag() TagRepository + Playlist() PlaylistRepository + PlayQueue() PlayQueueRepository + Transcoding() TranscodingRepository + Player() PlayerRepository + Radio() RadioRepository + Share() ShareRepository + Property() PropertyRepository + User() UserRepository + UserProps() UserPropsRepository + ScrobbleBuffer() ScrobbleBufferRepository + Scrobble() ScrobbleRepository + Plugin() PluginRepository + Artwork() ArtworkRepository + ArtworkQueue() ArtworkQueueRepository WithTx(block func(tx DataStore) error, scope ...string) error WithTxImmediate(block func(tx DataStore) error, scope ...string) error diff --git a/model/folder.go b/model/folder.go index 5207a9db0..701fca8cd 100644 --- a/model/folder.go +++ b/model/folder.go @@ -1,6 +1,7 @@ package model import ( + "context" "fmt" "iter" "os" @@ -83,19 +84,19 @@ type FolderUpdateInfo struct { } type FolderRepository interface { - Get(id string) (*Folder, error) - GetByPath(lib Library, path string) (*Folder, error) - GetAll(...QueryOptions) ([]Folder, error) - CountAll(...QueryOptions) (int64, error) - GetFolderUpdateInfo(lib Library, targetPaths ...string) (map[string]FolderUpdateInfo, error) + Get(ctx context.Context, id string) (*Folder, error) + GetByPath(ctx context.Context, lib Library, path string) (*Folder, error) + GetAll(ctx context.Context, options ...QueryOptions) ([]Folder, error) + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + GetFolderUpdateInfo(ctx context.Context, lib Library, targetPaths ...string) (map[string]FolderUpdateInfo, error) // HasAudioOutsideFolders reports whether any folder in parent's subtree // (including parent itself) contains audio files and is not one of the // given folder IDs. - HasAudioOutsideFolders(parent Folder, excludeFolderIDs []string) (bool, error) - Put(*Folder) error - MarkMissing(missing bool, ids ...string) error - GetTouchedWithPlaylists() (FolderCursor, error) + HasAudioOutsideFolders(ctx context.Context, parent Folder, excludeFolderIDs []string) (bool, error) + Put(ctx context.Context, f *Folder) error + MarkMissing(ctx context.Context, missing bool, ids ...string) error + GetTouchedWithPlaylists(ctx context.Context) (FolderCursor, error) // GetAllWithPlaylists returns all non-missing folders with playlists, ignoring // the scan-timestamp gate used by GetTouchedWithPlaylists. - GetAllWithPlaylists() (FolderCursor, error) + GetAllWithPlaylists(ctx context.Context) (FolderCursor, error) } diff --git a/model/genre.go b/model/genre.go index fa5b6ec62..147422e67 100644 --- a/model/genre.go +++ b/model/genre.go @@ -1,5 +1,11 @@ package model +import ( + "context" + + "github.com/deluan/rest" +) + type Genre struct { ID string `structs:"id" json:"id,omitempty" toml:"id,omitempty" yaml:"id,omitempty"` Name string `structs:"name" json:"name"` @@ -10,6 +16,7 @@ type Genre struct { type Genres []Genre type GenreRepository interface { - GetAll(...QueryOptions) (Genres, error) - Get(id string) (*Genre, error) + rest.Repository[Genre] + GetAll(ctx context.Context, options ...QueryOptions) (Genres, error) + Get(ctx context.Context, id string) (*Genre, error) } diff --git a/model/get_entity.go b/model/get_entity.go index e5d41f1be..cfb3968ef 100644 --- a/model/get_entity.go +++ b/model/get_entity.go @@ -23,11 +23,11 @@ func getEntity(ctx context.Context, ds DataStore, id string) (any, Kind, error) kind Kind get func() (any, error) }{ - {KindArtistArtwork, func() (any, error) { return ds.Artist(ctx).Get(id) }}, - {KindAlbumArtwork, func() (any, error) { return ds.Album(ctx).Get(id) }}, - {KindPlaylistArtwork, func() (any, error) { return ds.Playlist(ctx).Get(id) }}, - {KindMediaFileArtwork, func() (any, error) { return ds.MediaFile(ctx).Get(id) }}, - {KindRadioArtwork, func() (any, error) { return ds.Radio(ctx).Get(id) }}, + {KindArtistArtwork, func() (any, error) { return ds.Artist().Get(ctx, id) }}, + {KindAlbumArtwork, func() (any, error) { return ds.Album().Get(ctx, id) }}, + {KindPlaylistArtwork, func() (any, error) { return ds.Playlist().Get(ctx, id) }}, + {KindMediaFileArtwork, func() (any, error) { return ds.MediaFile().Get(ctx, id) }}, + {KindRadioArtwork, func() (any, error) { return ds.Radio().Get(ctx, id) }}, } for _, g := range getters { entity, err := g.get() diff --git a/model/get_entity_test.go b/model/get_entity_test.go index 4e589406a..73e21a036 100644 --- a/model/get_entity_test.go +++ b/model/get_entity_test.go @@ -19,7 +19,7 @@ var _ = Describe("GetEntityByID", func() { }) It("returns the entity matching the id", func() { - ds.Album(ctx).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "a1", Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "a1", Name: "One"}}) entity, err := model.GetEntityByID(ctx, ds, "a1") Expect(err).ToNot(HaveOccurred()) Expect(entity).To(BeAssignableToTypeOf(&model.Album{})) @@ -32,7 +32,7 @@ var _ = Describe("GetEntityByID", func() { }) It("propagates unexpected repository errors instead of reporting not-found", func() { - ds.Album(ctx).(*tests.MockAlbumRepo).SetError(true) + ds.Album().(*tests.MockAlbumRepo).SetError(true) _, err := model.GetEntityByID(ctx, ds, "a1") Expect(err).To(HaveOccurred()) Expect(err).ToNot(MatchError(model.ErrNotFound)) @@ -49,7 +49,7 @@ var _ = Describe("GetEntityKindByID", func() { }) It("returns the artwork kind for the matching id", func() { - ds.Album(ctx).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "a1"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "a1"}}) kind, err := model.GetEntityKindByID(ctx, ds, "a1") Expect(err).ToNot(HaveOccurred()) Expect(kind).To(Equal(model.KindAlbumArtwork)) diff --git a/model/library.go b/model/library.go index aceab533a..1e33222ac 100644 --- a/model/library.go +++ b/model/library.go @@ -1,8 +1,10 @@ package model import ( + "context" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/utils/slice" ) @@ -39,23 +41,24 @@ func (l Libraries) IDs() []int { } type LibraryRepository interface { - Get(id int) (*Library, error) + rest.Repository[Library] + Get(ctx context.Context, id int) (*Library, error) // GetPath returns the path of the library with the given ID. // Its implementation must be optimized to avoid unnecessary queries. - GetPath(id int) (string, error) - GetAll(...QueryOptions) (Libraries, error) - CountAll(...QueryOptions) (int64, error) - Put(l *Library, colsToUpdate ...string) error - Delete(id int) error - StoreMusicFolder() error - AddArtist(id int, artistID string) error + GetPath(ctx context.Context, id int) (string, error) + GetAll(ctx context.Context, options ...QueryOptions) (Libraries, error) + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Put(ctx context.Context, l *Library, colsToUpdate ...string) error + Delete(ctx context.Context, id int) error + StoreMusicFolder(ctx context.Context) error + AddArtist(ctx context.Context, id int, artistID string) error // User-library association methods - GetUsersWithLibraryAccess(libraryID int) (Users, error) + GetUsersWithLibraryAccess(ctx context.Context, libraryID int) (Users, error) // TODO These methods should be moved to a core service - ScanBegin(id int, fullScan bool) error - ScanEnd(id int) error - ScanInProgress() (bool, error) - RefreshStats(id int) error + ScanBegin(ctx context.Context, id int, fullScan bool) error + ScanEnd(ctx context.Context, id int) error + ScanInProgress(ctx context.Context) (bool, error) + RefreshStats(ctx context.Context, id int) error } diff --git a/model/mediafile.go b/model/mediafile.go index 1147eaa35..0c56b4825 100644 --- a/model/mediafile.go +++ b/model/mediafile.go @@ -2,6 +2,7 @@ package model import ( "cmp" + "context" "encoding/json" "fmt" "iter" @@ -11,6 +12,7 @@ import ( "strings" "time" + "github.com/deluan/rest" "github.com/gohugoio/hashstructure" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" @@ -537,42 +539,43 @@ func (mfs MediaFiles) ToM3U8(title string, absolutePaths bool) string { type MediaFileCursor iter.Seq2[MediaFile, error] type MediaFileRepository interface { - CountAll(options ...QueryOptions) (int64, error) - CountBySuffix(options ...QueryOptions) (map[string]int64, error) - Exists(id string) (bool, error) - Put(m *MediaFile) error - UpdateProbeData(id string, data string) error - Get(id string) (*MediaFile, error) - GetWithParticipants(id string) (*MediaFile, error) - GetAll(options ...QueryOptions) (MediaFiles, error) + rest.Repository[MediaFile] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + CountBySuffix(ctx context.Context, options ...QueryOptions) (map[string]int64, error) + Exists(ctx context.Context, id string) (bool, error) + Put(ctx context.Context, m *MediaFile) error + UpdateProbeData(ctx context.Context, id string, data string) error + Get(ctx context.Context, id string) (*MediaFile, error) + GetWithParticipants(ctx context.Context, id string) (*MediaFile, error) + GetAll(ctx context.Context, options ...QueryOptions) (MediaFiles, error) // GetRandom returns up to options.Max media files in random order, applying the same // filters as GetAll. Sort/Order are ignored. - GetRandom(options ...QueryOptions) (MediaFiles, error) - GetAllByTags(tag TagName, values []string, options ...QueryOptions) (MediaFiles, error) + GetRandom(ctx context.Context, options ...QueryOptions) (MediaFiles, error) + GetAllByTags(ctx context.Context, tag TagName, values []string, options ...QueryOptions) (MediaFiles, error) // MatchesCriteria reports whether the media file matches the criteria's rule // expression, using the logged user's annotations. Limit and offset are ignored. - MatchesCriteria(id string, c criteria.Criteria) (bool, error) - GetCursor(options ...QueryOptions) (MediaFileCursor, error) + MatchesCriteria(ctx context.Context, id string, c criteria.Criteria) (bool, error) + GetCursor(ctx context.Context, options ...QueryOptions) (MediaFileCursor, error) // GetAlbumIDsByFolder returns the distinct IDs of albums with non-missing tracks in the given // folders or their direct children. - GetAlbumIDsByFolder(lib Library, folderIDs ...string) ([]string, error) + GetAlbumIDsByFolder(ctx context.Context, lib Library, folderIDs ...string) ([]string, error) // GetCursorWithArtwork streams like GetCursor, hydrated, so callers that render images don't // pay the scanner's per-row cost; it uses the same id pre-pass as the other cursors. - GetCursorWithArtwork(options ...QueryOptions) (MediaFileCursor, error) - Delete(id string) error - DeleteMissing(ids []string) error - DeleteAllMissing() (int64, error) - FindByPaths(paths []string) (MediaFiles, error) + GetCursorWithArtwork(ctx context.Context, options ...QueryOptions) (MediaFileCursor, error) + Delete(ctx context.Context, id string) error + DeleteMissing(ctx context.Context, ids []string) error + DeleteAllMissing(ctx context.Context) (int64, error) + FindByPaths(ctx context.Context, paths []string) (MediaFiles, error) // ReassignReferences moves annotations, bookmarks and playlist entries from prevID to newID, // keeping newID's own row wherever a user has both. - ReassignReferences(prevID, newID string) error + ReassignReferences(ctx context.Context, prevID, newID string) error // The following methods are used exclusively by the scanner: - MarkMissing(bool, ...*MediaFile) error - MarkMissingByFolder(missing bool, folderIDs ...string) error - GetMissingAndMatching(libId int) (MediaFileCursor, error) - FindRecentFilesByMBZTrackID(missing MediaFile, since time.Time) (MediaFiles, error) - FindRecentFilesByProperties(missing MediaFile, since time.Time) (MediaFiles, error) + MarkMissing(ctx context.Context, missing bool, mfs ...*MediaFile) error + MarkMissingByFolder(ctx context.Context, missing bool, folderIDs ...string) error + GetMissingAndMatching(ctx context.Context, libId int) (MediaFileCursor, error) + FindRecentFilesByMBZTrackID(ctx context.Context, missing MediaFile, since time.Time) (MediaFiles, error) + FindRecentFilesByProperties(ctx context.Context, missing MediaFile, since time.Time) (MediaFiles, error) AnnotatedRepository BookmarkableRepository diff --git a/model/player.go b/model/player.go index 39ea99d1a..2e4484a10 100644 --- a/model/player.go +++ b/model/player.go @@ -1,7 +1,10 @@ package model import ( + "context" "time" + + "github.com/deluan/rest" ) type Player struct { @@ -23,9 +26,11 @@ type Player struct { type Players []Player type PlayerRepository interface { - Get(id string) (*Player, error) - FindMatch(userId, client, userAgent string) (*Player, error) - Put(p *Player) error - CountAll(...QueryOptions) (int64, error) - CountByClient(...QueryOptions) (map[string]int64, error) + rest.Repository[Player] + rest.Persistable[Player] + Get(ctx context.Context, id string) (*Player, error) + FindMatch(ctx context.Context, userId, client, userAgent string) (*Player, error) + Put(ctx context.Context, p *Player) error + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + CountByClient(ctx context.Context, options ...QueryOptions) (map[string]int64, error) } diff --git a/model/playlist.go b/model/playlist.go index 19fedc9fa..d2ed97682 100644 --- a/model/playlist.go +++ b/model/playlist.go @@ -1,6 +1,7 @@ package model import ( + "context" "iter" "maps" "os" @@ -9,6 +10,7 @@ import ( "strconv" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/log" @@ -206,19 +208,19 @@ type Playlists []Playlist type PlaylistCursor iter.Seq2[Playlist, error] type PlaylistRepository interface { - ResourceRepository + rest.Repository[Playlist] + rest.Persistable[Playlist] AnnotatedRepository - CountAll(options ...QueryOptions) (int64, error) - Exists(id string) (bool, error) - Put(pls *Playlist, cols ...string) error - Get(id string) (*Playlist, error) - GetWithTracks(id string, refreshSmartPlaylist, includeMissing bool) (*Playlist, error) - GetAll(options ...QueryOptions) (Playlists, error) - GetCursor(options ...QueryOptions) (PlaylistCursor, error) - FindByPath(path string) (*Playlist, error) - Delete(id string) error - Tracks(playlistId string, refreshSmartPlaylist bool) PlaylistTrackRepository - GetPlaylists(mediaFileId string) (Playlists, error) + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Exists(ctx context.Context, id string) (bool, error) + Put(ctx context.Context, pls *Playlist, cols ...string) error + Get(ctx context.Context, id string) (*Playlist, error) + GetWithTracks(ctx context.Context, id string, refreshSmartPlaylist, includeMissing bool) (*Playlist, error) + GetAll(ctx context.Context, options ...QueryOptions) (Playlists, error) + GetCursor(ctx context.Context, options ...QueryOptions) (PlaylistCursor, error) + FindByPath(ctx context.Context, path string) (*Playlist, error) + Tracks(ctx context.Context, playlistId string, refreshSmartPlaylist bool) PlaylistTrackRepository + GetPlaylists(ctx context.Context, mediaFileId string) (Playlists, error) } type PlaylistTrack struct { @@ -241,18 +243,18 @@ func (plt PlaylistTracks) MediaFiles() MediaFiles { type PlaylistTrackCursor iter.Seq2[PlaylistTrack, error] type PlaylistTrackRepository interface { - ResourceRepository - CountAll(options ...QueryOptions) (int64, error) - GetAll(options ...QueryOptions) (PlaylistTracks, error) - GetCursor(options ...QueryOptions) (PlaylistTrackCursor, error) - GetAlbumIDs(options ...QueryOptions) ([]string, error) - GetMediaFileIDs(options ...QueryOptions) ([]string, error) - Add(mediaFileIds []string) (int, error) - Insert(mediaFileIds []string, pos int) (int, error) - AddAlbums(albumIds []string) (int, error) - AddArtists(artistIds []string) (int, error) - AddDiscs(discs []DiscID) (int, error) - Delete(id ...string) error - DeleteAll() error - Reorder(pos int, newPos int) error + rest.Repository[PlaylistTrack] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + GetAll(ctx context.Context, options ...QueryOptions) (PlaylistTracks, error) + GetCursor(ctx context.Context, options ...QueryOptions) (PlaylistTrackCursor, error) + GetAlbumIDs(ctx context.Context, options ...QueryOptions) ([]string, error) + GetMediaFileIDs(ctx context.Context, options ...QueryOptions) ([]string, error) + Add(ctx context.Context, mediaFileIds []string) (int, error) + Insert(ctx context.Context, mediaFileIds []string, pos int) (int, error) + AddAlbums(ctx context.Context, albumIds []string) (int, error) + AddArtists(ctx context.Context, artistIds []string) (int, error) + AddDiscs(ctx context.Context, discs []DiscID) (int, error) + Delete(ctx context.Context, id ...string) error + DeleteAll(ctx context.Context) error + Reorder(ctx context.Context, pos int, newPos int) error } diff --git a/model/playqueue.go b/model/playqueue.go index 03b562253..ddc53d7da 100644 --- a/model/playqueue.go +++ b/model/playqueue.go @@ -1,6 +1,7 @@ package model import ( + "context" "time" ) @@ -18,11 +19,11 @@ type PlayQueue struct { type PlayQueues []PlayQueue type PlayQueueRepository interface { - Store(queue *PlayQueue, colNames ...string) error + Store(ctx context.Context, queue *PlayQueue, colNames ...string) error // Retrieve returns the playqueue without loading the full MediaFiles // (Items only contain IDs) - Retrieve(userId string) (*PlayQueue, error) + Retrieve(ctx context.Context, userId string) (*PlayQueue, error) // RetrieveWithMediaFiles returns the playqueue with full MediaFiles loaded - RetrieveWithMediaFiles(userId string) (*PlayQueue, error) - Clear(userId string) error + RetrieveWithMediaFiles(ctx context.Context, userId string) (*PlayQueue, error) + Clear(ctx context.Context, userId string) error } diff --git a/model/plugin.go b/model/plugin.go index 18d66e305..a448a9633 100644 --- a/model/plugin.go +++ b/model/plugin.go @@ -1,6 +1,11 @@ package model -import "time" +import ( + "context" + "time" + + "github.com/deluan/rest" +) type Plugin struct { ID string `structs:"id" json:"id"` @@ -22,11 +27,11 @@ type Plugin struct { type Plugins []Plugin type PluginRepository interface { - ResourceRepository - ClearErrors() error - CountAll(options ...QueryOptions) (int64, error) - Delete(id string) error - Get(id string) (*Plugin, error) - GetAll(options ...QueryOptions) (Plugins, error) - Put(p *Plugin) error + rest.Repository[Plugin] + ClearErrors(ctx context.Context) error + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Delete(ctx context.Context, id string) error + Get(ctx context.Context, id string) (*Plugin, error) + GetAll(ctx context.Context, options ...QueryOptions) (Plugins, error) + Put(ctx context.Context, p *Plugin) error } diff --git a/model/properties.go b/model/properties.go index 06bb9ebec..24b56db26 100644 --- a/model/properties.go +++ b/model/properties.go @@ -1,8 +1,10 @@ package model +import "context" + type PropertyRepository interface { - Put(id string, value string) error - Get(id string) (string, error) - Delete(id string) error - DefaultGet(id string, defaultValue string) (string, error) + Put(ctx context.Context, id string, value string) error + Get(ctx context.Context, id string) (string, error) + Delete(ctx context.Context, id string) error + DefaultGet(ctx context.Context, id string, defaultValue string) (string, error) } diff --git a/model/radio.go b/model/radio.go index 013a24beb..a7e0e9218 100644 --- a/model/radio.go +++ b/model/radio.go @@ -1,8 +1,10 @@ package model import ( + "context" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/consts" ) @@ -29,11 +31,11 @@ func (r Radio) UploadedImagePath() string { type Radios []Radio type RadioRepository interface { - ResourceRepository - CountAll(options ...QueryOptions) (int64, error) - Delete(id string) error - Exists(id string) (bool, error) - Get(id string) (*Radio, error) - GetAll(options ...QueryOptions) (Radios, error) - Put(u *Radio, colsToUpdate ...string) error + rest.Repository[Radio] + rest.Persistable[Radio] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Exists(ctx context.Context, id string) (bool, error) + Get(ctx context.Context, id string) (*Radio, error) + GetAll(ctx context.Context, options ...QueryOptions) (Radios, error) + Put(ctx context.Context, u *Radio, colsToUpdate ...string) error } diff --git a/model/scrobble.go b/model/scrobble.go index a8022fc16..b99d25b6a 100644 --- a/model/scrobble.go +++ b/model/scrobble.go @@ -1,6 +1,11 @@ package model -import "time" +import ( + "context" + "time" + + "github.com/deluan/rest" +) type Scrobble struct { ID int64 `structs:"id" json:"id"` @@ -10,10 +15,11 @@ type Scrobble struct { } type ScrobbleRepository interface { - CountAll(options ...QueryOptions) (int64, error) - Get(id string) (*Scrobble, error) - GetAll(options ...QueryOptions) (Scrobbles, error) - RecordScrobble(mediaFileID string, submissionTime time.Time) error + rest.Repository[Scrobble] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Get(ctx context.Context, id string) (*Scrobble, error) + GetAll(ctx context.Context, options ...QueryOptions) (Scrobbles, error) + RecordScrobble(ctx context.Context, mediaFileID string, submissionTime time.Time) error } type Scrobbles []Scrobble diff --git a/model/scrobble_buffer.go b/model/scrobble_buffer.go index 43ee2cc01..dda8dfce3 100644 --- a/model/scrobble_buffer.go +++ b/model/scrobble_buffer.go @@ -1,6 +1,9 @@ package model -import "time" +import ( + "context" + "time" +) type ScrobbleEntry struct { ID string @@ -15,10 +18,10 @@ type ScrobbleEntry struct { type ScrobbleEntries []ScrobbleEntry type ScrobbleBufferRepository interface { - UserIDs(service string) ([]string, error) - Enqueue(service, userId, mediaFileId string, playTime time.Time) error - Next(service string, userId string) (*ScrobbleEntry, error) - Dequeue(entry *ScrobbleEntry) error - Length() (int64, error) - Discard(service string) error + UserIDs(ctx context.Context, service string) ([]string, error) + Enqueue(ctx context.Context, service, userId, mediaFileId string, playTime time.Time) error + Next(ctx context.Context, service string, userId string) (*ScrobbleEntry, error) + Dequeue(ctx context.Context, entry *ScrobbleEntry) error + Length(ctx context.Context) (int64, error) + Discard(ctx context.Context, service string) error } diff --git a/model/searchable.go b/model/searchable.go index a64a0171c..6dba989c1 100644 --- a/model/searchable.go +++ b/model/searchable.go @@ -1,5 +1,7 @@ package model +import "context" + type SearchableRepository[T any] interface { - Search(q string, options ...QueryOptions) (T, error) + Search(ctx context.Context, q string, options ...QueryOptions) (T, error) } diff --git a/model/share.go b/model/share.go index cf1f4cb34..952fd9759 100644 --- a/model/share.go +++ b/model/share.go @@ -2,9 +2,11 @@ package model import ( "cmp" + "context" "strings" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/utils/random" ) @@ -56,8 +58,10 @@ func (s Share) ToM3U8() string { } type ShareRepository interface { - Exists(id string) (bool, error) - Get(id string) (*Share, error) - GetAll(options ...QueryOptions) (Shares, error) - CountAll(options ...QueryOptions) (int64, error) + rest.Repository[Share] + rest.Persistable[Share] + Exists(ctx context.Context, id string) (bool, error) + Get(ctx context.Context, id string) (*Share, error) + GetAll(ctx context.Context, options ...QueryOptions) (Shares, error) + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) } diff --git a/model/tag.go b/model/tag.go index 152b7164e..60fa252dd 100644 --- a/model/tag.go +++ b/model/tag.go @@ -2,10 +2,12 @@ package model import ( "cmp" + "context" "fmt" "slices" "strings" + "github.com/deluan/rest" "github.com/navidrome/navidrome/model/id" "github.com/navidrome/navidrome/utils/slice" "github.com/zeebo/xxh3" @@ -162,9 +164,10 @@ func (t Tags) Add(name TagName, v string) { } type TagRepository interface { - Add(libraryID int, tags ...Tag) error - UpdateCounts() error - GetAll(name TagName, options ...QueryOptions) (TagList, error) + rest.Repository[Tag] + Add(ctx context.Context, libraryID int, tags ...Tag) error + UpdateCounts(ctx context.Context) error + GetAll(ctx context.Context, name TagName, options ...QueryOptions) (TagList, error) } type TagName string diff --git a/model/transcoding.go b/model/transcoding.go index 9b81a7c9c..da5e09e91 100644 --- a/model/transcoding.go +++ b/model/transcoding.go @@ -1,5 +1,11 @@ package model +import ( + "context" + + "github.com/deluan/rest" +) + type Transcoding struct { ID string `structs:"id" json:"id"` Name string `structs:"name" json:"name"` @@ -11,8 +17,10 @@ type Transcoding struct { type Transcodings []Transcoding type TranscodingRepository interface { - Get(id string) (*Transcoding, error) - CountAll(...QueryOptions) (int64, error) - Put(*Transcoding) error - FindByFormat(format string) (*Transcoding, error) + rest.Repository[Transcoding] + rest.Persistable[Transcoding] + Get(ctx context.Context, id string) (*Transcoding, error) + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Put(ctx context.Context, t *Transcoding) error + FindByFormat(ctx context.Context, format string) (*Transcoding, error) } diff --git a/model/user.go b/model/user.go index 37bdca33d..bbeb4a05d 100644 --- a/model/user.go +++ b/model/user.go @@ -1,7 +1,10 @@ package model import ( + "context" "time" + + "github.com/deluan/rest" ) type User struct { @@ -46,21 +49,21 @@ func (u User) HasLibraryAccess(libraryID int) bool { type Users []User type UserRepository interface { - ResourceRepository - CountAll(...QueryOptions) (int64, error) - Delete(id string) error - Get(id string) (*User, error) - GetAll(options ...QueryOptions) (Users, error) - Put(*User) error - UpdateLastLoginAt(id string) error - UpdateLastAccessAt(id string) error - FindFirstAdmin() (*User, error) + rest.Repository[User] + rest.Persistable[User] + CountAll(ctx context.Context, options ...QueryOptions) (int64, error) + Get(ctx context.Context, id string) (*User, error) + GetAll(ctx context.Context, options ...QueryOptions) (Users, error) + Put(ctx context.Context, u *User) error + UpdateLastLoginAt(ctx context.Context, id string) error + UpdateLastAccessAt(ctx context.Context, id string) error + FindFirstAdmin(ctx context.Context) (*User, error) // FindByUsername must be case-insensitive - FindByUsername(username string) (*User, error) + FindByUsername(ctx context.Context, username string) (*User, error) // FindByUsernameWithPassword is the same as above, but also returns the decrypted password - FindByUsernameWithPassword(username string) (*User, error) + FindByUsernameWithPassword(ctx context.Context, username string) (*User, error) // Library association methods - GetUserLibraries(userID string) (Libraries, error) - SetUserLibraries(userID string, libraryIDs []int) error + GetUserLibraries(ctx context.Context, userID string) (Libraries, error) + SetUserLibraries(ctx context.Context, userID string, libraryIDs []int) error } diff --git a/model/user_props.go b/model/user_props.go index c2eb536ec..af19f34f1 100644 --- a/model/user_props.go +++ b/model/user_props.go @@ -1,8 +1,10 @@ package model +import "context" + type UserPropsRepository interface { - Put(userId, key string, value string) error - Get(userId, key string) (string, error) - Delete(userId, key string) error - DefaultGet(userId, key string, defaultValue string) (string, error) + Put(ctx context.Context, userId, key string, value string) error + Get(ctx context.Context, userId, key string) (string, error) + Delete(ctx context.Context, userId, key string) error + DefaultGet(ctx context.Context, userId, key string, defaultValue string) (string, error) } diff --git a/persistence/album_repository.go b/persistence/album_repository.go index 808c880fb..8e7662b8f 100644 --- a/persistence/album_repository.go +++ b/persistence/album_repository.go @@ -102,9 +102,8 @@ func (as dbAlbums) toModels() model.Albums { return slice.Map(as, func(a dbAlbum) model.Album { return *a.Album }) } -func NewAlbumRepository(ctx context.Context, db dbx.Builder) model.AlbumRepository { +func NewAlbumRepository(db dbx.Builder) model.AlbumRepository { r := &albumRepository{} - r.ctx = ctx r.db = db r.tableName = "album" r.registerModel(&model.Album{}, albumFilters()) @@ -190,49 +189,49 @@ func allRolesFilter(_ string, value any) Sqlizer { return ParticipantIDFilter("album", value) } -func (r *albumRepository) CountAll(options ...model.QueryOptions) (int64, error) { - query := r.newSelect() - query = r.applyLibraryFilter(query) +func (r *albumRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + query := r.newSelect(ctx) + query = r.applyLibraryFilter(ctx, query) if filtersNeedAnnotation(r.applyFilters(query, options...)) { - query = r.withAnnotation(query, "album.id") + query = r.withAnnotation(ctx, query, "album.id") } - return r.count(query, options...) + return r.count(ctx, query, options...) } -func (r *albumRepository) Exists(id string) (bool, error) { +func (r *albumRepository) Exists(ctx context.Context, id string) (bool, error) { // The exists() helper applies no library filter, so it would report rows the caller cannot see. - c, err := r.count(r.applyLibraryFilter(r.newSelect().Where(Eq{"album.id": id}))) + c, err := r.count(ctx, r.applyLibraryFilter(ctx, r.newSelect(ctx).Where(Eq{"album.id": id}))) return c > 0, err } -func (r *albumRepository) Put(al *model.Album) error { +func (r *albumRepository) Put(ctx context.Context, al *model.Album) error { al.ImportedAt = time.Now() - id, err := r.put(al.ID, &dbAlbum{Album: al}) + id, err := r.put(ctx, al.ID, &dbAlbum{Album: al}) if err != nil { return err } al.ID = id - if err := r.updateParticipants(al.ID, al.Participants); err != nil { + if err := r.updateParticipants(ctx, al.ID, al.Participants); err != nil { return err } - return r.updateTags(al.ID, al.Tags) + return r.updateTags(ctx, al.ID, al.Tags) } // TODO Move external metadata to a separated table -func (r *albumRepository) UpdateExternalInfo(al *model.Album) error { - _, err := r.put(al.ID, &dbAlbum{Album: al}, "description", "small_image_url", "medium_image_url", "large_image_url", "external_url", "external_info_updated_at") +func (r *albumRepository) UpdateExternalInfo(ctx context.Context, al *model.Album) error { + _, err := r.put(ctx, al.ID, &dbAlbum{Album: al}, "description", "small_image_url", "medium_image_url", "large_image_url", "external_url", "external_info_updated_at") return err } -func (r *albumRepository) selectAlbum(options ...model.QueryOptions) SelectBuilder { - sql := r.newSelect(options...).Columns("album.*", "library.path as library_path", "library.name as library_name"). +func (r *albumRepository) selectAlbum(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + sql := r.newSelect(ctx, options...).Columns("album.*", "library.path as library_path", "library.name as library_name"). LeftJoin("library on album.library_id = library.id") - sql = r.withAnnotation(sql, "album.id") - return r.applyLibraryFilter(sql) + sql = r.withAnnotation(ctx, sql, "album.id") + return r.applyLibraryFilter(ctx, sql) } -func (r *albumRepository) Get(id string) (*model.Album, error) { - res, err := r.GetAll(model.QueryOptions{Filters: Eq{"album.id": id}}) +func (r *albumRepository) Get(ctx context.Context, id string) (*model.Album, error) { + res, err := r.GetAll(ctx, model.QueryOptions{Filters: Eq{"album.id": id}}) if err != nil { return nil, err } @@ -242,31 +241,31 @@ func (r *albumRepository) Get(id string) (*model.Album, error) { return &res[0], nil } -func (r *albumRepository) GetAll(options ...model.QueryOptions) (model.Albums, error) { - sq := r.selectAlbum(options...) +func (r *albumRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Albums, error) { + sq := r.selectAlbum(ctx, options...) var res dbAlbums - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) if err != nil { return nil, err } albums := res.toModels() - r.hydrateArtwork(albums) + r.hydrateArtwork(ctx, albums) return albums, nil } -func (r *albumRepository) hydrateArtwork(albums model.Albums) { - hydrateItems(r.ctx, r.db, model.KindAlbumArtwork, albums, +func (r *albumRepository) hydrateArtwork(ctx context.Context, albums model.Albums) { + hydrateItems(ctx, r.db, model.KindAlbumArtwork, albums, func(a *model.Album) (string, *model.ItemImage) { return a.ID, &a.ItemImage }) } // getAllIDs returns the IDs of GetAll's row set, skipping its column projection and JSON decoding. -func (r *albumRepository) getAllIDs(options ...model.QueryOptions) ([]string, error) { - sq := r.applyLibraryFilter(r.newSelect(options...).Columns("album.id")) +func (r *albumRepository) getAllIDs(ctx context.Context, options ...model.QueryOptions) ([]string, error) { + sq := r.applyLibraryFilter(ctx, r.newSelect(ctx, options...).Columns("album.id")) if filtersNeedAnnotation(sq) { - sq = r.withAnnotation(sq, "album.id") + sq = r.withAnnotation(ctx, sq, "album.id") } ids := []string{} - err := r.queryAllSlice(sq, &ids) + err := r.queryAllSlice(ctx, sq, &ids) return ids, err } @@ -282,7 +281,7 @@ func SoleAlbumArtistFilter(artistID string) Sqlizer { // GetSoleAlbumArtistIDsInSubtrees matches albums by their own folder_ids, which is the resolver's // notion of an album's folders. -func (r *albumRepository) GetSoleAlbumArtistIDsInSubtrees(lib model.Library, paths ...string) ([]string, error) { +func (r *albumRepository) GetSoleAlbumArtistIDsInSubtrees(ctx context.Context, lib model.Library, paths ...string) ([]string, error) { if len(paths) == 0 { return nil, nil } @@ -295,7 +294,7 @@ func (r *albumRepository) GetSoleAlbumArtistIDsInSubtrees(lib model.Library, pat sq := Select("distinct json_extract(participants, '$.albumartist[0].id')").From("album"). Where(And{soleAlbumArtistFilter, inSubtree}) var chunkIDs []string - if err := r.queryAllSlice(sq, &chunkIDs); err != nil { + if err := r.queryAllSlice(ctx, sq, &chunkIDs); err != nil { return nil, err } ids = append(ids, chunkIDs...) @@ -303,37 +302,37 @@ func (r *albumRepository) GetSoleAlbumArtistIDsInSubtrees(lib model.Library, pat return ids, nil } -func (r *albumRepository) GetCursor(options ...model.QueryOptions) (model.AlbumCursor, error) { - ids, err := r.getAllIDs(options...) +func (r *albumRepository) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.AlbumCursor, error) { + ids, err := r.getAllIDs(ctx, options...) if err != nil { return nil, err } opts := chunkOptions(options, "album.id") return model.AlbumCursor(streamByIDs(ids, func(chunk []string) (model.Albums, error) { - return r.GetAll(opts(chunk)) + return r.GetAll(ctx, opts(chunk)) })), nil } -func (r *albumRepository) GetYears(libraryIDs ...int) ([]int, error) { +func (r *albumRepository) GetYears(ctx context.Context, libraryIDs ...int) ([]int, error) { cond := And{Gt{"max_year": 0}, Eq{"missing": false}} if len(libraryIDs) > 0 { cond = append(cond, Eq{"library_id": libraryIDs}) } - sq := r.applyLibraryFilter(Select("distinct max_year").From("album").Where(cond).OrderBy("max_year")) + sq := r.applyLibraryFilter(ctx, Select("distinct max_year").From("album").Where(cond).OrderBy("max_year")) years := []int{} - err := r.queryAllSlice(sq, &years) + err := r.queryAllSlice(ctx, sq, &years) if err != nil && !errors.Is(err, model.ErrNotFound) { return nil, err } return years, nil } -func (r *albumRepository) CopyAttributes(fromID, toID string, columns ...string) error { +func (r *albumRepository) CopyAttributes(ctx context.Context, fromID, toID string, columns ...string) error { // Cast values to text so go-sqlite3 does not decode datetime columns as time.Time // and reformat them as RFC3339 when written back. sel := slice.Map(columns, func(c string) string { return fmt.Sprintf("cast(%[1]s as text) as %[1]s", c) }) var from dbx.NullStringMap - err := r.queryOne(Select(sel...).From(r.tableName).Where(Eq{"id": fromID}), &from) + err := r.queryOne(ctx, Select(sel...).From(r.tableName).Where(Eq{"id": fromID}), &from) if err != nil { return fmt.Errorf("getting album to copy fields from: %w", err) } @@ -351,35 +350,35 @@ func (r *albumRepository) CopyAttributes(fromID, toID string, columns ...string) if len(to) == 0 { return nil } - _, err = r.executeSQL(Update(r.tableName).SetMap(to).Where(Eq{"id": toID})) + _, err = r.executeSQL(ctx, Update(r.tableName).SetMap(to).Where(Eq{"id": toID})) return err } // Touch flags an album as being scanned by the scanner, but not necessarily updated. // This is used for when missing tracks are detected for an album during scan. -func (r *albumRepository) Touch(ids ...string) error { +func (r *albumRepository) Touch(ctx context.Context, ids ...string) error { if len(ids) == 0 { return nil } for ids := range slices.Chunk(ids, 200) { upd := Update(r.tableName).Set("imported_at", time.Now()).Where(Eq{"id": ids}) - c, err := r.executeSQL(upd) + c, err := r.executeSQL(ctx, upd) if err != nil { return fmt.Errorf("error touching albums: %w", err) } - log.Debug(r.ctx, "Touching albums", "ids", ids, "updated", c) + log.Debug(ctx, "Touching albums", "ids", ids, "updated", c) } return nil } // TouchByMissingFolder touches all albums that have missing folders -func (r *albumRepository) TouchByMissingFolder() (int64, error) { +func (r *albumRepository) TouchByMissingFolder(ctx context.Context) (int64, error) { upd := Update(r.tableName).Set("imported_at", time.Now()). Where(And{ NotEq{"folder_ids": nil}, ConcatExpr("EXISTS (SELECT 1 FROM json_each(folder_ids) AS je JOIN main.folder AS f ON je.value = f.id WHERE f.missing = true)"), }) - c, err := r.executeSQL(upd) + c, err := r.executeSQL(ctx, upd) if err != nil { return 0, fmt.Errorf("error touching albums by missing folder: %w", err) } @@ -389,13 +388,13 @@ func (r *albumRepository) TouchByMissingFolder() (int64, error) { // GetTouchedAlbums returns all albums that were touched by the scanner for a given library, in the // current library scan run. // It does not need to load participants, as they are not used by the scanner. -func (r *albumRepository) GetTouchedAlbums(libID int) (model.AlbumCursor, error) { - query := r.selectAlbum(). +func (r *albumRepository) GetTouchedAlbums(ctx context.Context, libID int) (model.AlbumCursor, error) { + query := r.selectAlbum(ctx). Where(And{ Eq{"library.id": libID}, ConcatExpr("album.imported_at > library.last_scan_at"), }) - cursor, err := queryWithStableResults[dbAlbum](r.sqlRepository, query) + cursor, err := queryWithStableResults[dbAlbum](ctx, r.sqlRepository, query) if err != nil { return nil, err } @@ -408,7 +407,7 @@ func wrapAlbumCursor(cursor iter.Seq2[dbAlbum, error]) model.AlbumCursor { // RefreshPlayCounts updates the play count and last play date annotations for all albums, based // on the media files associated with them. -func (r *albumRepository) RefreshPlayCounts() (int64, error) { +func (r *albumRepository) RefreshPlayCounts(ctx context.Context) (int64, error) { query := Expr(` with play_counts as ( select user_id, album_id, sum(play_count) as total_play_count, max(play_date) as last_play_date @@ -424,21 +423,21 @@ on conflict (user_id, item_id, item_type) do update set play_count = excluded.play_count, play_date = excluded.play_date; `) - return r.executeSQL(query) + return r.executeSQL(ctx, query) } -func (r *albumRepository) purgeEmpty(libraryIDs ...int) error { +func (r *albumRepository) purgeEmpty(ctx context.Context, libraryIDs ...int) error { del := Delete(r.tableName).Where("id not in (select distinct(album_id) from media_file)") // If libraryIDs are specified, only purge albums from those libraries if len(libraryIDs) > 0 { del = del.Where(Eq{"library_id": libraryIDs}) } - c, err := r.executeSQL(del) + c, err := r.executeSQL(ctx, del) if err != nil { return fmt.Errorf("purging empty albums: %w", err) } if c > 0 { - log.Debug(r.ctx, "Purged empty albums", "totalDeleted", c) + log.Debug(ctx, "Purged empty albums", "totalDeleted", c) } return nil } @@ -449,40 +448,32 @@ var albumSearchConfig = searchConfig{ MBIDFields: []string{"mbz_album_id", "mbz_release_group_id"}, } -func (r *albumRepository) Search(q string, options ...model.QueryOptions) (model.Albums, error) { +func (r *albumRepository) Search(ctx context.Context, q string, options ...model.QueryOptions) (model.Albums, error) { var opts model.QueryOptions if len(options) > 0 { opts = options[0] } var res dbAlbums - err := r.doSearch(r.selectAlbum(options...), q, &res, albumSearchConfig, opts) + err := r.doSearch(ctx, r.selectAlbum(ctx, options...), q, &res, albumSearchConfig, opts) if err != nil { return nil, fmt.Errorf("searching album %q: %w", q, err) } albums := res.toModels() - r.hydrateArtwork(albums) + r.hydrateArtwork(ctx, albums) return albums, nil } -func (r *albumRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *albumRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *albumRepository) Read(id string) (any, error) { - return r.Get(id) +func (r *albumRepository) Read(ctx context.Context, id string) (*model.Album, error) { + return r.Get(ctx, id) } -func (r *albumRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) -} - -func (r *albumRepository) EntityName() string { - return "album" -} - -func (r *albumRepository) NewInstance() any { - return &model.Album{} +func (r *albumRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Album, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } var _ model.AlbumRepository = (*albumRepository)(nil) -var _ model.ResourceRepository = (*albumRepository)(nil) +var _ rest.Repository[model.Album] = (*albumRepository)(nil) diff --git a/persistence/album_repository_test.go b/persistence/album_repository_test.go index c3d7f018f..7652e703d 100644 --- a/persistence/album_repository_test.go +++ b/persistence/album_repository_test.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "sort" + "sync" "time" "github.com/Masterminds/squirrel" @@ -22,11 +23,11 @@ import ( // rawColumn returns a column exactly as stored, bypassing go-sqlite3's decoding of // `datetime` columns into time.Time. -func rawColumn(r sqlRepository, id, column string) string { +func rawColumn(ctx context.Context, r sqlRepository, id, column string) string { var res struct{ Value string } sel := squirrel.Select("cast(" + column + " as text) as value"). From(r.tableName).Where(squirrel.Eq{"id": id}) - ExpectWithOffset(1, r.queryOne(sel, &res)).To(Succeed()) + ExpectWithOffset(1, r.queryOne(ctx, sel, &res)).To(Succeed()) return res.Value } @@ -36,7 +37,7 @@ var _ = Describe("AlbumRepository", func() { BeforeEach(func() { ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "userid", UserName: "johndoe"}) - albumRepo = NewAlbumRepository(ctx, GetDBXBuilder()).(*albumRepository) + albumRepo = NewAlbumRepository(GetDBXBuilder()).(*albumRepository) }) Describe("natural sorting", func() { @@ -48,12 +49,12 @@ var _ = Describe("AlbumRepository", func() { for _, n := range []string{"foo 1", "foo 10", "foo 2", "foo 20", "foo 3"} { aid := "nat-" + n ids = append(ids, aid) - Expect(albumRepo.Put(&model.Album{ + Expect(albumRepo.Put(ctx, &model.Album{ ID: aid, LibraryID: 1, Name: n, OrderAlbumName: n, })).To(Succeed()) } DeferCleanup(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": ids})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": ids})) }) }) @@ -61,8 +62,8 @@ var _ = Describe("AlbumRepository", func() { func(naturalSorting, preferSortTags bool, expected []string) { conf.Server.EnableNaturalSorting = naturalSorting conf.Server.PreferSortTags = preferSortTags - albumRepo = NewAlbumRepository(ctx, GetDBXBuilder()).(*albumRepository) - albums, err := albumRepo.GetAll(model.QueryOptions{ + albumRepo = NewAlbumRepository(GetDBXBuilder()).(*albumRepository) + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Sort: "name", Filters: squirrel.Eq{"album.id": ids}, }) Expect(err).ToNot(HaveOccurred()) @@ -79,7 +80,7 @@ var _ = Describe("AlbumRepository", func() { Describe("Get", func() { var Get = func(id string) (*model.Album, error) { - album, err := albumRepo.Get(id) + album, err := albumRepo.Get(ctx, id) if album != nil { album.ImportedAt = time.Time{} } @@ -99,29 +100,29 @@ var _ = Describe("AlbumRepository", func() { BeforeEach(func() { srcTime = time.Date(2020, 1, 2, 3, 4, 5, 0, time.UTC) dstTime = time.Date(2024, 6, 7, 8, 9, 10, 0, time.UTC) - Expect(albumRepo.Put(&model.Album{ID: "copy-src", Name: "src", LibraryID: 1, CreatedAt: srcTime})).To(Succeed()) - Expect(albumRepo.Put(&model.Album{ID: "copy-dst", Name: "dst", LibraryID: 1, CreatedAt: dstTime})).To(Succeed()) - Expect(albumRepo.Put(&model.Album{ID: "copy-zero", Name: "zero", LibraryID: 1})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{ID: "copy-src", Name: "src", LibraryID: 1, CreatedAt: srcTime})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{ID: "copy-dst", Name: "dst", LibraryID: 1, CreatedAt: dstTime})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{ID: "copy-zero", Name: "zero", LibraryID: 1})).To(Succeed()) DeferCleanup(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"copy-src", "copy-dst", "copy-zero"}})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"copy-src", "copy-dst", "copy-zero"}})) }) }) It("copies a valid created_at from source to destination", func() { - Expect(albumRepo.CopyAttributes("copy-src", "copy-dst", "created_at")).To(Succeed()) - got, err := albumRepo.Get("copy-dst") + Expect(albumRepo.CopyAttributes(ctx, "copy-src", "copy-dst", "created_at")).To(Succeed()) + got, err := albumRepo.Get(ctx, "copy-dst") Expect(err).ToNot(HaveOccurred()) Expect(got.CreatedAt).To(BeTemporally("~", srcTime, time.Second)) }) It("leaves destination untouched when source created_at is zero", func() { - Expect(albumRepo.CopyAttributes("copy-zero", "copy-dst", "created_at")).To(Succeed()) - got, err := albumRepo.Get("copy-dst") + Expect(albumRepo.CopyAttributes(ctx, "copy-zero", "copy-dst", "created_at")).To(Succeed()) + got, err := albumRepo.Get(ctx, "copy-dst") Expect(err).ToNot(HaveOccurred()) Expect(got.CreatedAt).To(BeTemporally("~", dstTime, time.Second)) }) It("returns not found and leaves destination untouched when source does not exist", func() { - err := albumRepo.CopyAttributes("copy-missing", "copy-dst", "created_at") + err := albumRepo.CopyAttributes(ctx, "copy-missing", "copy-dst", "created_at") Expect(errors.Is(err, model.ErrNotFound)).To(BeTrue()) - got, getErr := albumRepo.Get("copy-dst") + got, getErr := albumRepo.Get(ctx, "copy-dst") Expect(getErr).ToNot(HaveOccurred()) Expect(got.CreatedAt).To(BeTemporally("~", dstTime, time.Second)) }) @@ -129,35 +130,35 @@ var _ = Describe("AlbumRepository", func() { // Copying through a Go string would rewrite it as RFC3339 ("2020-01-02T03:04:05Z"), // which string-sorts above every space-format timestamp and pins the album to the // top of "Recently Added". - Expect(albumRepo.CopyAttributes("copy-src", "copy-dst", "created_at")).To(Succeed()) - Expect(rawColumn(albumRepo.sqlRepository, "copy-dst", "created_at")). - To(Equal(rawColumn(albumRepo.sqlRepository, "copy-src", "created_at"))) - Expect(rawColumn(albumRepo.sqlRepository, "copy-dst", "created_at")).ToNot(ContainSubstring("T")) + Expect(albumRepo.CopyAttributes(ctx, "copy-src", "copy-dst", "created_at")).To(Succeed()) + Expect(rawColumn(ctx, albumRepo.sqlRepository, "copy-dst", "created_at")). + To(Equal(rawColumn(ctx, albumRepo.sqlRepository, "copy-src", "created_at"))) + Expect(rawColumn(ctx, albumRepo.sqlRepository, "copy-dst", "created_at")).ToNot(ContainSubstring("T")) }) }) Describe("GetCursor", func() { It("yields the same albums as GetAll", func() { opts := model.QueryOptions{Sort: "name"} - want, err := albumRepo.GetAll(opts) + want, err := albumRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(albumRepo.GetCursor(opts))).To(Equal([]model.Album(want))) + Expect(collectCursor(albumRepo.GetCursor(ctx, opts))).To(Equal([]model.Album(want))) }) It("honors Max/Offset like GetAll", func() { opts := model.QueryOptions{Sort: "name", Max: 2, Offset: 1} - want, err := albumRepo.GetAll(opts) + want, err := albumRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(albumRepo.GetCursor(opts))).To(Equal([]model.Album(want))) + Expect(collectCursor(albumRepo.GetCursor(ctx, opts))).To(Equal([]model.Album(want))) }) }) Describe("getAllIDs", func() { It("returns the same id set as GetAll", func() { - want, err := albumRepo.GetAll() + want, err := albumRepo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - ids, err := albumRepo.getAllIDs() + ids, err := albumRepo.getAllIDs(ctx) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(ConsistOf(slice.Map(want, func(a model.Album) string { return a.ID }))) }) @@ -165,13 +166,13 @@ var _ = Describe("AlbumRepository", func() { Describe("GetSoleAlbumArtistIDsInSubtrees", func() { It("returns the sole album artists of albums with folders in the subtree", func() { - folderRepo := newFolderRepository(ctx, GetDBXBuilder()) - lib, err := NewLibraryRepository(ctx, GetDBXBuilder()).Get(1) + folderRepo := newFolderRepository(GetDBXBuilder()) + lib, err := NewLibraryRepository(GetDBXBuilder()).Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) inTree := model.NewFolder(*lib, "SubtreeAlbums/Artist") outTree := model.NewFolder(*lib, "OtherTree/Artist") - Expect(folderRepo.Put(inTree)).To(Succeed()) - Expect(folderRepo.Put(outTree)).To(Succeed()) + Expect(folderRepo.Put(ctx, inTree)).To(Succeed()) + Expect(folderRepo.Put(ctx, outTree)).To(Succeed()) // album_artist_id is deliberately wrong: the artist must come from participants inAl := model.Album{ID: "subtree-in-al", Name: "In", LibraryID: 1, AlbumArtistID: "999", FolderIDs: []string{inTree.ID}, @@ -181,34 +182,34 @@ var _ = Describe("AlbumRepository", func() { duoAl := model.Album{ID: "subtree-duo-al", Name: "Duo", LibraryID: 1, AlbumArtistID: "5", FolderIDs: []string{inTree.ID}, Participants: model.Participants{model.RoleAlbumArtist: []model.Participant{{Artist: artistPunctuation}, {Artist: artistBeatles}}}} for _, al := range []model.Album{inAl, outAl, duoAl} { - Expect(albumRepo.Put(&al)).To(Succeed()) + Expect(albumRepo.Put(ctx, &al)).To(Succeed()) } DeferCleanup(func() { _, _ = GetDBXBuilder().NewQuery("DELETE FROM album WHERE id LIKE 'subtree-%'").Execute() _, _ = GetDBXBuilder().NewQuery("DELETE FROM folder WHERE path LIKE 'SubtreeAlbums%' OR path LIKE 'OtherTree%'").Execute() }) - ids, err := albumRepo.GetSoleAlbumArtistIDsInSubtrees(*lib, "SubtreeAlbums") + ids, err := albumRepo.GetSoleAlbumArtistIDsInSubtrees(ctx, *lib, "SubtreeAlbums") Expect(err).ToNot(HaveOccurred()) Expect(ids).To(ConsistOf("2")) // sole artist in the subtree; the duo and the outside album are excluded }) It("stays under SQLite's expression tree depth limit with many paths", func() { - lib, err := NewLibraryRepository(ctx, GetDBXBuilder()).Get(1) + lib, err := NewLibraryRepository(GetDBXBuilder()).Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) paths := make([]string, 200) for i := range paths { paths[i] = fmt.Sprintf("DepthProbe/Folder%d", i) } - _, err = albumRepo.GetSoleAlbumArtistIDsInSubtrees(*lib, paths...) + _, err = albumRepo.GetSoleAlbumArtistIDsInSubtrees(ctx, *lib, paths...) Expect(err).ToNot(HaveOccurred()) }) It("returns nothing when given no paths", func() { - lib, err := NewLibraryRepository(ctx, GetDBXBuilder()).Get(1) + lib, err := NewLibraryRepository(GetDBXBuilder()).Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) - ids, err := albumRepo.GetSoleAlbumArtistIDsInSubtrees(*lib) + ids, err := albumRepo.GetSoleAlbumArtistIDsInSubtrees(ctx, *lib) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(BeEmpty()) }) @@ -221,13 +222,13 @@ var _ = Describe("AlbumRepository", func() { Participants: model.Participants{model.RoleAlbumArtist: []model.Participant{{Artist: artistKraftwerk}}}} duo := model.Album{ID: "duo-artist-al", Name: "Duo", LibraryID: 1, AlbumArtistID: "999", Participants: model.Participants{model.RoleAlbumArtist: []model.Participant{{Artist: artistKraftwerk}, {Artist: artistBeatles}}}} - Expect(albumRepo.Put(&sole)).To(Succeed()) - Expect(albumRepo.Put(&duo)).To(Succeed()) + Expect(albumRepo.Put(ctx, &sole)).To(Succeed()) + Expect(albumRepo.Put(ctx, &duo)).To(Succeed()) DeferCleanup(func() { _, _ = GetDBXBuilder().NewQuery("DELETE FROM album WHERE id IN ('sole-artist-al', 'duo-artist-al')").Execute() }) - als, err := albumRepo.GetAll(model.QueryOptions{Filters: SoleAlbumArtistFilter("2")}) + als, err := albumRepo.GetAll(ctx, model.QueryOptions{Filters: SoleAlbumArtistFilter("2")}) Expect(err).ToNot(HaveOccurred()) Expect(als).To(HaveLen(1)) Expect(als[0].ID).To(Equal(sole.ID)) @@ -236,7 +237,7 @@ var _ = Describe("AlbumRepository", func() { Describe("GetAll", func() { var GetAll = func(opts ...model.QueryOptions) (model.Albums, error) { - albums, err := albumRepo.GetAll(opts...) + albums, err := albumRepo.GetAll(ctx, opts...) for i := range albums { albums[i].ImportedAt = time.Time{} } @@ -280,7 +281,7 @@ var _ = Describe("AlbumRepository", func() { Describe("recently_added sort", func() { AfterEach(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album"). + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album"). Where(squirrel.Like{"id": "ra-%"})) }) @@ -299,18 +300,18 @@ var _ = Describe("AlbumRepository", func() { // Same second, different nanoseconds: datetime() would tie these. earlier := &model.Album{LibraryID: 1, ID: "ra-earlier", Name: "Earlier"} later := &model.Album{LibraryID: 1, ID: "ra-later", Name: "Later"} - Expect(albumRepo.Put(earlier)).To(Succeed()) - Expect(albumRepo.Put(later)).To(Succeed()) - _, err := albumRepo.executeSQL(squirrel.Update("album"). + Expect(albumRepo.Put(ctx, earlier)).To(Succeed()) + Expect(albumRepo.Put(ctx, later)).To(Succeed()) + _, err := albumRepo.executeSQL(ctx, squirrel.Update("album"). Set("created_at", "2024-01-15 10:00:00.100000000+00:00"). Where(squirrel.Eq{"id": "ra-earlier"})) Expect(err).ToNot(HaveOccurred()) - _, err = albumRepo.executeSQL(squirrel.Update("album"). + _, err = albumRepo.executeSQL(ctx, squirrel.Update("album"). Set("created_at", "2024-01-15 10:00:00.900000000+00:00"). Where(squirrel.Eq{"id": "ra-later"})) Expect(err).ToNot(HaveOccurred()) - albums, err := albumRepo.GetAll(model.QueryOptions{Sort: "recently_added", Order: "desc"}) + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{Sort: "recently_added", Order: "desc"}) Expect(err).ToNot(HaveOccurred()) Expect(indexOf(albums, "ra-later")).To(BeNumerically("<", indexOf(albums, "ra-earlier")), ".900 should sort before .100 in desc order") @@ -321,17 +322,17 @@ var _ = Describe("AlbumRepository", func() { // match the unfiltered order (the inversion mechanism in #5673). ids := []string{"ra-t1", "ra-t2", "ra-t3", "ra-t4"} for _, aid := range ids { - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: aid, Name: aid})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: aid, Name: aid})).To(Succeed()) } - _, err := albumRepo.executeSQL(squirrel.Update("album"). + _, err := albumRepo.executeSQL(ctx, squirrel.Update("album"). Set("created_at", "2024-02-20 12:00:00+00:00"). Where(squirrel.Eq{"id": ids})) Expect(err).ToNot(HaveOccurred()) - all, err := albumRepo.GetAll(model.QueryOptions{Sort: "recently_added", Order: "desc"}) + all, err := albumRepo.GetAll(ctx, model.QueryOptions{Sort: "recently_added", Order: "desc"}) Expect(err).ToNot(HaveOccurred()) - subset, err := albumRepo.GetAll(model.QueryOptions{ + subset, err := albumRepo.GetAll(ctx, model.QueryOptions{ Sort: "recently_added", Order: "desc", Filters: squirrel.Eq{"album.id": []string{"ra-t1", "ra-t3"}}}) Expect(err).ToNot(HaveOccurred()) @@ -348,20 +349,20 @@ var _ = Describe("AlbumRepository", func() { BeforeEach(func() { // Create album without any annotation (no star, no rating) albumWithoutAnnotation = model.Album{ID: "no-annotation-album", Name: "No Annotation", LibraryID: 1} - Expect(albumRepo.Put(&albumWithoutAnnotation)).To(Succeed()) + Expect(albumRepo.Put(ctx, &albumWithoutAnnotation)).To(Succeed()) }) AfterEach(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": albumWithoutAnnotation.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": albumWithoutAnnotation.ID})) }) Describe("starred", func() { It("false includes items without annotations", func() { - res, err := albumRepo.ReadAll(rest.QueryOptions{ + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "false"}, }) Expect(err).ToNot(HaveOccurred()) - albums := res.(model.Albums) + albums := res var found bool for _, a := range albums { @@ -374,11 +375,11 @@ var _ = Describe("AlbumRepository", func() { }) It("true excludes items without annotations", func() { - res, err := albumRepo.ReadAll(rest.QueryOptions{ + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "true"}, }) Expect(err).ToNot(HaveOccurred()) - albums := res.(model.Albums) + albums := res for _, a := range albums { Expect(a.ID).ToNot(Equal(albumWithoutAnnotation.ID)) @@ -388,11 +389,11 @@ var _ = Describe("AlbumRepository", func() { Describe("has_rating", func() { It("false includes items without annotations", func() { - res, err := albumRepo.ReadAll(rest.QueryOptions{ + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"has_rating": "false"}, }) Expect(err).ToNot(HaveOccurred()) - albums := res.(model.Albums) + albums := res var found bool for _, a := range albums { @@ -405,11 +406,11 @@ var _ = Describe("AlbumRepository", func() { }) It("true excludes items without annotations", func() { - res, err := albumRepo.ReadAll(rest.QueryOptions{ + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"has_rating": "true"}, }) Expect(err).ToNot(HaveOccurred()) - albums := res.(model.Albums) + albums := res for _, a := range albums { Expect(a.ID).ToNot(Equal(albumWithoutAnnotation.ID)) @@ -425,12 +426,12 @@ var _ = Describe("AlbumRepository", func() { conf.Server.AlbumPlayCountMode = consts.AlbumPlayCountModeAbsolute newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "name", SongCount: songCount})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "name", SongCount: songCount})).To(Succeed()) for range playCount { - Expect(albumRepo.IncPlayCount(newID, time.Now())).To(Succeed()) + Expect(albumRepo.IncPlayCount(ctx, newID, time.Now())).To(Succeed()) } - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.PlayCount).To(Equal(int64(expected))) }, @@ -448,12 +449,12 @@ var _ = Describe("AlbumRepository", func() { conf.Server.AlbumPlayCountMode = consts.AlbumPlayCountModeNormalized newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "name", SongCount: songCount})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "name", SongCount: songCount})).To(Succeed()) for range playCount { - Expect(albumRepo.IncPlayCount(newID, time.Now())).To(Succeed()) + Expect(albumRepo.IncPlayCount(ctx, newID, time.Now())).To(Succeed()) } - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.PlayCount).To(Equal(int64(expected))) }, @@ -470,83 +471,79 @@ var _ = Describe("AlbumRepository", func() { Describe("Album.AverageRating", func() { It("returns 0 when no ratings exist", func() { newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "no ratings album"})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "no ratings album"})).To(Succeed()) - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.AverageRating).To(Equal(0.0)) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) }) It("returns the user's rating as average when only one user rated", func() { newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "single rating album"})).To(Succeed()) - Expect(albumRepo.SetRating(4, newID)).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "single rating album"})).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 4, newID)).To(Succeed()) - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.AverageRating).To(Equal(4.0)) - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) }) It("calculates average across multiple users", func() { newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "multi rating album"})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "multi rating album"})).To(Succeed()) - Expect(albumRepo.SetRating(4, newID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 4, newID)).To(Succeed()) user2Ctx := request.WithUser(GinkgoT().Context(), regularUser) - user2Repo := NewAlbumRepository(user2Ctx, GetDBXBuilder()).(*albumRepository) - Expect(user2Repo.SetRating(5, newID)).To(Succeed()) + Expect(albumRepo.SetRating(user2Ctx, 5, newID)).To(Succeed()) - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.AverageRating).To(Equal(4.5)) - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) }) It("excludes zero ratings from average calculation", func() { newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "zero rating excluded album"})).To(Succeed()) - Expect(albumRepo.SetRating(3, newID)).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "zero rating excluded album"})).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 3, newID)).To(Succeed()) user2Ctx := request.WithUser(GinkgoT().Context(), regularUser) - user2Repo := NewAlbumRepository(user2Ctx, GetDBXBuilder()).(*albumRepository) - Expect(user2Repo.SetRating(0, newID)).To(Succeed()) + Expect(albumRepo.SetRating(user2Ctx, 0, newID)).To(Succeed()) - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.AverageRating).To(Equal(3.0)) - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) }) It("rounds to 2 decimal places", func() { newID := id.NewRandom() - Expect(albumRepo.Put(&model.Album{LibraryID: 1, ID: newID, Name: "rounding test album"})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{LibraryID: 1, ID: newID, Name: "rounding test album"})).To(Succeed()) - Expect(albumRepo.SetRating(5, newID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 5, newID)).To(Succeed()) user2Ctx := request.WithUser(GinkgoT().Context(), regularUser) - user2Repo := NewAlbumRepository(user2Ctx, GetDBXBuilder()).(*albumRepository) - Expect(user2Repo.SetRating(4, newID)).To(Succeed()) + Expect(albumRepo.SetRating(user2Ctx, 4, newID)).To(Succeed()) user3Ctx := request.WithUser(GinkgoT().Context(), thirdUser) - user3Repo := NewAlbumRepository(user3Ctx, GetDBXBuilder()).(*albumRepository) - Expect(user3Repo.SetRating(4, newID)).To(Succeed()) + Expect(albumRepo.SetRating(user3Ctx, 4, newID)).To(Succeed()) - album, err := albumRepo.Get(newID) + album, err := albumRepo.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(album.AverageRating).To(Equal(4.33)) // (5 + 4 + 4) / 3 = 4.333... - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": newID})) }) }) @@ -715,10 +712,11 @@ var _ = Describe("AlbumRepository", func() { } var artistRepo *artistRepository + var artistCtx context.Context BeforeEach(func() { - ctx := request.WithUser(GinkgoT().Context(), adminUser) - artistRepo = NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository) + artistCtx = request.WithUser(ctx, adminUser) + artistRepo = NewArtistRepository(GetDBXBuilder()).(*artistRepository) }) // Helper to verify album_artists records @@ -730,7 +728,7 @@ var _ = Describe("AlbumRepository", func() { Where(squirrel.Eq{"album_id": albumID}). OrderBy("role", "artist_id", "sub_role") - err := albumRepo.queryAll(sq, &actual) + err := albumRepo.queryAll(ctx, sq, &actual) Expect(err).ToNot(HaveOccurred()) Expect(actual).To(Equal(expected)) } @@ -743,7 +741,7 @@ var _ = Describe("AlbumRepository", func() { OrderArtistName: "real artist", SortArtistName: "Artist, Real", } - err := createArtistWithLibrary(artistRepo, artist, 1) + err := createArtistWithLibrary(artistCtx, artistRepo, artist, 1) Expect(err).ToNot(HaveOccurred()) // Create an album with participants that reference the real artist @@ -764,7 +762,7 @@ var _ = Describe("AlbumRepository", func() { } // Insert the album - err = albumRepo.Put(album) + err = albumRepo.Put(ctx, album) Expect(err).ToNot(HaveOccurred()) // Verify that participant records were actually inserted into album_artists table @@ -775,13 +773,13 @@ var _ = Describe("AlbumRepository", func() { verifyAlbumArtists(album.ID, expected) // Clean up the test artist and album created for this test - _, _ = artistRepo.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = artistRepo.executeSQL(artistCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) It("finds albums through the participant-based filters", func() { artist := &model.Artist{ID: "filter-artist-1", Name: "Filter Artist", OrderArtistName: "filter artist"} - Expect(createArtistWithLibrary(artistRepo, artist, 1)).To(Succeed()) + Expect(createArtistWithLibrary(artistCtx, artistRepo, artist, 1)).To(Succeed()) album := &model.Album{ LibraryID: 1, @@ -794,36 +792,36 @@ var _ = Describe("AlbumRepository", func() { model.RoleComposer: {{Artist: model.Artist{ID: artist.ID, Name: artist.Name}}}, }, } - Expect(albumRepo.Put(album)).To(Succeed()) + Expect(albumRepo.Put(ctx, album)).To(Succeed()) - byArtist, err := albumRepo.GetAll(model.QueryOptions{Filters: artistFilter("artist_id", artist.ID)}) + byArtist, err := albumRepo.GetAll(ctx, model.QueryOptions{Filters: artistFilter("artist_id", artist.ID)}) Expect(err).ToNot(HaveOccurred()) Expect(byArtist).To(HaveLen(1)) Expect(byArtist[0].ID).To(Equal(album.ID)) - byComposer, err := albumRepo.GetAll(model.QueryOptions{Filters: artistRoleFilter("role_composer_id", artist.ID)}) + byComposer, err := albumRepo.GetAll(ctx, model.QueryOptions{Filters: artistRoleFilter("role_composer_id", artist.ID)}) Expect(err).ToNot(HaveOccurred()) Expect(byComposer).To(HaveLen(1)) - byLyricist, err := albumRepo.GetAll(model.QueryOptions{Filters: artistRoleFilter("role_lyricist_id", artist.ID)}) + byLyricist, err := albumRepo.GetAll(ctx, model.QueryOptions{Filters: artistRoleFilter("role_lyricist_id", artist.ID)}) Expect(err).ToNot(HaveOccurred()) Expect(byLyricist).To(BeEmpty()) - byAnyRole, err := albumRepo.GetAll(model.QueryOptions{Filters: allRolesFilter("role_total_id", artist.ID)}) + byAnyRole, err := albumRepo.GetAll(ctx, model.QueryOptions{Filters: allRolesFilter("role_total_id", artist.ID)}) Expect(err).ToNot(HaveOccurred()) Expect(byAnyRole).To(HaveLen(1)) - count, err := albumRepo.CountAll(model.QueryOptions{Filters: artistFilter("artist_id", artist.ID)}) + count, err := albumRepo.CountAll(ctx, model.QueryOptions{Filters: artistFilter("artist_id", artist.ID)}) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(1))) - _, _ = artistRepo.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = artistRepo.executeSQL(artistCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) It("clears album_artists rows when saved with empty participants", func() { artist := &model.Artist{ID: "clear-artist-1", Name: "Clear Artist", OrderArtistName: "clear artist"} - Expect(createArtistWithLibrary(artistRepo, artist, 1)).To(Succeed()) + Expect(createArtistWithLibrary(artistCtx, artistRepo, artist, 1)).To(Succeed()) album := &model.Album{ LibraryID: 1, @@ -836,14 +834,14 @@ var _ = Describe("AlbumRepository", func() { }, } DeferCleanup(func() { - _, _ = artistRepo.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = artistRepo.executeSQL(artistCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) - Expect(albumRepo.Put(album)).To(Succeed()) + Expect(albumRepo.Put(ctx, album)).To(Succeed()) verifyAlbumArtists(album.ID, []albumArtistRecord{{ArtistID: artist.ID, Role: "albumartist", SubRole: ""}}) album.Participants = model.Participants{} - Expect(albumRepo.Put(album)).To(Succeed()) + Expect(albumRepo.Put(ctx, album)).To(Succeed()) verifyAlbumArtists(album.ID, []albumArtistRecord{}) }) @@ -859,9 +857,9 @@ var _ = Describe("AlbumRepository", func() { Name: "Real Artist 2", OrderArtistName: "real artist 2", } - err := createArtistWithLibrary(artistRepo, artist1, 1) + err := createArtistWithLibrary(artistCtx, artistRepo, artist1, 1) Expect(err).ToNot(HaveOccurred()) - err = createArtistWithLibrary(artistRepo, artist2, 1) + err = createArtistWithLibrary(artistCtx, artistRepo, artist2, 1) Expect(err).ToNot(HaveOccurred()) // Create an album with mix of valid and invalid artist IDs @@ -885,7 +883,7 @@ var _ = Describe("AlbumRepository", func() { } // This should not fail - only valid artists should be inserted - err = albumRepo.Put(album) + err = albumRepo.Put(ctx, album) Expect(err).ToNot(HaveOccurred()) // Verify that only valid artist IDs were inserted into album_artists table @@ -899,8 +897,8 @@ var _ = Describe("AlbumRepository", func() { // Clean up the test artists and album created for this test artistIDs := []string{artist1.ID, artist2.ID} - _, _ = artistRepo.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": artistIDs})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = artistRepo.executeSQL(artistCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": artistIDs})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) It("handles complex nested JSON with multiple roles and sub-roles", func() { @@ -913,7 +911,7 @@ var _ = Describe("AlbumRepository", func() { } for _, artist := range artists { - err := createArtistWithLibrary(artistRepo, artist, 1) + err := createArtistWithLibrary(artistCtx, artistRepo, artist, 1) Expect(err).ToNot(HaveOccurred()) } @@ -940,7 +938,7 @@ var _ = Describe("AlbumRepository", func() { }, } - err := albumRepo.Put(album) + err := albumRepo.Put(ctx, album) Expect(err).ToNot(HaveOccurred()) // Verify complex JSON structure was correctly parsed and inserted @@ -959,8 +957,8 @@ var _ = Describe("AlbumRepository", func() { for i, artist := range artists { artistIDs[i] = artist.ID } - _, _ = artistRepo.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": artistIDs})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = artistRepo.executeSQL(artistCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": artistIDs})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) It("handles albums with non-existent artist IDs without constraint errors", func() { @@ -991,7 +989,7 @@ var _ = Describe("AlbumRepository", func() { // This should not fail with foreign key constraint error // The updateParticipants method should handle non-existent artist IDs gracefully - err := albumRepo.Put(album) + err := albumRepo.Put(ctx, album) Expect(err).ToNot(HaveOccurred()) // Verify that no participant records were inserted since all artist IDs were invalid @@ -999,7 +997,7 @@ var _ = Describe("AlbumRepository", func() { verifyAlbumArtists(album.ID, []albumArtistRecord{}) // Clean up the test album created for this test - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) It("removes stale role associations when artist role changes", func() { @@ -1013,7 +1011,7 @@ var _ = Describe("AlbumRepository", func() { Name: "Role Change Artist", OrderArtistName: "role change artist", } - err := createArtistWithLibrary(artistRepo, artist, 1) + err := createArtistWithLibrary(artistCtx, artistRepo, artist, 1) Expect(err).ToNot(HaveOccurred()) // Create album with artist as both albumartist and composer @@ -1033,7 +1031,7 @@ var _ = Describe("AlbumRepository", func() { }, } - err = albumRepo.Put(album) + err = albumRepo.Put(ctx, album) Expect(err).ToNot(HaveOccurred()) // Verify initial state: artist has both albumartist and composer roles @@ -1050,7 +1048,7 @@ var _ = Describe("AlbumRepository", func() { }, } - err = albumRepo.Put(album) + err = albumRepo.Put(ctx, album) Expect(err).ToNot(HaveOccurred()) // Verify that the albumartist role was removed - only composer should remain @@ -1062,14 +1060,14 @@ var _ = Describe("AlbumRepository", func() { verifyAlbumArtists(album.ID, expectedAfter) // Clean up - _, _ = artistRepo.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) + _, _ = artistRepo.executeSQL(artistCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": artist.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": album.ID})) }) }) Describe("GetYears", func() { It("returns distinct album years ascending, excluding zero", func() { - years, err := albumRepo.GetYears() + years, err := albumRepo.GetYears(ctx) Expect(err).ToNot(HaveOccurred()) // Sorted ascending, no duplicates, no zero-year entries. Expect(sort.IsSorted(sort.IntSlice(years))).To(BeTrue()) @@ -1084,13 +1082,13 @@ var _ = Describe("AlbumRepository", func() { // Insert two albums with the same non-zero max_year (2005). album1 := &model.Album{LibraryID: 1, ID: "dedup-test-1", Name: "Album 1", MaxYear: 2005} album2 := &model.Album{LibraryID: 1, ID: "dedup-test-2", Name: "Album 2", MaxYear: 2005} - Expect(albumRepo.Put(album1)).To(Succeed()) - Expect(albumRepo.Put(album2)).To(Succeed()) + Expect(albumRepo.Put(ctx, album1)).To(Succeed()) + Expect(albumRepo.Put(ctx, album2)).To(Succeed()) DeferCleanup(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"dedup-test-1", "dedup-test-2"}})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"dedup-test-1", "dedup-test-2"}})) }) - years, err := albumRepo.GetYears() + years, err := albumRepo.GetYears(ctx) Expect(err).ToNot(HaveOccurred()) // Count occurrences of 2005 in the result @@ -1104,10 +1102,10 @@ var _ = Describe("AlbumRepository", func() { }) It("scopes years to the given libraries", func() { - all, err := albumRepo.GetYears() + all, err := albumRepo.GetYears(ctx) Expect(err).ToNot(HaveOccurred()) // A library with no albums yields no years. - scoped, err := albumRepo.GetYears(99999) + scoped, err := albumRepo.GetYears(ctx, 99999) Expect(err).ToNot(HaveOccurred()) Expect(scoped).To(BeEmpty()) Expect(all).ToNot(BeEmpty()) @@ -1115,12 +1113,12 @@ var _ = Describe("AlbumRepository", func() { It("excludes years that belong only to missing albums", func() { gone := &model.Album{LibraryID: 1, ID: "missing-year-1", Name: "Gone", MaxYear: 1911, Missing: true} - Expect(albumRepo.Put(gone)).To(Succeed()) + Expect(albumRepo.Put(ctx, gone)).To(Succeed()) DeferCleanup(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": "missing-year-1"})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": "missing-year-1"})) }) - years, err := albumRepo.GetYears() + years, err := albumRepo.GetYears(ctx) Expect(err).ToNot(HaveOccurred()) Expect(years).ToNot(ContainElement(1911)) }) @@ -1169,15 +1167,15 @@ var _ = Describe("AlbumRepository", func() { Describe("ReplayGain", func() { BeforeEach(func() { DeferCleanup(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"rg-1", "rg-2"}})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"rg-1", "rg-2"}})) }) }) It("round-trips album ReplayGain gain and peak", func() { - Expect(albumRepo.Put(&model.Album{ + Expect(albumRepo.Put(ctx, &model.Album{ ID: "rg-1", Name: "rg", LibraryID: 1, RGAlbumGain: new(-7.5), RGAlbumPeak: new(0.98), })).To(Succeed()) - got, err := albumRepo.Get("rg-1") + got, err := albumRepo.Get(ctx, "rg-1") Expect(err).ToNot(HaveOccurred()) Expect(got.RGAlbumGain).ToNot(BeNil()) Expect(*got.RGAlbumGain).To(Equal(-7.5)) @@ -1185,8 +1183,8 @@ var _ = Describe("AlbumRepository", func() { Expect(*got.RGAlbumPeak).To(Equal(0.98)) }) It("reads nil when ReplayGain is unset", func() { - Expect(albumRepo.Put(&model.Album{ID: "rg-2", Name: "rg2", LibraryID: 1})).To(Succeed()) - got, err := albumRepo.Get("rg-2") + Expect(albumRepo.Put(ctx, &model.Album{ID: "rg-2", Name: "rg2", LibraryID: 1})).To(Succeed()) + got, err := albumRepo.Get(ctx, "rg-2") Expect(err).ToNot(HaveOccurred()) Expect(got.RGAlbumGain).To(BeNil()) Expect(got.RGAlbumPeak).To(BeNil()) @@ -1196,16 +1194,46 @@ var _ = Describe("AlbumRepository", func() { // Exists must apply the same library filter as Get/GetAll/CountAll. Describe("Exists library visibility", func() { It("hides an album the user has no library access to", func() { - Expect(albumRepo.Put(&model.Album{ID: "vis-album", Name: "Vis", LibraryID: 1})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{ID: "vis-album", Name: "Vis", LibraryID: 1})).To(Succeed()) DeferCleanup(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": "vis-album"})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": "vis-album"})) }) - Expect(albumRepo.Exists("vis-album")).To(BeTrue(), "admin sees it") + Expect(albumRepo.Exists(ctx, "vis-album")).To(BeTrue(), "admin sees it") restricted := model.User{ID: "restricted_album_user", UserName: "ra", Name: "RA", Email: "ra@t.com"} rctx := request.WithUser(GinkgoT().Context(), restricted) - Expect(NewAlbumRepository(rctx, GetDBXBuilder()).Exists("vis-album")).To(BeFalse()) + Expect(albumRepo.Exists(rctx, "vis-album")).To(BeFalse()) + }) + + It("keeps per-user library visibility separate on a shared repository", func() { + adminCtx := request.WithUser(ctx, adminUser) + adminCount, err := albumRepo.CountAll(adminCtx) + Expect(err).ToNot(HaveOccurred()) + Expect(adminCount).To(BeNumerically(">", 0)) + + // A user with no library grants, so its visibility can't drift with other specs + restrictedCtx := request.WithUser(ctx, model.User{ID: "shared-repo-restricted"}) + restrictedCount, err := albumRepo.CountAll(restrictedCtx) + Expect(err).ToNot(HaveOccurred()) + Expect(restrictedCount).To(BeZero()) + + var wg sync.WaitGroup + for i := range 20 { + wg.Add(1) + go func(i int) { + defer GinkgoRecover() + defer wg.Done() + c, want := adminCtx, adminCount + if i%2 == 1 { + c, want = restrictedCtx, restrictedCount + } + got, err := albumRepo.CountAll(c) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(want)) + }(i) + } + wg.Wait() }) }) }) diff --git a/persistence/artist_repository.go b/persistence/artist_repository.go index 97c1452cb..b90a48c48 100644 --- a/persistence/artist_repository.go +++ b/persistence/artist_repository.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "fmt" + "maps" "os" "slices" "strings" @@ -129,9 +130,8 @@ func (dba dbArtists) toModels() model.Artists { return res } -func NewArtistRepository(ctx context.Context, db dbx.Builder) model.ArtistRepository { +func NewArtistRepository(db dbx.Builder) model.ArtistRepository { r := &artistRepository{} - r.ctx = ctx r.db = db r.indexGroups = utils.ParseIndexGroups(conf.Server.IndexGroups) r.tableName = "artist" // To be used by the idFilter below @@ -188,8 +188,8 @@ func artistLibraryIdFilter(_ string, value any) Sqlizer { } // applyLibraryFilterToArtistQuery applies library filtering to artist queries through the library_artist junction table -func (r *artistRepository) applyLibraryFilterToArtistQuery(query SelectBuilder) SelectBuilder { - user := loggedUser(r.ctx) +func (r *artistRepository) applyLibraryFilterToArtistQuery(ctx context.Context, query SelectBuilder) SelectBuilder { + user := loggedUser(ctx) // Join with library_artist first to ensure only artists with content in libraries are included // Exclude artists with empty stats (no actual content in the library) query = query.Join("library_artist on library_artist.artist_id = artist.id") @@ -204,106 +204,106 @@ func (r *artistRepository) applyLibraryFilterToArtistQuery(query SelectBuilder) return query } -func (r *artistRepository) selectArtist(options ...model.QueryOptions) SelectBuilder { +func (r *artistRepository) selectArtist(ctx context.Context, options ...model.QueryOptions) SelectBuilder { // Stats Format: {"1": {"albumartist": {"m": 10, "a": 5, "s": 1024}, "artist": {...}}, "2": {...}} - query := r.newSelect(options...).Columns("artist.*", + query := r.newSelect(ctx, options...).Columns("artist.*", "JSON_GROUP_OBJECT(library_artist.library_id, JSONB(library_artist.stats)) as library_stats_json") - query = r.applyLibraryFilterToArtistQuery(query) + query = r.applyLibraryFilterToArtistQuery(ctx, query) query = query.GroupBy("artist.id") - return r.withAnnotation(query, "artist.id") + return r.withAnnotation(ctx, query, "artist.id") } -func (r *artistRepository) CountAll(options ...model.QueryOptions) (int64, error) { - query := r.newSelect() - query = r.applyLibraryFilterToArtistQuery(query) +func (r *artistRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + query := r.newSelect(ctx) + query = r.applyLibraryFilterToArtistQuery(ctx, query) // Only the annotation join is gated; the library_artist join above (and its count(distinct)) // must stay, since an artist can span multiple libraries. if filtersNeedAnnotation(r.applyFilters(query, options...)) { - query = r.withAnnotation(query, "artist.id") + query = r.withAnnotation(ctx, query, "artist.id") } - return r.count(query, options...) + return r.count(ctx, query, options...) } // Exists checks if an artist with the given ID exists in the database and is accessible by the current user. -func (r *artistRepository) Exists(id string) (bool, error) { +func (r *artistRepository) Exists(ctx context.Context, id string) (bool, error) { // Create a query using the same library filtering logic as selectArtist() - query := r.newSelect().Columns("count(distinct artist.id) as exist").Where(Eq{"artist.id": id}) - query = r.applyLibraryFilterToArtistQuery(query) + query := r.newSelect(ctx).Columns("count(distinct artist.id) as exist").Where(Eq{"artist.id": id}) + query = r.applyLibraryFilterToArtistQuery(ctx, query) var res struct{ Exist int64 } - err := r.queryOne(query, &res) + err := r.queryOne(ctx, query, &res) return res.Exist > 0, err } -func (r *artistRepository) Put(a *model.Artist, colsToUpdate ...string) error { +func (r *artistRepository) Put(ctx context.Context, a *model.Artist, colsToUpdate ...string) error { dba := &dbArtist{Artist: a} dba.CreatedAt = new(time.Now()) dba.UpdatedAt = dba.CreatedAt - _, err := r.put(dba.ID, dba, colsToUpdate...) + _, err := r.put(ctx, dba.ID, dba, colsToUpdate...) return err } -func (r *artistRepository) UpdateExternalInfo(a *model.Artist) error { +func (r *artistRepository) UpdateExternalInfo(ctx context.Context, a *model.Artist) error { dba := &dbArtist{Artist: a} - _, err := r.put(a.ID, dba, + _, err := r.put(ctx, a.ID, dba, "biography", "small_image_url", "medium_image_url", "large_image_url", "similar_artists", "external_url", "external_info_updated_at") return err } -func (r *artistRepository) Get(id string) (*model.Artist, error) { - sel := r.selectArtist().Where(Eq{"artist.id": id}) +func (r *artistRepository) Get(ctx context.Context, id string) (*model.Artist, error) { + sel := r.selectArtist(ctx).Where(Eq{"artist.id": id}) var dba dbArtists - if err := r.queryAll(sel, &dba); err != nil { + if err := r.queryAll(ctx, sel, &dba); err != nil { return nil, err } if len(dba) == 0 { return nil, model.ErrNotFound } res := dba.toModels() - r.hydrateArtwork(res) + r.hydrateArtwork(ctx, res) return &res[0], nil } -func (r *artistRepository) GetAll(options ...model.QueryOptions) (model.Artists, error) { - sel := r.selectArtist(options...) +func (r *artistRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Artists, error) { + sel := r.selectArtist(ctx, options...) var dba dbArtists - err := r.queryAll(sel, &dba) + err := r.queryAll(ctx, sel, &dba) if err != nil { return nil, err } res := dba.toModels() - r.hydrateArtwork(res) + r.hydrateArtwork(ctx, res) return res, err } // getAllIDs returns just the artist IDs for the same row set as GetAll, skipping the // heavy stats columns and JSON post-processing. -func (r *artistRepository) getAllIDs(options ...model.QueryOptions) ([]string, error) { - sq := r.applyLibraryFilterToArtistQuery(r.newSelect(options...).Columns("artist.id")).GroupBy("artist.id") +func (r *artistRepository) getAllIDs(ctx context.Context, options ...model.QueryOptions) ([]string, error) { + sq := r.applyLibraryFilterToArtistQuery(ctx, r.newSelect(ctx, options...).Columns("artist.id")).GroupBy("artist.id") if filtersNeedAnnotation(sq) { - sq = r.withAnnotation(sq, "artist.id") + sq = r.withAnnotation(ctx, sq, "artist.id") } ids := []string{} - err := r.queryAllSlice(sq, &ids) + err := r.queryAllSlice(ctx, sq, &ids) return ids, err } // hydrateArtwork fills each artist's ImageHash/ImageAbsent from one batched item_artwork lookup. -func (r *artistRepository) hydrateArtwork(artists model.Artists) { - hydrateItems(r.ctx, r.db, model.KindArtistArtwork, artists, +func (r *artistRepository) hydrateArtwork(ctx context.Context, artists model.Artists) { + hydrateItems(ctx, r.db, model.KindArtistArtwork, artists, func(a *model.Artist) (string, *model.ItemImage) { return a.ID, &a.ItemImage }) } -func (r *artistRepository) GetCursor(options ...model.QueryOptions) (model.ArtistCursor, error) { - ids, err := r.getAllIDs(options...) +func (r *artistRepository) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.ArtistCursor, error) { + ids, err := r.getAllIDs(ctx, options...) if err != nil { return nil, err } opts := chunkOptions(options, "artist.id") return model.ArtistCursor(streamByIDs(ids, func(chunk []string) (model.Artists, error) { - return r.GetAll(opts(chunk)) + return r.GetAll(ctx, opts(chunk)) })), nil } @@ -324,7 +324,7 @@ func (r *artistRepository) getIndexKey(a model.Artist) string { // GetIndex returns a list of artists grouped by the first letter of their name, or by the index group if configured. // It can filter by roles and libraries, and optionally include artists that are missing (i.e., have no albums). // TODO Cache the index (recalculate at scan time) -func (r *artistRepository) GetIndex(includeMissing bool, libraryIds []int, roles ...model.Role) (model.ArtistIndexes, error) { +func (r *artistRepository) GetIndex(ctx context.Context, includeMissing bool, libraryIds []int, roles ...model.Role) (model.ArtistIndexes, error) { // Validate library IDs. If no library IDs are provided, return an empty index. if len(libraryIds) == 0 { return nil, nil @@ -352,7 +352,7 @@ func (r *artistRepository) GetIndex(includeMissing bool, libraryIds []int, roles options.Filters = And{options.Filters, libFilter} } - artists, err := r.GetAll(options) + artists, err := r.GetAll(ctx, options) if err != nil { return nil, err } @@ -367,7 +367,7 @@ func (r *artistRepository) GetIndex(includeMissing bool, libraryIds []int, roles return result, nil } -func (r *artistRepository) purgeEmpty() error { +func (r *artistRepository) purgeEmpty(ctx context.Context) error { orphanFilter := "id not in (select artist_id from album_artists)" // Collect uploaded image filenames before deleting @@ -375,18 +375,18 @@ func (r *artistRepository) purgeEmpty() error { Where(orphanFilter). Where("uploaded_image != ''") var imageFiles []string - if err := r.queryAllSlice(sel, &imageFiles); err != nil && !errors.Is(err, model.ErrNotFound) { + if err := r.queryAllSlice(ctx, sel, &imageFiles); err != nil && !errors.Is(err, model.ErrNotFound) { return fmt.Errorf("collecting artist images for cleanup: %w", err) } // Delete orphan artists del := Delete(r.tableName).Where(orphanFilter) - c, err := r.executeSQL(del) + c, err := r.executeSQL(ctx, del) if err != nil { return fmt.Errorf("purging empty artists: %w", err) } if c > 0 { - log.Debug(r.ctx, "Purged empty artists", "totalDeleted", c) + log.Debug(ctx, "Purged empty artists", "totalDeleted", c) } if len(imageFiles) == 0 { @@ -394,11 +394,11 @@ func (r *artistRepository) purgeEmpty() error { } // Best-effort cleanup of uploaded image files - log.Debug(r.ctx, "Cleaning up artist images", "totalImages", len(imageFiles)) + log.Debug(ctx, "Cleaning up artist images", "totalImages", len(imageFiles)) for _, filename := range imageFiles { path := model.UploadedImagePath(consts.EntityArtist, filename) if err := os.Remove(path); err != nil && !os.IsNotExist(err) { - log.Warn(r.ctx, "Failed to remove artist image during GC", "path", path, err) + log.Warn(ctx, "Failed to remove artist image during GC", "path", path, err) } } return nil @@ -407,9 +407,9 @@ func (r *artistRepository) purgeEmpty() error { // markOrphansMissing flags as missing any non-missing artist with no library_artist row, keeping the // search fast-path's `missing = false` filter correct (see searchCfg). Called wherever such a row can // be dropped: RefreshStats cleanup and library deletion cascade. -func (r *artistRepository) markOrphansMissing() error { - _, err := r.executeSQL(Expr( - "update artist set missing = true where missing = false " + +func (r *artistRepository) markOrphansMissing(ctx context.Context) error { + _, err := r.executeSQL(ctx, Expr( + "update artist set missing = true where missing = false "+ "and not exists (select 1 from library_artist where library_artist.artist_id = artist.id)")) if err != nil { return fmt.Errorf("marking orphaned artists missing: %w", err) @@ -418,7 +418,7 @@ func (r *artistRepository) markOrphansMissing() error { } // markMissing marks artists as missing if all their albums are missing. -func (r *artistRepository) markMissing() error { +func (r *artistRepository) markMissing(ctx context.Context) error { q := Expr(` with artists_with_non_missing_albums as ( select distinct aa.artist_id @@ -429,7 +429,7 @@ with artists_with_non_missing_albums as ( update artist set missing = (artist.id not in (select artist_id from artists_with_non_missing_albums)); `) - _, err := r.executeSQL(q) + _, err := r.executeSQL(ctx, q) if err != nil { return fmt.Errorf("marking missing artists: %w", err) } @@ -438,7 +438,7 @@ set missing = (artist.id not in (select artist_id from artists_with_non_missing_ // RefreshPlayCounts updates the play count and last play date annotations for all artists, based // on the media files associated with them. -func (r *artistRepository) RefreshPlayCounts() (int64, error) { +func (r *artistRepository) RefreshPlayCounts(ctx context.Context) (int64, error) { query := Expr(` with play_counts as ( select user_id, atom as artist_id, sum(play_count) as total_play_count, max(play_date) as last_play_date @@ -456,13 +456,13 @@ on conflict (user_id, item_id, item_type) do update set play_count = excluded.play_count, play_date = excluded.play_date; `) - return r.executeSQL(query) + return r.executeSQL(ctx, query) } // RefreshStats updates the stats field for artists whose associated media files were updated after the oldest recorded library scan time. // When allArtists is true, it refreshes stats for all artists. It processes artists in batches to handle potentially large updates. // This method now calculates per-library statistics and stores them in the library_artist junction table. -func (r *artistRepository) RefreshStats(allArtists bool) (int64, error) { +func (r *artistRepository) RefreshStats(ctx context.Context, allArtists bool) (int64, error) { var allTouchedArtistIDs []string if allArtists { // Refresh stats for all artists @@ -470,7 +470,7 @@ func (r *artistRepository) RefreshStats(allArtists bool) (int64, error) { if err := r.db.NewQuery(allArtistsQuerySQL).Column(&allTouchedArtistIDs); err != nil { return 0, fmt.Errorf("fetching all artist IDs: %w", err) } - log.Debug(r.ctx, "RefreshStats: Refreshing all artists.", "count", len(allTouchedArtistIDs)) + log.Debug(ctx, "RefreshStats: Refreshing all artists.", "count", len(allTouchedArtistIDs)) } else { // Only refresh artists with updated timestamps touchedArtistsQuerySQL := ` @@ -481,11 +481,11 @@ func (r *artistRepository) RefreshStats(allArtists bool) (int64, error) { if err := r.db.NewQuery(touchedArtistsQuerySQL).Column(&allTouchedArtistIDs); err != nil { return 0, fmt.Errorf("fetching touched artist IDs: %w", err) } - log.Debug(r.ctx, "RefreshStats: Refreshing touched artists.", "count", len(allTouchedArtistIDs)) + log.Debug(ctx, "RefreshStats: Refreshing touched artists.", "count", len(allTouchedArtistIDs)) } if len(allTouchedArtistIDs) == 0 { - log.Debug(r.ctx, "RefreshStats: No artists to update.") + log.Debug(ctx, "RefreshStats: No artists to update.") return 0, nil } @@ -558,7 +558,7 @@ func (r *artistRepository) RefreshStats(allArtists bool) (int64, error) { batchCounter := 0 for artistIDBatch := range slice.CollectChunks(slices.Values(allTouchedArtistIDs), batchSize) { batchCounter++ - log.Trace(r.ctx, "RefreshStats: Processing batch", "batchNum", batchCounter, "batchSize", len(artistIDBatch)) + log.Trace(ctx, "RefreshStats: Processing batch", "batchNum", batchCounter, "batchSize", len(artistIDBatch)) // Create placeholders for each ID in the IN clauses placeholders := make([]string, len(artistIDBatch)) @@ -584,7 +584,7 @@ func (r *artistRepository) RefreshStats(allArtists bool) (int64, error) { // Now use Expr with the expanded SQL and all parameters sqlizer := Expr(batchSQL, args...) - rowsAffected, err := r.executeSQL(sqlizer) + rowsAffected, err := r.executeSQL(ctx, sqlizer) if err != nil { return totalRowsAffected, fmt.Errorf("executing batch update for artist stats (batch %d): %w", batchCounter, err) } @@ -593,23 +593,23 @@ func (r *artistRepository) RefreshStats(allArtists bool) (int64, error) { // Remove library_artist entries for artists that no longer have any content in a library. cleanupSQL := Delete("library_artist").Where("stats = '{}'") - cleanupRows, err := r.executeSQL(cleanupSQL) + cleanupRows, err := r.executeSQL(ctx, cleanupSQL) if err != nil { - log.Warn(r.ctx, "Failed to cleanup empty library_artist entries", err) + log.Warn(ctx, "Failed to cleanup empty library_artist entries", err) } else { if cleanupRows > 0 { - log.Debug(r.ctx, "Cleaned up empty library_artist entries", "rowsDeleted", cleanupRows) + log.Debug(ctx, "Cleaned up empty library_artist entries", "rowsDeleted", cleanupRows) } // Reconcile orphans whenever the cleanup removed rows, and on a full refresh so a full scan // also heals any left by older versions. if cleanupRows > 0 || allArtists { - if err := r.markOrphansMissing(); err != nil { - log.Warn(r.ctx, "Failed to mark orphaned artists missing after library_artist cleanup", err) + if err := r.markOrphansMissing(ctx); err != nil { + log.Warn(ctx, "Failed to mark orphaned artists missing after library_artist cleanup", err) } } } - log.Debug(r.ctx, "RefreshStats: Successfully updated stats.", "totalArtistsProcessed", len(allTouchedArtistIDs), "totalDBRowsAffected", totalRowsAffected) + log.Debug(ctx, "RefreshStats: Successfully updated stats.", "totalArtistsProcessed", len(allTouchedArtistIDs), "totalDBRowsAffected", totalRowsAffected) return totalRowsAffected, nil } @@ -649,39 +649,39 @@ func artistLibraryFilter(libraryIDs []int) Sqlizer { return Expr("EXISTS ("+sub+")", args...) } -func (r *artistRepository) Search(q string, options ...model.QueryOptions) (model.Artists, error) { +func (r *artistRepository) Search(ctx context.Context, q string, options ...model.QueryOptions) (model.Artists, error) { var opts model.QueryOptions if len(options) > 0 { opts = options[0] } // Artists have no library_id column, so the library_id filter callers pass (same as albums/songs) // can't be applied directly: consume it and realize it as a join-free Phase-1 scope (searchCfg). - scope := r.searchScope(opts.Filters) + scope := r.searchScope(ctx, opts.Filters) if isLibraryIDFilter(opts.Filters) { opts.Filters = nil } var res dbArtists - err := r.doSearch(r.selectArtist(opts), q, &res, r.searchCfg(scope), opts) + err := r.doSearch(ctx, r.selectArtist(ctx, opts), q, &res, r.searchCfg(scope), opts) if err != nil { return nil, fmt.Errorf("searching artist %q: %w", q, err) } artists := res.toModels() - r.hydrateArtwork(artists) + r.hydrateArtwork(ctx, artists) return artists, nil } // searchScope returns the library IDs the search must be restricted to, or nil to skip the filter // entirely (the fast-path: the user sees everything the search could return, so a filter would be // pure O(offset) overhead). It intersects the requested libraries with what the user can see. -func (r *artistRepository) searchScope(filter Sqlizer) []int { - visible, err := r.visibleLibraryIDs() +func (r *artistRepository) searchScope(ctx context.Context, filter Sqlizer) []int { + visible, err := r.visibleLibraryIDs(ctx) if err != nil { return r.requestedLibraryIDs(filter) // fail safe: narrow to the request rather than widen } requested := r.requestedLibraryIDs(filter) if requested == nil { // No explicit request: scope to the visible set, unless the user sees everything. - if r.userSeesAllLibraries(visible) { + if r.userSeesAllLibraries(ctx, visible) { return nil } return visible @@ -717,15 +717,17 @@ func isLibraryIDFilter(filter Sqlizer) bool { return ok } -func (r *artistRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *artistRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *artistRepository) Read(id string) (any, error) { - return r.Get(id) +func (r *artistRepository) Read(ctx context.Context, id string) (*model.Artist, error) { + return r.Get(ctx, id) } -func (r *artistRepository) ReadAll(options ...rest.QueryOptions) (any, error) { +// sortMappingsForRole copies the shared mappings so a role-specific sort never leaks into them. +// Anything sanitizeArtistStatsRole rejects falls back to the "total" aggregate. +func (r *artistRepository) sortMappingsForRole(options ...rest.QueryOptions) map[string]string { role := "total" if len(options) > 0 { if v, ok := options[0].Filters["role"].(string); ok { @@ -734,19 +736,18 @@ func (r *artistRepository) ReadAll(options ...rest.QueryOptions) (any, error) { } } } - r.sortMappings["song_count"] = "sum(stats->>'" + role + "'->>'m')" - r.sortMappings["album_count"] = "sum(stats->>'" + role + "'->>'a')" - r.sortMappings["size"] = "sum(stats->>'" + role + "'->>'s')" - return r.GetAll(r.parseRestOptions(r.ctx, options...)) + mappings := maps.Clone(r.sortMappings) + mappings["song_count"] = "sum(stats->>'" + role + "'->>'m')" + mappings["album_count"] = "sum(stats->>'" + role + "'->>'a')" + mappings["size"] = "sum(stats->>'" + role + "'->>'s')" + return mappings } -func (r *artistRepository) EntityName() string { - return "artist" -} - -func (r *artistRepository) NewInstance() any { - return &model.Artist{} +func (r *artistRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Artist, error) { + scoped := *r + scoped.sortMappings = r.sortMappingsForRole(options...) + return scoped.GetAll(ctx, scoped.parseRestOptions(ctx, options...)) } var _ model.ArtistRepository = (*artistRepository)(nil) -var _ model.ResourceRepository = (*artistRepository)(nil) +var _ rest.Repository[model.Artist] = (*artistRepository)(nil) diff --git a/persistence/artist_repository_test.go b/persistence/artist_repository_test.go index f6612acc4..914c2f89a 100644 --- a/persistence/artist_repository_test.go +++ b/persistence/artist_repository_test.go @@ -6,6 +6,7 @@ import ( "fmt" "os" "path/filepath" + "sync" "github.com/Masterminds/squirrel" "github.com/deluan/rest" @@ -49,6 +50,11 @@ func createUserWithLibraries(userID string, libraryIDs []int) model.User { } var _ = Describe("ArtistRepository", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) Context("Core Functionality", func() { Describe("GetIndexKey", func() { @@ -130,31 +136,70 @@ var _ = Describe("ArtistRepository", func() { }) }) - Describe("ReadAll role sort SQL injection", func() { - It("does not interpolate attacker-controlled role into ORDER BY", func() { - ctx := request.WithUser(GinkgoT().Context(), adminUser) - repo := NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository) - payload := "total') OR 1=1--" - _, err := repo.ReadAll(rest.QueryOptions{ - Sort: "songCount", - Order: "ASC", - Filters: map[string]any{"role": payload}, - }) - Expect(err).ToNot(HaveOccurred()) - Expect(repo.sortMappings["song_count"]).To(Equal("sum(stats->>'total'->>'m')")) - Expect(repo.sortMappings["song_count"]).ToNot(ContainSubstring(payload)) + Describe("ReadAll role sort", func() { + payload := "total') OR 1=1--" + + songCountSortFor := func(role any) string { + repo := NewArtistRepository(GetDBXBuilder()).(*artistRepository) + return repo.sortMappingsForRole(rest.QueryOptions{Filters: map[string]any{"role": role}})["song_count"] + } + + It("falls back to the total stats role for an attacker-controlled role", func() { + Expect(songCountSortFor(payload)).To(Equal("sum(stats->>'total'->>'m')")) + Expect(songCountSortFor("bogus")).To(Equal("sum(stats->>'total'->>'m')")) + Expect(songCountSortFor(42)).To(Equal("sum(stats->>'total'->>'m')")) }) It("keeps valid role sort paths", func() { + Expect(songCountSortFor("composer")).To(Equal("sum(stats->>'composer'->>'m')")) + Expect(songCountSortFor("albumartist")).To(Equal("sum(stats->>'albumartist'->>'m')")) + }) + + It("leaves the shared mappings untouched", func() { + repo := NewArtistRepository(GetDBXBuilder()).(*artistRepository) + Expect(repo.sortMappingsForRole(rest.QueryOptions{Filters: map[string]any{"role": "composer"}})). + ToNot(Equal(repo.sortMappings)) + Expect(repo.sortMappings["song_count"]).To(Equal("stats->>'total'->>'m'")) + }) + + It("orders by the requested role's stats, not the total", func() { ctx := request.WithUser(GinkgoT().Context(), adminUser) - repo := NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository) - _, err := repo.ReadAll(rest.QueryOptions{ + repo := NewArtistRepository(GetDBXBuilder()).(*artistRepository) + // Composer and total counts rank the two artists in opposite orders, so the + // resulting order alone proves which mapping the sort used. + seed := func(artistID, stats string) { + _, err := repo.executeSQL(ctx, squirrel.Insert("library_artist"). + Columns("library_id", "artist_id", "stats"). + Values(1, artistID, stats). + Suffix("ON CONFLICT(library_id, artist_id) DO UPDATE SET stats = excluded.stats")) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { + _, _ = repo.executeSQL(ctx, squirrel.Update("library_artist").Set("stats", "{}"). + Where(squirrel.Eq{"library_id": 1, "artist_id": artistID})) + }) + } + seed(artistBeatles.ID, `{"composer": {"s": 1, "m": 1, "a": 1}, "total": {"s": 1, "m": 100, "a": 1}}`) + seed(artistKraftwerk.ID, `{"composer": {"s": 1, "m": 9, "a": 1}, "total": {"s": 1, "m": 2, "a": 1}}`) + + res, err := repo.ReadAll(ctx, rest.QueryOptions{ Sort: "songCount", Order: "DESC", Filters: map[string]any{"role": "composer"}, }) Expect(err).ToNot(HaveOccurred()) - Expect(repo.sortMappings["song_count"]).To(Equal("sum(stats->>'composer'->>'m')")) + Expect(slice.Map(res, func(a model.Artist) string { return a.ID })). + To(Equal([]string{artistKraftwerk.ID, artistBeatles.ID})) + }) + + It("still returns results when the role is an injection payload", func() { + ctx := request.WithUser(GinkgoT().Context(), adminUser) + repo := NewArtistRepository(GetDBXBuilder()).(*artistRepository) + _, err := repo.ReadAll(ctx, rest.QueryOptions{ + Sort: "songCount", + Order: "ASC", + Filters: map[string]any{"role": payload}, + }) + Expect(err).ToNot(HaveOccurred()) }) }) @@ -163,8 +208,8 @@ var _ = Describe("ArtistRepository", func() { // the way Search() does, for a repo whose context carries the given user. scope := func(user model.User, filter squirrel.Sqlizer) []int { ctx := request.WithUser(GinkgoT().Context(), user) - r := NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository) - return r.searchScope(filter) + r := NewArtistRepository(GetDBXBuilder()).(*artistRepository) + return r.searchScope(ctx, filter) } subsetUser := model.User{ID: "u", Libraries: model.Libraries{{ID: 1}, {ID: 2}, {ID: 3}}} @@ -185,7 +230,7 @@ var _ = Describe("ArtistRepository", func() { // A restricted user (strictly fewer libs than exist) with no musicFolderId is still // confined to their granted libs. Build the user with total-1 libraries derived from // the real DB total, so the "sees all" fast-path can't kick in regardless of count. - total, err := NewLibraryRepository(GinkgoT().Context(), GetDBXBuilder()).CountAll() + total, err := NewLibraryRepository(GetDBXBuilder()).CountAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(total).To(BeNumerically(">", 0)) libs := make(model.Libraries, 0, total-1) @@ -201,8 +246,8 @@ var _ = Describe("ArtistRepository", func() { // Admins see every library, so the visible set is the whole library table — derive // it from the DB rather than assuming a count. var allLibs []int - Expect(NewLibraryRepository(GinkgoT().Context(), GetDBXBuilder()).(*libraryRepository). - queryAllSlice(squirrel.Select("id").From("library"), &allLibs)).To(Succeed()) + Expect(NewLibraryRepository(GetDBXBuilder()).(*libraryRepository). + queryAllSlice(ctx, squirrel.Select("id").From("library"), &allLibs)).To(Succeed()) admin := model.User{ID: "a", IsAdmin: true} Expect(scope(admin, squirrel.Eq{"library_id": allLibs})).To(BeNil()) Expect(scope(admin, nil)).To(BeNil()) @@ -310,33 +355,54 @@ var _ = Describe("ArtistRepository", func() { var repo model.ArtistRepository BeforeEach(func() { - ctx := GinkgoT().Context() ctx = request.WithUser(ctx, adminUser) - repo = NewArtistRepository(ctx, GetDBXBuilder()) + repo = NewArtistRepository(GetDBXBuilder()) + }) + + Describe("ReadAll with role sort", func() { + It("does not change the shared sort mappings", func() { + original := repo.(*artistRepository).sortMappings["song_count"] + var wg sync.WaitGroup + for i := 0; i < 20; i++ { + role := "artist" + if i%2 == 1 { + role = "composer" + } + wg.Add(1) + go func() { + defer GinkgoRecover() + defer wg.Done() + _, err := repo.ReadAll(ctx, rest.QueryOptions{Sort: "song_count", Filters: map[string]any{"role": role}}) + Expect(err).ToNot(HaveOccurred()) + }() + } + wg.Wait() + Expect(repo.(*artistRepository).sortMappings["song_count"]).To(Equal(original)) + }) }) Describe("GetCursor", func() { It("yields the same artists as GetAll", func() { opts := model.QueryOptions{Sort: "name"} - want, err := repo.GetAll(opts) + want, err := repo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(repo.GetCursor(opts))).To(Equal([]model.Artist(want))) + Expect(collectCursor(repo.GetCursor(ctx, opts))).To(Equal([]model.Artist(want))) }) It("honors Max/Offset like GetAll", func() { opts := model.QueryOptions{Sort: "name", Max: 2, Offset: 1} - want, err := repo.GetAll(opts) + want, err := repo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(repo.GetCursor(opts))).To(Equal([]model.Artist(want))) + Expect(collectCursor(repo.GetCursor(ctx, opts))).To(Equal([]model.Artist(want))) }) }) Describe("getAllIDs", func() { It("returns the same id set as GetAll", func() { - want, err := repo.GetAll() + want, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - ids, err := repo.(*artistRepository).getAllIDs() + ids, err := repo.(*artistRepository).getAllIDs(ctx) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(ConsistOf(slice.Map(want, func(a model.Artist) string { return a.ID }))) }) @@ -345,12 +411,12 @@ var _ = Describe("ArtistRepository", func() { Describe("Basic Operations", func() { Describe("Count", func() { It("returns the number of artists in the DB", func() { - Expect(repo.CountAll()).To(Equal(int64(4))) + Expect(repo.CountAll(ctx)).To(Equal(int64(4))) }) It("counts starred artists when an annotation filter is present", func() { // The Beatles (id 3) is starred for the admin user in the seed data - count, err := repo.CountAll(model.QueryOptions{ + count, err := repo.CountAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("starred")("starred", "true"), }) Expect(err).ToNot(HaveOccurred()) @@ -358,7 +424,7 @@ var _ = Describe("ArtistRepository", func() { }) It("counts with has_rating=false without a 'no such column' error (join kept)", func() { - count, err := repo.CountAll(model.QueryOptions{ + count, err := repo.CountAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("rating")("rating", "false"), }) Expect(err).ToNot(HaveOccurred()) @@ -368,16 +434,16 @@ var _ = Describe("ArtistRepository", func() { Describe("Exists", func() { It("returns true for an artist that is in the DB", func() { - Expect(repo.Exists("3")).To(BeTrue()) + Expect(repo.Exists(ctx, "3")).To(BeTrue()) }) It("returns false for an artist that is NOT in the DB", func() { - Expect(repo.Exists("666")).To(BeFalse()) + Expect(repo.Exists(ctx, "666")).To(BeFalse()) }) }) Describe("Get", func() { It("retrieves existing artist data", func() { - artist, err := repo.Get("2") + artist, err := repo.Get(ctx, "2") Expect(err).ToNot(HaveOccurred()) Expect(artist.Name).To(Equal(artistKraftwerk.Name)) }) @@ -392,10 +458,10 @@ var _ = Describe("ArtistRepository", func() { It("returns the index when PreferSortTags is true and SortArtistName is not empty", func() { // Set SortArtistName to "Foo" for Beatles artistBeatles.SortArtistName = "Foo" - er := repo.Put(&artistBeatles) + er := repo.Put(ctx, &artistBeatles) Expect(er).To(BeNil()) - idx, err := repo.GetIndex(false, []int{1}) + idx, err := repo.GetIndex(ctx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(4)) Expect(idx[0].ID).To(Equal("F")) @@ -413,13 +479,13 @@ var _ = Describe("ArtistRepository", func() { // Restore the original value artistBeatles.SortArtistName = "" - er = repo.Put(&artistBeatles) + er = repo.Put(ctx, &artistBeatles) Expect(er).To(BeNil()) }) // BFR Empty SortArtistName is not saved in the DB anymore XIt("returns the index when PreferSortTags is true and SortArtistName is empty", func() { - idx, err := repo.GetIndex(false, []int{1}) + idx, err := repo.GetIndex(ctx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(4)) Expect(idx[0].ID).To(Equal("B")) @@ -444,10 +510,10 @@ var _ = Describe("ArtistRepository", func() { It("returns the index when SortArtistName is NOT empty", func() { // Set SortArtistName to "Foo" for Beatles artistBeatles.SortArtistName = "Foo" - er := repo.Put(&artistBeatles) + er := repo.Put(ctx, &artistBeatles) Expect(er).To(BeNil()) - idx, err := repo.GetIndex(false, []int{1}) + idx, err := repo.GetIndex(ctx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(4)) Expect(idx[0].ID).To(Equal("B")) @@ -465,12 +531,12 @@ var _ = Describe("ArtistRepository", func() { // Restore the original value artistBeatles.SortArtistName = "" - er = repo.Put(&artistBeatles) + er = repo.Put(ctx, &artistBeatles) Expect(er).To(BeNil()) }) It("returns the index when SortArtistName is empty", func() { - idx, err := repo.GetIndex(false, []int{1}) + idx, err := repo.GetIndex(ctx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(4)) Expect(idx[0].ID).To(Equal("B")) @@ -498,14 +564,14 @@ var _ = Describe("ArtistRepository", func() { producerStats := `{"producer": {"s": 500, "m": 3, "a": 1}}` // Set Beatles as composer in library 1 - _, err := raw.executeSQL(squirrel.Insert("library_artist"). + _, err := raw.executeSQL(ctx, squirrel.Insert("library_artist"). Columns("library_id", "artist_id", "stats"). Values(1, artistBeatles.ID, composerStats). Suffix("ON CONFLICT(library_id, artist_id) DO UPDATE SET stats = excluded.stats")) Expect(err).ToNot(HaveOccurred()) // Set Kraftwerk as producer in library 1 - _, err = raw.executeSQL(squirrel.Insert("library_artist"). + _, err = raw.executeSQL(ctx, squirrel.Insert("library_artist"). Columns("library_id", "artist_id", "stats"). Values(1, artistKraftwerk.ID, producerStats). Suffix("ON CONFLICT(library_id, artist_id) DO UPDATE SET stats = excluded.stats")) @@ -514,16 +580,16 @@ var _ = Describe("ArtistRepository", func() { AfterEach(func() { // Clean up stats from library_artist table - _, _ = raw.executeSQL(squirrel.Update("library_artist"). + _, _ = raw.executeSQL(ctx, squirrel.Update("library_artist"). Set("stats", "{}"). Where(squirrel.Eq{"artist_id": artistBeatles.ID, "library_id": 1})) - _, _ = raw.executeSQL(squirrel.Update("library_artist"). + _, _ = raw.executeSQL(ctx, squirrel.Update("library_artist"). Set("stats", "{}"). Where(squirrel.Eq{"artist_id": artistKraftwerk.ID, "library_id": 1})) }) It("returns only artists with the specified role", func() { - idx, err := repo.GetIndex(false, []int{1}, model.RoleComposer) + idx, err := repo.GetIndex(ctx, false, []int{1}, model.RoleComposer) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(1)) Expect(idx[0].ID).To(Equal("B")) @@ -532,7 +598,7 @@ var _ = Describe("ArtistRepository", func() { }) It("returns artists with any of the specified roles", func() { - idx, err := repo.GetIndex(false, []int{1}, model.RoleComposer, model.RoleProducer) + idx, err := repo.GetIndex(ctx, false, []int{1}, model.RoleComposer, model.RoleProducer) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(2)) @@ -553,7 +619,7 @@ var _ = Describe("ArtistRepository", func() { }) It("returns empty index when no artists have the specified role", func() { - idx, err := repo.GetIndex(false, []int{1}, model.RoleDirector) + idx, err := repo.GetIndex(ctx, false, []int{1}, model.RoleDirector) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(0)) }) @@ -561,19 +627,19 @@ var _ = Describe("ArtistRepository", func() { When("validating library IDs", func() { It("returns nil when no library IDs are provided", func() { - idx, err := repo.GetIndex(false, []int{}) + idx, err := repo.GetIndex(ctx, false, []int{}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(BeNil()) }) It("returns artists when library IDs are provided (admin user sees all content)", func() { // Admin users can see all content when valid library IDs are provided - idx, err := repo.GetIndex(false, []int{1}) + idx, err := repo.GetIndex(ctx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(4)) // With non-existent library ID, admin users see no content because no artists are associated with that library - idx, err = repo.GetIndex(false, []int{999}) + idx, err = repo.GetIndex(ctx, false, []int{999}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(0)) // Even admin users need valid library associations }) @@ -586,23 +652,23 @@ var _ = Describe("ArtistRepository", func() { BeforeEach(func() { // Create artist without any annotation artistWithoutAnnotation = model.Artist{ID: "no-annotation-artist", Name: "No Annotation Artist"} - err := createArtistWithLibrary(repo, &artistWithoutAnnotation, 1) + err := createArtistWithLibrary(ctx, repo, &artistWithoutAnnotation, 1) Expect(err).ToNot(HaveOccurred()) }) AfterEach(func() { if raw, ok := repo.(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": artistWithoutAnnotation.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": artistWithoutAnnotation.ID})) } }) Describe("starred", func() { It("false includes items without annotations", func() { - res, err := repo.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := repo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "false"}, }) Expect(err).ToNot(HaveOccurred()) - artists := res.(model.Artists) + artists := res var found bool for _, a := range artists { @@ -615,11 +681,11 @@ var _ = Describe("ArtistRepository", func() { }) It("true excludes items without annotations", func() { - res, err := repo.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := repo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "true"}, }) Expect(err).ToNot(HaveOccurred()) - artists := res.(model.Artists) + artists := res for _, a := range artists { Expect(a.ID).ToNot(Equal(artistWithoutAnnotation.ID)) @@ -631,57 +697,63 @@ var _ = Describe("ArtistRepository", func() { Describe("MBID and Text Search", func() { var lib2 model.Library var lr model.LibraryRepository + var lrCtx context.Context var restrictedUser model.User var restrictedRepo model.ArtistRepository + var restrictedCtx context.Context var headlessRepo model.ArtistRepository + var headlessCtx context.Context BeforeEach(func() { // Set up headless repo (no user context) - headlessRepo = NewArtistRepository(context.Background(), GetDBXBuilder()) + headlessCtx = GinkgoT().Context() + headlessRepo = NewArtistRepository(GetDBXBuilder()) // Create library for testing access restrictions lib2 = model.Library{ID: 0, Name: "Artist Test Library", Path: "/artist/test/lib"} - lr = NewLibraryRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - err := lr.Put(&lib2) + lrCtx = request.WithUser(ctx, adminUser) + lr = NewLibraryRepository(GetDBXBuilder()) + err := lr.Put(lrCtx, &lib2) Expect(err).ToNot(HaveOccurred()) // Create a user with access to only library 1 restrictedUser = createUserWithLibraries("search_user", []int{1}) // Create repository context for the restricted user - ctx := request.WithUser(GinkgoT().Context(), restrictedUser) - restrictedRepo = NewArtistRepository(ctx, GetDBXBuilder()) + restrictedCtx = request.WithUser(ctx, restrictedUser) + restrictedRepo = NewArtistRepository(GetDBXBuilder()) // Ensure both test artists are associated with library 1 - err = lr.AddArtist(1, artistBeatles.ID) + err = lr.AddArtist(lrCtx, 1, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) - err = lr.AddArtist(1, artistKraftwerk.ID) + err = lr.AddArtist(lrCtx, 1, artistKraftwerk.ID) Expect(err).ToNot(HaveOccurred()) // Create the restricted user in the database - ur := NewUserRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - err = ur.Put(&restrictedUser) + urCtx := request.WithUser(ctx, adminUser) + ur := NewUserRepository(GetDBXBuilder()) + err = ur.Put(urCtx, &restrictedUser) Expect(err).ToNot(HaveOccurred()) - err = ur.SetUserLibraries(restrictedUser.ID, []int{1}) + err = ur.SetUserLibraries(urCtx, restrictedUser.ID, []int{1}) Expect(err).ToNot(HaveOccurred()) }) AfterEach(func() { // Clean up library 2 - lr := NewLibraryRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - _ = lr.(*libraryRepository).delete(squirrel.Eq{"id": lib2.ID}) + lr := NewLibraryRepository(GetDBXBuilder()) + _ = lr.(*libraryRepository).delete(ctx, squirrel.Eq{"id": lib2.ID}) }) DescribeTable("MBID search behavior across different user types", - func(testRepo *model.ArtistRepository, shouldFind bool, testDesc string) { + func(testRepo *model.ArtistRepository, testCtx *context.Context, shouldFind bool, testDesc string) { // Create test artist with MBID artistWithMBID := createTestArtistWithMBID("test-mbid-artist", "Test MBID Artist", "550e8400-e29b-41d4-a716-446655440010") - err := createArtistWithLibrary(*testRepo, &artistWithMBID, 1) + err := createArtistWithLibrary(*testCtx, *testRepo, &artistWithMBID, 1) Expect(err).ToNot(HaveOccurred()) // Test the search - results, err := (*testRepo).Search("550e8400-e29b-41d4-a716-446655440010", model.QueryOptions{Max: 10}) + results, err := (*testRepo).Search(*testCtx, "550e8400-e29b-41d4-a716-446655440010", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) if shouldFind { @@ -693,43 +765,43 @@ var _ = Describe("ArtistRepository", func() { // Clean up if raw, ok := (*testRepo).(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": artistWithMBID.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": artistWithMBID.ID})) } }, - Entry("Admin user can find artist by MBID", &repo, true, "Admin should find MBID artist"), - Entry("Restricted user can find artist by MBID in accessible library", &restrictedRepo, true, "Restricted user should find MBID artist in accessible library"), - Entry("Headless process can find artist by MBID", &headlessRepo, true, "Headless process should find MBID artist"), + Entry("Admin user can find artist by MBID", &repo, &ctx, true, "Admin should find MBID artist"), + Entry("Restricted user can find artist by MBID in accessible library", &restrictedRepo, &restrictedCtx, true, "Restricted user should find MBID artist in accessible library"), + Entry("Headless process can find artist by MBID", &headlessRepo, &headlessCtx, true, "Headless process should find MBID artist"), ) It("prevents restricted user from finding artist by MBID when not in accessible library", func() { // Create an artist in library 2 (not accessible to restricted user) inaccessibleArtist := createTestArtistWithMBID("inaccessible-mbid-artist", "Inaccessible MBID Artist", "a74b1b7f-71a5-4011-9441-d0b5e4122711") - err := repo.Put(&inaccessibleArtist) + err := repo.Put(ctx, &inaccessibleArtist) Expect(err).ToNot(HaveOccurred()) // Add to library 2 (not accessible to restricted user) - err = lr.AddArtist(lib2.ID, inaccessibleArtist.ID) + err = lr.AddArtist(lrCtx, lib2.ID, inaccessibleArtist.ID) Expect(err).ToNot(HaveOccurred()) // Restricted user should not find this artist - results, err := restrictedRepo.Search("a74b1b7f-71a5-4011-9441-d0b5e4122711", model.QueryOptions{Max: 10}) + results, err := restrictedRepo.Search(restrictedCtx, "a74b1b7f-71a5-4011-9441-d0b5e4122711", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) // But admin should find it - results, err = repo.Search("a74b1b7f-71a5-4011-9441-d0b5e4122711", model.QueryOptions{Max: 10}) + results, err = repo.Search(ctx, "a74b1b7f-71a5-4011-9441-d0b5e4122711", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) // Clean up if raw, ok := repo.(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": inaccessibleArtist.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": inaccessibleArtist.ID})) } }) Context("Text Search", func() { It("allows admin to find artists by name regardless of library", func() { - results, err := repo.Search("Beatles", model.QueryOptions{Max: 10}) + results, err := repo.Search(ctx, "Beatles", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Name).To(Equal("The Beatles")) @@ -741,21 +813,21 @@ var _ = Describe("ArtistRepository", func() { ID: "inaccessible-text-artist", Name: "Unique Search Name Artist", } - err := repo.Put(&inaccessibleArtist) + err := repo.Put(ctx, &inaccessibleArtist) Expect(err).ToNot(HaveOccurred()) // Add to library 2 (not accessible to restricted user) - err = lr.AddArtist(lib2.ID, inaccessibleArtist.ID) + err = lr.AddArtist(lrCtx, lib2.ID, inaccessibleArtist.ID) Expect(err).ToNot(HaveOccurred()) // Restricted user should not find this artist - results, err := restrictedRepo.Search("Unique Search Name", model.QueryOptions{Max: 10}) + results, err := restrictedRepo.Search(restrictedCtx, "Unique Search Name", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty(), "Text search should respect library filtering") // Clean up if raw, ok := repo.(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": inaccessibleArtist.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": inaccessibleArtist.ID})) } }) }) @@ -764,15 +836,15 @@ var _ = Describe("ArtistRepository", func() { It("does not duplicate artists that belong to multiple libraries", func() { // An artist in two libraries has two library_artist rows; pagination // must still enumerate it exactly once, at a stable offset. - Expect(lr.AddArtist(lib2.ID, artistBeatles.ID)).To(Succeed()) + Expect(lr.AddArtist(lrCtx, lib2.ID, artistBeatles.ID)).To(Succeed()) - all, err := repo.Search("", model.QueryOptions{Max: 1000}) + all, err := repo.Search(ctx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) seen := map[string]bool{} var paged model.Artists for offset := range len(all) { - page, err := repo.Search("", model.QueryOptions{Max: 1, Offset: offset}) + page, err := repo.Search(ctx, "", model.QueryOptions{Max: 1, Offset: offset}) Expect(err).ToNot(HaveOccurred()) for _, a := range page { Expect(seen[a.ID]).To(BeFalse(), fmt.Sprintf("artist %s returned twice", a.ID)) @@ -784,14 +856,14 @@ var _ = Describe("ArtistRepository", func() { }) It("paginates all artists in natural order without overlaps or gaps", func() { - all, err := repo.Search("", model.QueryOptions{Max: 1000}) + all, err := repo.Search(ctx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) Expect(len(all)).To(BeNumerically(">", 1)) var paged model.Artists pageSize := 2 for offset := 0; offset < len(all); offset += pageSize { - page, err := repo.Search("", model.QueryOptions{Max: pageSize, Offset: offset}) + page, err := repo.Search(ctx, "", model.QueryOptions{Max: pageSize, Offset: offset}) Expect(err).ToNot(HaveOccurred()) paged = append(paged, page...) } @@ -804,10 +876,10 @@ var _ = Describe("ArtistRepository", func() { It("respects library filtering for restricted users", func() { // Create an artist only in library 2 (not accessible to restricted user) lib2Artist := model.Artist{ID: "empty-query-lib2-artist", Name: "Empty Query Lib2 Artist"} - Expect(repo.Put(&lib2Artist)).To(Succeed()) - Expect(lr.AddArtist(lib2.ID, lib2Artist.ID)).To(Succeed()) + Expect(repo.Put(ctx, &lib2Artist)).To(Succeed()) + Expect(lr.AddArtist(lrCtx, lib2.ID, lib2Artist.ID)).To(Succeed()) - results, err := restrictedRepo.Search("", model.QueryOptions{Max: 1000}) + results, err := restrictedRepo.Search(restrictedCtx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) for _, a := range results { Expect(a.ID).ToNot(Equal(lib2Artist.ID), "Empty query search should respect library filtering") @@ -815,7 +887,7 @@ var _ = Describe("ArtistRepository", func() { // Clean up if raw, ok := repo.(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": lib2Artist.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": lib2Artist.ID})) } }) @@ -823,15 +895,15 @@ var _ = Describe("ArtistRepository", func() { // ID "25" sorts between base fixtures "2" and "3", so this lib2-only artist lands // inside the restricted user's visible range — exercising the no-gap guarantee. lib2Artist := model.Artist{ID: "25", Name: "Restricted Lib2 Artist"} - Expect(repo.Put(&lib2Artist)).To(Succeed()) - Expect(lr.AddArtist(lib2.ID, lib2Artist.ID)).To(Succeed()) + Expect(repo.Put(ctx, &lib2Artist)).To(Succeed()) + Expect(lr.AddArtist(lrCtx, lib2.ID, lib2Artist.ID)).To(Succeed()) DeferCleanup(func() { if raw, ok := repo.(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": lib2Artist.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": lib2Artist.ID})) } }) - all, err := restrictedRepo.Search("", model.QueryOptions{Max: 1000}) + all, err := restrictedRepo.Search(restrictedCtx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) Expect(len(all)).To(BeNumerically(">", 1)) for _, a := range all { @@ -840,7 +912,7 @@ var _ = Describe("ArtistRepository", func() { var paged model.Artists for offset := range len(all) { - page, err := restrictedRepo.Search("", model.QueryOptions{Max: 1, Offset: offset}) + page, err := restrictedRepo.Search(restrictedCtx, "", model.QueryOptions{Max: 1, Offset: offset}) Expect(err).ToNot(HaveOccurred()) Expect(page).To(HaveLen(1), fmt.Sprintf("page at offset %d should be full", offset)) paged = append(paged, page...) @@ -855,11 +927,11 @@ var _ = Describe("ArtistRepository", func() { Context("Headless Processes (No User Context)", func() { It("should see all artists from all libraries when no user is in context", func() { // Add artists to different libraries - err := lr.AddArtist(lib2.ID, artistBeatles.ID) + err := lr.AddArtist(lrCtx, lib2.ID, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) // Headless processes should see all artists regardless of library - artists, err := headlessRepo.GetAll() + artists, err := headlessRepo.GetAll(headlessCtx) Expect(err).ToNot(HaveOccurred()) // Should see all artists from all libraries @@ -875,11 +947,11 @@ var _ = Describe("ArtistRepository", func() { It("should allow headless processes to apply explicit library_id filters", func() { // Add artists to different libraries - err := lr.AddArtist(lib2.ID, artistBeatles.ID) + err := lr.AddArtist(lrCtx, lib2.ID, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) // Filter by specific library - artists, err := headlessRepo.GetAll(model.QueryOptions{ + artists, err := headlessRepo.GetAll(headlessCtx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -895,11 +967,11 @@ var _ = Describe("ArtistRepository", func() { It("should get individual artists when no user is in context", func() { // Add artist to a library - err := lr.AddArtist(lib2.ID, artistBeatles.ID) + err := lr.AddArtist(lrCtx, lib2.ID, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) // Headless process should be able to get the artist - artist, err := headlessRepo.Get(artistBeatles.ID) + artist, err := headlessRepo.Get(headlessCtx, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) Expect(artist.ID).To(Equal(artistBeatles.ID)) }) @@ -908,15 +980,15 @@ var _ = Describe("ArtistRepository", func() { Describe("Admin User Library Access", func() { It("sees all artists regardless of library permissions", func() { - count, err := repo.CountAll() + count, err := repo.CountAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(4))) - artists, err := repo.GetAll() + artists, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(artists).To(HaveLen(4)) - exists, err := repo.Exists(artistBeatles.ID) + exists, err := repo.Exists(ctx, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeTrue()) }) @@ -931,25 +1003,25 @@ var _ = Describe("ArtistRepository", func() { missingArtist = model.Artist{ID: "missing_test", Name: "Missing Artist", OrderArtistName: "missing artist"} // Create and mark as missing - err := createArtistWithLibrary(repo, &missingArtist, 1) + err := createArtistWithLibrary(ctx, repo, &missingArtist, 1) Expect(err).ToNot(HaveOccurred()) - _, err = raw.executeSQL(squirrel.Update(raw.tableName).Set("missing", true).Where(squirrel.Eq{"id": missingArtist.ID})) + _, err = raw.executeSQL(ctx, squirrel.Update(raw.tableName).Set("missing", true).Where(squirrel.Eq{"id": missingArtist.ID})) Expect(err).ToNot(HaveOccurred()) }) AfterEach(func() { - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": missingArtist.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": missingArtist.ID})) }) It("missing artists are never returned by search", func() { // Should see missing artist in GetAll by default for admin users - artists, err := repo.GetAll() + artists, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(artists).To(HaveLen(5)) // Including the missing artist // Search never returns missing artists (hardcoded behavior) - results, err := repo.Search("Missing Artist", model.QueryOptions{Max: 10}) + results, err := repo.Search(ctx, "Missing Artist", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -958,6 +1030,7 @@ var _ = Describe("ArtistRepository", func() { Context("Regular User Operations", func() { var restrictedRepo model.ArtistRepository + var restrictedCtx context.Context var unauthorizedUser model.User BeforeEach(func() { @@ -965,55 +1038,54 @@ var _ = Describe("ArtistRepository", func() { unauthorizedUser = model.User{ID: "restricted_user", UserName: "restricted", Name: "Restricted User", Email: "restricted@test.com", IsAdmin: false} // Create repository context for the unauthorized user - ctx := GinkgoT().Context() - ctx = request.WithUser(ctx, unauthorizedUser) - restrictedRepo = NewArtistRepository(ctx, GetDBXBuilder()) + restrictedCtx = request.WithUser(ctx, unauthorizedUser) + restrictedRepo = NewArtistRepository(GetDBXBuilder()) }) Describe("Library Access Restrictions", func() { It("CountAll returns 0 for users without library access", func() { - count, err := restrictedRepo.CountAll() + count, err := restrictedRepo.CountAll(restrictedCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(0))) }) It("GetAll returns empty list for users without library access", func() { - artists, err := restrictedRepo.GetAll() + artists, err := restrictedRepo.GetAll(restrictedCtx) Expect(err).ToNot(HaveOccurred()) Expect(artists).To(BeEmpty()) }) It("Exists returns false for existing artists when user has no library access", func() { // These artists exist in the DB but the user has no access to them - exists, err := restrictedRepo.Exists(artistBeatles.ID) + exists, err := restrictedRepo.Exists(restrictedCtx, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeFalse()) - exists, err = restrictedRepo.Exists(artistKraftwerk.ID) + exists, err = restrictedRepo.Exists(restrictedCtx, artistKraftwerk.ID) Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeFalse()) }) It("Get returns ErrNotFound for existing artists when user has no library access", func() { - _, err := restrictedRepo.Get(artistBeatles.ID) + _, err := restrictedRepo.Get(restrictedCtx, artistBeatles.ID) Expect(err).To(Equal(model.ErrNotFound)) - _, err = restrictedRepo.Get(artistKraftwerk.ID) + _, err = restrictedRepo.Get(restrictedCtx, artistKraftwerk.ID) Expect(err).To(Equal(model.ErrNotFound)) }) It("Search returns empty results for users without library access", func() { - results, err := restrictedRepo.Search("Beatles", model.QueryOptions{Max: 10}) + results, err := restrictedRepo.Search(restrictedCtx, "Beatles", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) - results, err = restrictedRepo.Search("Kraftwerk", model.QueryOptions{Max: 10}) + results, err = restrictedRepo.Search(restrictedCtx, "Kraftwerk", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) It("GetIndex returns empty index for users without library access", func() { - idx, err := restrictedRepo.GetIndex(false, []int{1}) + idx, err := restrictedRepo.GetIndex(restrictedCtx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(0)) }) @@ -1021,42 +1093,43 @@ var _ = Describe("ArtistRepository", func() { Context("when user gains library access", func() { BeforeEach(func() { - ctx := GinkgoT().Context() // Give the user access to library 1 - ur := NewUserRepository(request.WithUser(ctx, adminUser), GetDBXBuilder()) + urCtx := request.WithUser(ctx, adminUser) + ur := NewUserRepository(GetDBXBuilder()) // First create the user if not exists - err := ur.Put(&unauthorizedUser) + err := ur.Put(urCtx, &unauthorizedUser) Expect(err).ToNot(HaveOccurred()) // Then add library access - err = ur.SetUserLibraries(unauthorizedUser.ID, []int{1}) + err = ur.SetUserLibraries(urCtx, unauthorizedUser.ID, []int{1}) Expect(err).ToNot(HaveOccurred()) // Update the user object with the libraries to simulate middleware behavior - libraries, err := ur.GetUserLibraries(unauthorizedUser.ID) + libraries, err := ur.GetUserLibraries(urCtx, unauthorizedUser.ID) Expect(err).ToNot(HaveOccurred()) unauthorizedUser.Libraries = libraries // Recreate repository context with updated user - ctx = request.WithUser(ctx, unauthorizedUser) - restrictedRepo = NewArtistRepository(ctx, GetDBXBuilder()) + restrictedCtx = request.WithUser(ctx, unauthorizedUser) + restrictedRepo = NewArtistRepository(GetDBXBuilder()) }) AfterEach(func() { // Clean up: remove the user's library access - ur := NewUserRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - _ = ur.SetUserLibraries(unauthorizedUser.ID, []int{}) + urCtx := request.WithUser(ctx, adminUser) + ur := NewUserRepository(GetDBXBuilder()) + _ = ur.SetUserLibraries(urCtx, unauthorizedUser.ID, []int{}) }) It("CountAll returns correct count after gaining access", func() { - count, err := restrictedRepo.CountAll() + count, err := restrictedRepo.CountAll(restrictedCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(4))) // Beatles, Kraftwerk, Seatbelts, and The Roots }) It("GetAll returns artists after gaining access", func() { - artists, err := restrictedRepo.GetAll() + artists, err := restrictedRepo.GetAll(restrictedCtx) Expect(err).ToNot(HaveOccurred()) Expect(artists).To(HaveLen(4)) @@ -1068,23 +1141,23 @@ var _ = Describe("ArtistRepository", func() { }) It("Exists returns true for accessible artists", func() { - exists, err := restrictedRepo.Exists(artistBeatles.ID) + exists, err := restrictedRepo.Exists(restrictedCtx, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeTrue()) - exists, err = restrictedRepo.Exists(artistKraftwerk.ID) + exists, err = restrictedRepo.Exists(restrictedCtx, artistKraftwerk.ID) Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeTrue()) }) It("GetIndex returns artists with proper library filtering", func() { // With valid library access, should see artists - idx, err := restrictedRepo.GetIndex(false, []int{1}) + idx, err := restrictedRepo.GetIndex(restrictedCtx, false, []int{1}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(4)) // With non-existent library ID, should see nothing (non-admin user) - idx, err = restrictedRepo.GetIndex(false, []int{999}) + idx, err = restrictedRepo.GetIndex(restrictedCtx, false, []int{999}) Expect(err).ToNot(HaveOccurred()) Expect(idx).To(HaveLen(0)) }) @@ -1092,11 +1165,12 @@ var _ = Describe("ArtistRepository", func() { It("takes the unfiltered fast-path when the user can access every library", func() { // The fixture DB has a single library and the user was granted it, so it has access // to all libraries: search results must match what an admin sees. - adminRepo := NewArtistRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - adminAll, err := adminRepo.Search("", model.QueryOptions{Max: 1000}) + adminCtx := request.WithUser(ctx, adminUser) + adminRepo := NewArtistRepository(GetDBXBuilder()) + adminAll, err := adminRepo.Search(adminCtx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) - userAll, err := restrictedRepo.Search("", model.QueryOptions{Max: 1000}) + userAll, err := restrictedRepo.Search(restrictedCtx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) ids := func(artists model.Artists) []string { @@ -1115,7 +1189,7 @@ var _ = Describe("ArtistRepository", func() { // visible-library count reaches the DB total. Derive the total from the DB so the // assertion doesn't depend on how many libraries other specs left behind. raw := restrictedRepo.(*artistRepository) // context carries a non-admin user - total, err := NewLibraryRepository(GinkgoT().Context(), GetDBXBuilder()).CountAll() + total, err := NewLibraryRepository(GetDBXBuilder()).CountAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(total).To(BeNumerically(">", 0)) @@ -1123,9 +1197,9 @@ var _ = Describe("ArtistRepository", func() { for i := range allLibs { allLibs[i] = i + 1 } - Expect(raw.userSeesAllLibraries(allLibs)).To(BeTrue()) - Expect(raw.userSeesAllLibraries(allLibs[:total-1])).To(BeFalse()) - Expect(raw.userSeesAllLibraries([]int{})).To(BeFalse()) + Expect(raw.userSeesAllLibraries(restrictedCtx, allLibs)).To(BeTrue()) + Expect(raw.userSeesAllLibraries(restrictedCtx, allLibs[:total-1])).To(BeFalse()) + Expect(raw.userSeesAllLibraries(restrictedCtx, []int{})).To(BeFalse()) }) }) }) @@ -1139,8 +1213,8 @@ var _ = Describe("ArtistRepository", func() { tmpDir = GinkgoT().TempDir() conf.Server.DataFolder = conf.NewDir(tmpDir) - ctx := request.WithUser(GinkgoT().Context(), adminUser) - repo = NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository) + ctx = request.WithUser(ctx, adminUser) + repo = NewArtistRepository(GetDBXBuilder()).(*artistRepository) }) // Helper to create an artist image file on disk and return its path @@ -1155,13 +1229,13 @@ var _ = Describe("ArtistRepository", func() { It("removes uploaded image files for purged artists", func() { // Create an orphan artist (not in album_artists) with an uploaded image orphanArtist := model.Artist{ID: "orphan-with-image", Name: "Orphan Artist", UploadedImage: "orphan-with-image_Orphan_Artist.jpg"} - Expect(repo.Put(&orphanArtist)).To(Succeed()) + Expect(repo.Put(ctx, &orphanArtist)).To(Succeed()) imgPath := createImageFile("orphan-with-image_Orphan_Artist.jpg") - Expect(repo.purgeEmpty()).To(Succeed()) + Expect(repo.purgeEmpty(ctx)).To(Succeed()) // Artist should be gone from DB - exists, err := repo.Exists("orphan-with-image") + exists, err := repo.Exists(ctx, "orphan-with-image") Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeFalse()) @@ -1173,12 +1247,12 @@ var _ = Describe("ArtistRepository", func() { It("handles missing image files gracefully", func() { // Artist has UploadedImage set but no actual file on disk orphanArtist := model.Artist{ID: "orphan-no-file", Name: "Ghost Image", UploadedImage: "orphan-no-file_Ghost_Image.jpg"} - Expect(repo.Put(&orphanArtist)).To(Succeed()) + Expect(repo.Put(ctx, &orphanArtist)).To(Succeed()) - Expect(repo.purgeEmpty()).To(Succeed()) + Expect(repo.purgeEmpty(ctx)).To(Succeed()) // Artist should be gone from DB - exists, err := repo.Exists("orphan-no-file") + exists, err := repo.Exists(ctx, "orphan-no-file") Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeFalse()) }) @@ -1186,24 +1260,24 @@ var _ = Describe("ArtistRepository", func() { It("does not delete images for artists that are kept", func() { // Create an artist with an uploaded image AND an album_artists entry so it won't be purged keptArtist := model.Artist{ID: "kept-artist", Name: "Kept Artist", UploadedImage: "kept-artist_Kept_Artist.jpg"} - Expect(repo.Put(&keptArtist)).To(Succeed()) + Expect(repo.Put(ctx, &keptArtist)).To(Succeed()) imgPath := createImageFile("kept-artist_Kept_Artist.jpg") // Insert an album_artists record to keep this artist from being purged - _, err := repo.executeSQL(squirrel.Insert("album_artists"). + _, err := repo.executeSQL(ctx, squirrel.Insert("album_artists"). SetMap(map[string]any{"album_id": "101", "artist_id": "kept-artist", "role": "artist", "sub_role": ""})) Expect(err).ToNot(HaveOccurred()) DeferCleanup(func() { - _, _ = repo.executeSQL(squirrel.Delete("album_artists").Where(squirrel.Eq{"artist_id": "kept-artist"})) - _ = repo.delete(squirrel.Eq{"id": "kept-artist"}) + _, _ = repo.executeSQL(ctx, squirrel.Delete("album_artists").Where(squirrel.Eq{"artist_id": "kept-artist"})) + _ = repo.delete(ctx, squirrel.Eq{"id": "kept-artist"}) }) - Expect(repo.purgeEmpty()).To(Succeed()) + Expect(repo.purgeEmpty(ctx)).To(Succeed()) // Artist should still exist (check directly, bypassing library filter) var ids []string - err = repo.queryAllSlice(squirrel.Select("id").From("artist").Where(squirrel.Eq{"id": "kept-artist"}), &ids) + err = repo.queryAllSlice(ctx, squirrel.Select("id").From("artist").Where(squirrel.Eq{"id": "kept-artist"}), &ids) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(HaveLen(1)) @@ -1218,37 +1292,37 @@ var _ = Describe("ArtistRepository", func() { missing := func(id string) bool { var vals []bool - Expect(repo.queryAllSlice(squirrel.Select("missing").From("artist").Where(squirrel.Eq{"id": id}), &vals)).To(Succeed()) + Expect(repo.queryAllSlice(ctx, squirrel.Select("missing").From("artist").Where(squirrel.Eq{"id": id}), &vals)).To(Succeed()) Expect(vals).To(HaveLen(1)) return vals[0] } BeforeEach(func() { - ctx := request.WithUser(GinkgoT().Context(), adminUser) - repo = NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository) + ctx = request.WithUser(ctx, adminUser) + repo = NewArtistRepository(GetDBXBuilder()).(*artistRepository) }) It("marks artists missing when the empty-stats cleanup drops their last library_artist row", func() { // A library_artist row with stats '{}' (no content) gets deleted by the cleanup, // which would orphan this non-missing artist. emptyArtist := model.Artist{ID: "refresh-empty", Name: "No Content Artist"} - Expect(repo.Put(&emptyArtist)).To(Succeed()) - _, err := repo.executeSQL(squirrel.Insert("library_artist"). + Expect(repo.Put(ctx, &emptyArtist)).To(Succeed()) + _, err := repo.executeSQL(ctx, squirrel.Insert("library_artist"). SetMap(map[string]any{"library_id": 1, "artist_id": emptyArtist.ID, "stats": "{}"})) Expect(err).ToNot(HaveOccurred()) DeferCleanup(func() { - _, _ = repo.executeSQL(squirrel.Delete("library_artist").Where(squirrel.Eq{"artist_id": emptyArtist.ID})) - _ = repo.delete(squirrel.Eq{"id": emptyArtist.ID}) + _, _ = repo.executeSQL(ctx, squirrel.Delete("library_artist").Where(squirrel.Eq{"artist_id": emptyArtist.ID})) + _ = repo.delete(ctx, squirrel.Eq{"id": emptyArtist.ID}) }) Expect(missing(emptyArtist.ID)).To(BeFalse()) - _, err = repo.RefreshStats(true) + _, err = repo.RefreshStats(ctx, true) Expect(err).ToNot(HaveOccurred()) Expect(missing(emptyArtist.ID)).To(BeTrue()) var orphanIDs []string - Expect(repo.queryAllSlice(squirrel.Select("id").From("artist"). + Expect(repo.queryAllSlice(ctx, squirrel.Select("id").From("artist"). Where("missing = false"). Where("id not in (select artist_id from library_artist)"), &orphanIDs)).To(Succeed()) Expect(orphanIDs).ToNot(ContainElement(emptyArtist.ID)) @@ -1259,14 +1333,14 @@ var _ = Describe("ArtistRepository", func() { // all. The cleanup deletes nothing for it, so a full refresh (allArtists) must still // reconcile it. legacyOrphan := model.Artist{ID: "refresh-legacy-orphan", Name: "Legacy Orphan"} - Expect(repo.Put(&legacyOrphan)).To(Succeed()) + Expect(repo.Put(ctx, &legacyOrphan)).To(Succeed()) DeferCleanup(func() { - _ = repo.delete(squirrel.Eq{"id": legacyOrphan.ID}) + _ = repo.delete(ctx, squirrel.Eq{"id": legacyOrphan.ID}) }) Expect(missing(legacyOrphan.ID)).To(BeFalse()) - _, err := repo.RefreshStats(true) + _, err := repo.RefreshStats(ctx, true) Expect(err).ToNot(HaveOccurred()) Expect(missing(legacyOrphan.ID)).To(BeTrue()) @@ -1276,13 +1350,13 @@ var _ = Describe("ArtistRepository", func() { // Helper function to create an artist with proper library association. // This ensures test artists always have library_artist associations to avoid orphaned artists in tests. -func createArtistWithLibrary(repo model.ArtistRepository, artist *model.Artist, libraryID int) error { - err := repo.Put(artist) +func createArtistWithLibrary(ctx context.Context, repo model.ArtistRepository, artist *model.Artist, libraryID int) error { + err := repo.Put(ctx, artist) if err != nil { return err } // Add the artist to the specified library - lr := NewLibraryRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - return lr.AddArtist(libraryID, artist.ID) + lr := NewLibraryRepository(GetDBXBuilder()) + return lr.AddArtist(request.WithUser(ctx, adminUser), libraryID, artist.ID) } diff --git a/persistence/artwork_hydration.go b/persistence/artwork_hydration.go index fa9920f27..76646ec74 100644 --- a/persistence/artwork_hydration.go +++ b/persistence/artwork_hydration.go @@ -59,7 +59,7 @@ func hydrateItemImages(ctx context.Context, db dbx.Builder, kind model.Kind, ids if len(ids) == 0 { return map[string]model.ItemArtworkInfo{} } - infos, err := NewArtworkRepository(ctx, db).GetInfoForItems(kind, ids) + infos, err := NewArtworkRepository(db).GetInfoForItems(ctx, kind, ids) if err != nil { log.Error(ctx, "Failed to hydrate artwork info onto page", "kind", kind, err) return map[string]model.ItemArtworkInfo{} diff --git a/persistence/artwork_hydration_test.go b/persistence/artwork_hydration_test.go index bb9cda04d..89182d6cc 100644 --- a/persistence/artwork_hydration_test.go +++ b/persistence/artwork_hydration_test.go @@ -45,7 +45,7 @@ var _ = Describe("Artwork hydration", func() { var aw model.ArtworkRepository putInfo := func(kind, id, hash string) { - Expect(aw.PutItemArtwork(&model.ItemArtwork{ + Expect(aw.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: kind, ItemID: id, ImageType: model.ImageTypePrimary, Hash: hash, })).To(Succeed()) } @@ -54,19 +54,19 @@ var _ = Describe("Artwork hydration", func() { clearArtworkTables() DeferCleanup(clearArtworkTables) ctx = request.WithUser(log.NewContext(context.Background()), adminUser) - aw = NewArtworkRepository(ctx, GetDBXBuilder()) + aw = NewArtworkRepository(GetDBXBuilder()) }) Describe("albums", func() { var repo model.AlbumRepository - BeforeEach(func() { repo = NewAlbumRepository(ctx, GetDBXBuilder()) }) + BeforeEach(func() { repo = NewAlbumRepository(GetDBXBuilder()) }) It("hydrates the found / known-absent / unresolved states", func() { putInfo("al", albumSgtPeppers.ID, "althash11111111") putInfo("al", albumAbbeyRoad.ID, "") // albumRadioactivity: no row -> unresolved - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) byID := slice.ToMap(all, func(a model.Album) (string, model.Album) { return a.ID, a }) @@ -80,14 +80,14 @@ var _ = Describe("Artwork hydration", func() { It("hydrates Get", func() { putInfo("al", albumSgtPeppers.ID, "gethash22222222") - got, err := repo.Get(albumSgtPeppers.ID) + got, err := repo.Get(ctx, albumSgtPeppers.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.ImageHash).To(Equal("gethash22222222")) }) It("hydrates Search", func() { putInfo("al", albumSgtPeppers.ID, "srchash33333333") - res, err := repo.Search("Peppers", model.QueryOptions{}) + res, err := repo.Search(ctx, "Peppers", model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) Expect(res).ToNot(BeEmpty()) Expect(res[0].ImageHash).To(Equal("srchash33333333")) @@ -97,9 +97,9 @@ var _ = Describe("Artwork hydration", func() { al := albumSgtPeppers al.ImageHash = "shouldnotpersist" al.ImageAbsent = true - Expect(repo.(*albumRepository).Put(&al)).To(Succeed()) + Expect(repo.(*albumRepository).Put(ctx, &al)).To(Succeed()) - got, err := repo.Get(al.ID) + got, err := repo.Get(ctx, al.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.ImageHash).To(BeEmpty()) Expect(got.ImageAbsent).To(BeFalse()) @@ -108,14 +108,14 @@ var _ = Describe("Artwork hydration", func() { Describe("artists", func() { var repo model.ArtistRepository - BeforeEach(func() { repo = NewArtistRepository(ctx, GetDBXBuilder()) }) + BeforeEach(func() { repo = NewArtistRepository(GetDBXBuilder()) }) It("hydrates the found / known-absent / unresolved states", func() { putInfo("ar", artistBeatles.ID, "arhash444444444") putInfo("ar", artistKraftwerk.ID, "") // artistCJK: no row -> unresolved - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) byID := slice.ToMap(all, func(a model.Artist) (string, model.Artist) { return a.ID, a }) @@ -129,14 +129,14 @@ var _ = Describe("Artwork hydration", func() { It("hydrates Get", func() { putInfo("ar", artistBeatles.ID, "arget5555555555") - got, err := repo.Get(artistBeatles.ID) + got, err := repo.Get(ctx, artistBeatles.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.ImageHash).To(Equal("arget5555555555")) }) It("hydrates Search", func() { putInfo("ar", artistBeatles.ID, "arsrch666666666") - res, err := repo.Search("Beatles", model.QueryOptions{}) + res, err := repo.Search(ctx, "Beatles", model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) Expect(res).ToNot(BeEmpty()) Expect(res[0].ImageHash).To(Equal("arsrch666666666")) @@ -145,13 +145,13 @@ var _ = Describe("Artwork hydration", func() { Describe("playlists", func() { var repo model.PlaylistRepository - BeforeEach(func() { repo = NewPlaylistRepository(ctx, GetDBXBuilder()) }) + BeforeEach(func() { repo = NewPlaylistRepository(GetDBXBuilder()) }) It("hydrates the found / known-absent states", func() { putInfo("pl", plsBest.ID, "plhash777777777") putInfo("pl", plsCool.ID, "") - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) byID := slice.ToMap(all, func(p model.Playlist) (string, model.Playlist) { return p.ID, p }) @@ -163,16 +163,16 @@ var _ = Describe("Artwork hydration", func() { It("hydrates Get", func() { putInfo("pl", plsBest.ID, "plget8888888888") - got, err := repo.Get(plsBest.ID) + got, err := repo.Get(ctx, plsBest.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.ImageHash).To(Equal("plget8888888888")) }) It("hydrates the tracks reached through a playlist", func() { - Expect(aw.PutImage(&model.Artwork{Hash: "pltrackhash1234", Mime: "image/jpeg", BlurHash: "LPLBLURhash"})).To(Succeed()) + Expect(aw.PutImage(ctx, &model.Artwork{Hash: "pltrackhash1234", Mime: "image/jpeg", BlurHash: "LPLBLURhash"})).To(Succeed()) putInfo("al", songDayInALife.AlbumID, "pltrackhash1234") - pls, err := repo.GetWithTracks(plsBest.ID, true, false) + pls, err := repo.GetWithTracks(ctx, plsBest.ID, true, false) Expect(err).ToNot(HaveOccurred()) tracks := pls.Tracks Expect(tracks).ToNot(BeEmpty()) @@ -181,7 +181,7 @@ var _ = Describe("Artwork hydration", func() { Expect(byID[songDayInALife.ID].AlbumImage.ImageHash).To(Equal("pltrackhash1234")) Expect(byID[songDayInALife.ID].BlurHash).To(Equal("LPLBLURhash")) - cursor, err := repo.Tracks(plsBest.ID, true).GetCursor() + cursor, err := repo.Tracks(ctx, plsBest.ID, true).GetCursor(ctx) Expect(err).ToNot(HaveOccurred()) var streamed *model.PlaylistTrack for t, err := range cursor { @@ -198,13 +198,13 @@ var _ = Describe("Artwork hydration", func() { Describe("radios", func() { var repo model.RadioRepository - BeforeEach(func() { repo = NewRadioRepository(ctx, GetDBXBuilder()) }) + BeforeEach(func() { repo = NewRadioRepository(GetDBXBuilder()) }) It("hydrates the found / known-absent states", func() { putInfo("ra", radioWithHomePage.ID, "rahash999999999") putInfo("ra", radioWithoutHomePage.ID, "") - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) byID := slice.ToMap(all, func(rd model.Radio) (string, model.Radio) { return rd.ID, rd }) @@ -216,7 +216,7 @@ var _ = Describe("Artwork hydration", func() { It("hydrates Get", func() { putInfo("ra", radioWithHomePage.ID, "ragetaaaaaaaaaa") - got, err := repo.Get(radioWithHomePage.ID) + got, err := repo.Get(ctx, radioWithHomePage.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.ImageHash).To(Equal("ragetaaaaaaaaaa")) }) @@ -232,13 +232,13 @@ var _ = Describe("Artwork hydration", func() { } getByID := func() map[string]model.MediaFile { - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) return slice.ToMap(all, func(mf model.MediaFile) (string, model.MediaFile) { return mf.ID, mf }) } BeforeEach(func() { - repo = NewMediaFileRepository(ctx, GetDBXBuilder()) + repo = NewMediaFileRepository(GetDBXBuilder()) DeferCleanup(configtest.SetupConfig()) conf.Server.EnableMediaFileCoverArt = true }) @@ -293,8 +293,8 @@ var _ = Describe("Artwork hydration", func() { setCover("1001", true) // eligible, resolves its own art -> own-art-wins branch DeferCleanup(func() { setCover("1001", false) }) - Expect(aw.PutImage(&model.Artwork{Hash: "mfh1001blurxxxxx", Mime: "image/jpeg", BlurHash: "LTRACKblur", ThumbHash: "THtrack", Width: 640, Height: 480})).To(Succeed()) - Expect(aw.PutImage(&model.Artwork{Hash: "alh102blurxxxxxx", Mime: "image/jpeg", BlurHash: "LALBUMblur", ThumbHash: "THalbum", Width: 1200, Height: 800})).To(Succeed()) + Expect(aw.PutImage(ctx, &model.Artwork{Hash: "mfh1001blurxxxxx", Mime: "image/jpeg", BlurHash: "LTRACKblur", ThumbHash: "THtrack", Width: 640, Height: 480})).To(Succeed()) + Expect(aw.PutImage(ctx, &model.Artwork{Hash: "alh102blurxxxxxx", Mime: "image/jpeg", BlurHash: "LALBUMblur", ThumbHash: "THalbum", Width: 1200, Height: 800})).To(Succeed()) putInfo("mf", "1001", "mfh1001blurxxxxx") putInfo("al", "102", "alh102blurxxxxxx") // 1002's album: single-disc inheritance branch @@ -374,7 +374,7 @@ var _ = Describe("Artwork hydration", func() { It("hydrates Search", func() { putInfo("al", "101", "alsrchhhhhhhhhhh") - res, err := repo.Search("A Day In A Life", model.QueryOptions{}) + res, err := repo.Search(ctx, "A Day In A Life", model.QueryOptions{}) Expect(err).ToNot(HaveOccurred()) Expect(res).ToNot(BeEmpty()) Expect(res[0].ImageHash).To(Equal("alsrchhhhhhhhhhh")) @@ -397,9 +397,9 @@ var _ = Describe("Artwork hydration", func() { } BeforeEach(func() { - albumRepo = NewAlbumRepository(ctx, GetDBXBuilder()) - artistRepo = NewArtistRepository(ctx, GetDBXBuilder()) - playlistRepo = NewPlaylistRepository(ctx, GetDBXBuilder()) + albumRepo = NewAlbumRepository(GetDBXBuilder()) + artistRepo = NewArtistRepository(GetDBXBuilder()) + playlistRepo = NewPlaylistRepository(GetDBXBuilder()) // Other specs leave rows behind, so scope every cursor spec to the fixtures. onlyAlbums = squirrel.Eq{"album.id": []string{albumSgtPeppers.ID, albumAbbeyRoad.ID, albumRadioactivity.ID, albumMultiDisc.ID, albumCJK.ID, albumPunctuation.ID}} @@ -408,8 +408,8 @@ var _ = Describe("Artwork hydration", func() { // Both fixture playlists share an owner, leaving the owner_name sort a single value to // order by; this one is also private, which the non-admin visibility spec needs. foreign := model.Playlist{Name: "Foreign", OwnerID: thirdUser.ID, OwnerName: thirdUser.UserName} - Expect(playlistRepo.Put(&foreign)).To(Succeed()) - DeferCleanup(func() { Expect(playlistRepo.Delete(foreign.ID)).To(Succeed()) }) + Expect(playlistRepo.Put(ctx, &foreign)).To(Succeed()) + DeferCleanup(func() { Expect(playlistRepo.Delete(ctx, foreign.ID)).To(Succeed()) }) onlyPlaylists = squirrel.Eq{"playlist.id": []string{plsBest.ID, plsCool.ID, foreign.ID}} // The suite annotates a single album and artist, leaving the annotation-backed sorts @@ -417,7 +417,7 @@ var _ = Describe("Artwork hydration", func() { seedAnnotations("album", albumSgtPeppers.ID, albumAbbeyRoad.ID) seedAnnotations("artist", artistKraftwerk.ID, artistCJK.ID) - Expect(aw.PutImage(&model.Artwork{ + Expect(aw.PutImage(ctx, &model.Artwork{ Hash: "curhash11111111", Mime: "image/jpeg", BlurHash: "LEHV6nWB2yk8", })).To(Succeed()) putInfo("al", albumSgtPeppers.ID, "curhash11111111") @@ -428,10 +428,10 @@ var _ = Describe("Artwork hydration", func() { It("hydrates every streamed album, like GetAll", func() { opts := model.QueryOptions{Sort: "name", Filters: onlyAlbums} - want, err := albumRepo.GetAll(opts) + want, err := albumRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - got := collectCursor(albumRepo.GetCursor(opts)) + got := collectCursor(albumRepo.GetCursor(ctx, opts)) Expect(got).To(ConsistOf(want)) byID := map[string]model.Album{} @@ -447,10 +447,10 @@ var _ = Describe("Artwork hydration", func() { It("hydrates every streamed artist, like GetAll", func() { opts := model.QueryOptions{Sort: "name", Filters: onlyArtists} - want, err := artistRepo.GetAll(opts) + want, err := artistRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - got := collectCursor(artistRepo.GetCursor(opts)) + got := collectCursor(artistRepo.GetCursor(ctx, opts)) Expect(got).To(ConsistOf(want)) Expect(slice.Map(got, func(a model.Artist) string { return a.ImageHash })). @@ -459,10 +459,10 @@ var _ = Describe("Artwork hydration", func() { It("hydrates every streamed playlist, like GetAll", func() { opts := model.QueryOptions{Sort: "name", Filters: onlyPlaylists} - want, err := playlistRepo.GetAll(opts) + want, err := playlistRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - got := collectCursor(playlistRepo.GetCursor(opts)) + got := collectCursor(playlistRepo.GetCursor(ctx, opts)) Expect(got).To(ConsistOf(want)) Expect(slice.Map(got, func(p model.Playlist) string { return p.ImageHash })). @@ -471,11 +471,11 @@ var _ = Describe("Artwork hydration", func() { It("honors Max and Offset exactly once", func() { opts := model.QueryOptions{Sort: "name", Filters: onlyAlbums, Max: 2, Offset: 1} - want, err := albumRepo.GetAll(opts) + want, err := albumRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).To(HaveLen(2)) - got := collectCursor(albumRepo.GetCursor(opts)) + got := collectCursor(albumRepo.GetCursor(ctx, opts)) Expect(slice.Map(got, func(a model.Album) string { return a.ID })). To(Equal(slice.Map(want, func(a model.Album) string { return a.ID }))) @@ -485,11 +485,11 @@ var _ = Describe("Artwork hydration", func() { DescribeTable("orders albums like GetAll", func(opts model.QueryOptions, key func(model.Album) string) { opts = scoped(opts, onlyAlbums) - want, err := albumRepo.GetAll(opts) + want, err := albumRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - got := collectCursor(albumRepo.GetCursor(opts)) + got := collectCursor(albumRepo.GetCursor(ctx, opts)) Expect(slice.Map(got, key)).To(Equal(slice.Map(want, key))) Expect(slice.Map(got, func(a model.Album) string { return a.ID })). @@ -517,11 +517,11 @@ var _ = Describe("Artwork hydration", func() { DescribeTable("orders artists like GetAll", func(opts model.QueryOptions, key func(model.Artist) string) { opts = scoped(opts, onlyArtists) - want, err := artistRepo.GetAll(opts) + want, err := artistRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - got := collectCursor(artistRepo.GetCursor(opts)) + got := collectCursor(artistRepo.GetCursor(ctx, opts)) Expect(slice.Map(got, key)).To(Equal(slice.Map(want, key))) Expect(slice.Map(got, func(a model.Artist) string { return a.ID })). @@ -543,11 +543,11 @@ var _ = Describe("Artwork hydration", func() { DescribeTable("orders playlists like GetAll", func(opts model.QueryOptions, key func(model.Playlist) string) { opts = scoped(opts, onlyPlaylists) - want, err := playlistRepo.GetAll(opts) + want, err := playlistRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - got := collectCursor(playlistRepo.GetCursor(opts)) + got := collectCursor(playlistRepo.GetCursor(ctx, opts)) Expect(slice.Map(got, key)).To(Equal(slice.Map(want, key))) Expect(slice.Map(got, func(p model.Playlist) string { return p.ID })). @@ -563,10 +563,10 @@ var _ = Describe("Artwork hydration", func() { It("streams every album exactly once when sorted randomly", func() { opts := model.QueryOptions{Sort: "random", Filters: onlyAlbums} - want, err := albumRepo.GetAll(opts) + want, err := albumRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - got := collectCursor(albumRepo.GetCursor(opts)) + got := collectCursor(albumRepo.GetCursor(ctx, opts)) Expect(slice.Map(got, func(a model.Album) string { return a.ID })). To(ConsistOf(slice.Map(want, func(a model.Album) string { return a.ID }))) @@ -574,16 +574,16 @@ var _ = Describe("Artwork hydration", func() { It("keeps a non-admin from streaming another user's private playlists", func() { otherCtx := request.WithUser(log.NewContext(context.Background()), regularUser) - repo := NewPlaylistRepository(otherCtx, GetDBXBuilder()) + repo := NewPlaylistRepository(GetDBXBuilder()) opts := model.QueryOptions{Sort: "name", Filters: onlyPlaylists} // Both phases must filter on their own: the id pre-pass and the chunk fetch. - Expect(repo.(*playlistRepository).getAllIDs(opts)).To(ConsistOf(plsBest.ID)) - all, err := repo.GetAll(model.QueryOptions{Filters: onlyPlaylists}) + Expect(repo.(*playlistRepository).getAllIDs(otherCtx, opts)).To(ConsistOf(plsBest.ID)) + all, err := repo.GetAll(otherCtx, model.QueryOptions{Filters: onlyPlaylists}) Expect(err).ToNot(HaveOccurred()) Expect(slice.Map(all, func(p model.Playlist) string { return p.ID })).To(ConsistOf(plsBest.ID)) - got := collectCursor(repo.GetCursor(opts)) + got := collectCursor(repo.GetCursor(otherCtx, opts)) Expect(slice.Map(got, func(p model.Playlist) string { return p.Name })). To(ConsistOf(plsBest.Name)) @@ -595,7 +595,7 @@ var _ = Describe("Artwork hydration", func() { var onlySongs squirrel.Eq BeforeEach(func() { - mfRepo = NewMediaFileRepository(ctx, GetDBXBuilder()) + mfRepo = NewMediaFileRepository(GetDBXBuilder()) putInfo("al", albumSgtPeppers.ID, "curhash11111111") // Distinct titles only: other fixture songs share titles (e.g. "Antenna" x3), which // would make the positional comparisons against GetAll pass by tie-order coincidence. @@ -605,11 +605,11 @@ var _ = Describe("Artwork hydration", func() { It("hydrates artwork onto every streamed track, unlike GetCursor", func() { opts := model.QueryOptions{Sort: "title", Filters: onlySongs} - want, err := mfRepo.GetAll(opts) + want, err := mfRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - cursor, err := mfRepo.GetCursorWithArtwork(opts) + cursor, err := mfRepo.GetCursorWithArtwork(ctx, opts) Expect(err).ToNot(HaveOccurred()) var got model.MediaFiles cursor(func(mf model.MediaFile, err error) bool { @@ -631,7 +631,7 @@ var _ = Describe("Artwork hydration", func() { }) It("leaves the scanner's GetCursor unhydrated", func() { - cursor, err := mfRepo.GetCursor(model.QueryOptions{Sort: "title"}) + cursor, err := mfRepo.GetCursor(ctx, model.QueryOptions{Sort: "title"}) Expect(err).ToNot(HaveOccurred()) var seen int cursor(func(mf model.MediaFile, err error) bool { @@ -645,11 +645,11 @@ var _ = Describe("Artwork hydration", func() { It("streams the same ids in the same order as GetAll", func() { opts := model.QueryOptions{Sort: "title", Filters: onlySongs} - want, err := mfRepo.GetAll(opts) + want, err := mfRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - got := collectCursor(mfRepo.GetCursorWithArtwork(opts)) + got := collectCursor(mfRepo.GetCursorWithArtwork(ctx, opts)) Expect(slice.Map(got, func(mf model.MediaFile) string { return mf.ID })). To(Equal(slice.Map(want, func(mf model.MediaFile) string { return mf.ID }))) @@ -657,11 +657,11 @@ var _ = Describe("Artwork hydration", func() { It("honors Max and Offset exactly once", func() { opts := model.QueryOptions{Sort: "title", Filters: onlySongs, Max: 2, Offset: 1} - want, err := mfRepo.GetAll(opts) + want, err := mfRepo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).To(HaveLen(2)) - got := collectCursor(mfRepo.GetCursorWithArtwork(opts)) + got := collectCursor(mfRepo.GetCursorWithArtwork(ctx, opts)) Expect(slice.Map(got, func(mf model.MediaFile) string { return mf.ID })). To(Equal(slice.Map(want, func(mf model.MediaFile) string { return mf.ID }))) diff --git a/persistence/artwork_queue_repository.go b/persistence/artwork_queue_repository.go index 88b6f6f80..321fe8a95 100644 --- a/persistence/artwork_queue_repository.go +++ b/persistence/artwork_queue_repository.go @@ -25,17 +25,16 @@ type artworkQueueRepository struct { sqlRepository } -func NewArtworkQueueRepository(ctx context.Context, db dbx.Builder) model.ArtworkQueueRepository { +func NewArtworkQueueRepository(db dbx.Builder) model.ArtworkQueueRepository { r := &artworkQueueRepository{} - r.ctx = ctx r.db = db r.tableName = "artwork_queue" return r } -func (r *artworkQueueRepository) Get(kind model.Kind, id, imageType string) (*model.ArtworkQueueItem, error) { +func (r *artworkQueueRepository) Get(ctx context.Context, kind model.Kind, id, imageType string) (*model.ArtworkQueueItem, error) { var res model.ArtworkQueueItem - err := r.queryOne(Select("*").From(r.tableName). + err := r.queryOne(ctx, Select("*").From(r.tableName). Where(Eq{"item_kind": kind.Prefix(), "item_id": id, "image_type": imageType}), &res) if err != nil { return nil, err @@ -45,30 +44,30 @@ func (r *artworkQueueRepository) Get(kind model.Kind, id, imageType string) (*mo // Enqueue starts a fresh lifecycle: it resets enqueued_at (so a fresh request does not inherit an old // row's spent retry budget) and clears trace (so explain does not show a prior failure at attempts 0). -func (r *artworkQueueRepository) Enqueue(items ...model.ArtworkQueueItem) error { - return r.enqueue(`ON CONFLICT (item_kind, item_id, image_type) DO UPDATE SET +func (r *artworkQueueRepository) Enqueue(ctx context.Context, items ...model.ArtworkQueueItem) error { + return r.enqueue(ctx, `ON CONFLICT (item_kind, item_id, image_type) DO UPDATE SET priority = MAX(priority, excluded.priority), retry_at = excluded.retry_at, attempts = 0, enqueued_at = excluded.enqueued_at, trace = '[]'`, items) } -func (r *artworkQueueRepository) EnqueuePreservingBackoff(items ...model.ArtworkQueueItem) error { - return r.enqueue(`ON CONFLICT (item_kind, item_id, image_type) DO UPDATE SET +func (r *artworkQueueRepository) EnqueuePreservingBackoff(ctx context.Context, items ...model.ArtworkQueueItem) error { + return r.enqueue(ctx, `ON CONFLICT (item_kind, item_id, image_type) DO UPDATE SET priority = MAX(priority, excluded.priority)`, items) } -func (r *artworkQueueRepository) EnqueueAllMissing(kind model.Kind, priority int) (int64, error) { +func (r *artworkQueueRepository) EnqueueAllMissing(ctx context.Context, kind model.Kind, priority int) (int64, error) { entityTable, ok := artworkOwnerTables[kind] if !ok { return 0, fmt.Errorf("artwork queue: no entity table for kind %q", kind.Prefix()) } now := time.Now() - return r.insertIfNotQueued("", `SELECT ?, id, ?, ?, 0, ?, ? + return r.insertIfNotQueued(ctx, "", `SELECT ?, id, ?, ?, 0, ?, ? FROM `+entityTable+` WHERE id NOT IN (SELECT item_id FROM `+itemArtworkTable+` WHERE item_kind = ?)`, kind.Prefix(), model.ImageTypePrimary, priority, now, now, kind.Prefix()) } -func (r *artworkQueueRepository) EnqueueIfMissing(items ...model.ArtworkQueueItem) error { +func (r *artworkQueueRepository) EnqueueIfMissing(ctx context.Context, items ...model.ArtworkQueueItem) error { now := time.Now() for chunk := range slices.Chunk(items, enqueueChunkSize) { rows := make([]string, 0, len(chunk)) @@ -78,7 +77,7 @@ func (r *artworkQueueRepository) EnqueueIfMissing(items ...model.ArtworkQueueIte args = append(args, it.ItemKind, it.ItemID, cmp.Or(it.ImageType, model.ImageTypePrimary), it.Priority) } args = append(args, now, now) - _, err := r.insertIfNotQueued( + _, err := r.insertIfNotQueued(ctx, `WITH new_items(item_kind, item_id, image_type, priority) AS (VALUES `+strings.Join(rows, ",")+`) `, `SELECT n.item_kind, n.item_id, n.image_type, n.priority, 0, ?, ? FROM new_items n @@ -97,8 +96,8 @@ func (r *artworkQueueRepository) EnqueueIfMissing(items ...model.ArtworkQueueIte const skipIfQueued = ` ON CONFLICT (item_kind, item_id, image_type) DO NOTHING` // insertIfNotQueued inserts the rows selected by the given SQL, optionally prefixed by a CTE. -func (r *artworkQueueRepository) insertIfNotQueued(with, sql string, args ...any) (int64, error) { - return r.executeSQL(Expr(with+`INSERT INTO `+r.tableName+ +func (r *artworkQueueRepository) insertIfNotQueued(ctx context.Context, with, sql string, args ...any) (int64, error) { + return r.executeSQL(ctx, Expr(with+`INSERT INTO `+r.tableName+ ` (`+strings.Join(enqueueColumns, ", ")+`) `+sql+skipIfQueued, args...)) } @@ -121,16 +120,16 @@ func artworkSourceFilter(kind model.Kind, sources []string) Sqlizer { return append(f, match) } -func (r *artworkQueueRepository) CountBySource(kind model.Kind, sources []string) (int64, error) { +func (r *artworkQueueRepository) CountBySource(ctx context.Context, kind model.Kind, sources []string) (int64, error) { var res struct{ Count int64 } - err := r.queryOne(Select("count(*) as count").From(itemArtworkTable). + err := r.queryOne(ctx, Select("count(*) as count").From(itemArtworkTable). Where(artworkSourceFilter(kind, sources)), &res) return res.Count, err } -func (r *artworkQueueRepository) SourcesInUse(kind model.Kind) ([]string, error) { +func (r *artworkQueueRepository) SourcesInUse(ctx context.Context, kind model.Kind) ([]string, error) { var res []struct{ Source string } - err := r.queryAll(Select("distinct source").From(itemArtworkTable). + err := r.queryAll(ctx, Select("distinct source").From(itemArtworkTable). Where(Eq{"item_kind": kind.Prefix()}), &res) if err != nil { return nil, err @@ -140,15 +139,15 @@ func (r *artworkQueueRepository) SourcesInUse(kind model.Kind) ([]string, error) // EnqueueBySource deliberately leaves item_artwork alone: clearing state in bulk would blank the // library's artwork until every item is resolved again. -func (r *artworkQueueRepository) EnqueueBySource(kind model.Kind, sources []string, priority int) (int64, error) { +func (r *artworkQueueRepository) EnqueueBySource(ctx context.Context, kind model.Kind, sources []string, priority int) (int64, error) { now := time.Now() sel := Select("item_kind", "item_id", "image_type"). Column(Expr("?", priority)).Column("0").Column(Expr("?", now)).Column(Expr("?", now)). From(itemArtworkTable).Where(artworkSourceFilter(kind, sources)) - return r.executeSQL(Insert(r.tableName).Columns(enqueueColumns...).Select(sel).Suffix(skipIfQueued)) + return r.executeSQL(ctx, Insert(r.tableName).Columns(enqueueColumns...).Select(sel).Suffix(skipIfQueued)) } -func (r *artworkQueueRepository) enqueue(conflict string, items []model.ArtworkQueueItem) error { +func (r *artworkQueueRepository) enqueue(ctx context.Context, conflict string, items []model.ArtworkQueueItem) error { now := time.Now() for chunk := range slices.Chunk(items, enqueueChunkSize) { ins := Insert(r.tableName).Columns(enqueueColumns...) @@ -156,14 +155,14 @@ func (r *artworkQueueRepository) enqueue(conflict string, items []model.ArtworkQ ins = ins.Values(it.ItemKind, it.ItemID, cmp.Or(it.ImageType, model.ImageTypePrimary), it.Priority, 0, now, now) } ins = ins.Suffix(conflict) - if _, err := r.executeSQL(ins); err != nil { + if _, err := r.executeSQL(ctx, ins); err != nil { return err } } return nil } -func (r *artworkQueueRepository) DequeueBatch(n int, kinds ...string) ([]model.ArtworkQueueItem, error) { +func (r *artworkQueueRepository) DequeueBatch(ctx context.Context, n int, kinds ...string) ([]model.ArtworkQueueItem, error) { sel := Select(enqueueColumns...).From(r.tableName). Where(LtOrEq{"retry_at": time.Now()}). OrderBy("priority DESC", "enqueued_at ASC"). @@ -172,26 +171,26 @@ func (r *artworkQueueRepository) DequeueBatch(n int, kinds ...string) ([]model.A sel = sel.Where(Eq{"item_kind": kinds}) } var res []model.ArtworkQueueItem - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) return res, err } -func (r *artworkQueueRepository) MarkFailedIfUnchanged(kind, id, imageType string, seenRetryAt, retryAt time.Time, trace string) error { +func (r *artworkQueueRepository) MarkFailedIfUnchanged(ctx context.Context, kind, id, imageType string, seenRetryAt, retryAt time.Time, trace string) error { upd := Update(r.tableName). Set("attempts", Expr("attempts + 1")). Set("retry_at", retryAt). Set("trace", trace). Where(Eq{"item_kind": kind, "item_id": id, "image_type": imageType, "retry_at": seenRetryAt}) - _, err := r.executeSQL(upd) + _, err := r.executeSQL(ctx, upd) return err } -func (r *artworkQueueRepository) DeleteIfUnchanged(kind, id, imageType string, retryAt time.Time) error { - return r.delete(Eq{"item_kind": kind, "item_id": id, "image_type": imageType, "retry_at": retryAt}) +func (r *artworkQueueRepository) DeleteIfUnchanged(ctx context.Context, kind, id, imageType string, retryAt time.Time) error { + return r.delete(ctx, Eq{"item_kind": kind, "item_id": id, "image_type": imageType, "retry_at": retryAt}) } -func (r *artworkQueueRepository) PurgeDangling() (int64, error) { - return purgeDangling(r.sqlRepository) +func (r *artworkQueueRepository) PurgeDangling(ctx context.Context) (int64, error) { + return purgeDangling(ctx, r.sqlRepository) } // artworkQueueFilter returns no conditions for an empty filter, so an unfiltered DELETE keeps @@ -208,28 +207,28 @@ func artworkQueueFilter(kinds []model.Kind, priorities []int) And { } // CountQueued shares its filter with PurgeQueued, so a preview cannot count rows the delete misses. -func (r *artworkQueueRepository) CountQueued(kinds []model.Kind, priorities []int) ([]model.ArtworkQueueStat, error) { +func (r *artworkQueueRepository) CountQueued(ctx context.Context, kinds []model.Kind, priorities []int) ([]model.ArtworkQueueStat, error) { sel := Select("item_kind", "priority", "count(*) as count").From(r.tableName). GroupBy("item_kind", "priority").OrderBy("item_kind", "priority desc") if f := artworkQueueFilter(kinds, priorities); len(f) > 0 { sel = sel.Where(f) } var res []model.ArtworkQueueStat - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) return res, err } -func (r *artworkQueueRepository) PurgeQueued(kinds []model.Kind, priorities []int) (int64, error) { +func (r *artworkQueueRepository) PurgeQueued(ctx context.Context, kinds []model.Kind, priorities []int) (int64, error) { del := Delete(r.tableName) if f := artworkQueueFilter(kinds, priorities); len(f) > 0 { del = del.Where(f) } - return r.executeSQL(del) + return r.executeSQL(ctx, del) } -func (r *artworkQueueRepository) Count() (int64, error) { +func (r *artworkQueueRepository) Count(ctx context.Context) (int64, error) { var res struct{ Count int64 } - err := r.queryOne(Select("count(*) as count").From(r.tableName), &res) + err := r.queryOne(ctx, Select("count(*) as count").From(r.tableName), &res) return res.Count, err } diff --git a/persistence/artwork_queue_repository_test.go b/persistence/artwork_queue_repository_test.go index 0481d4193..6ce4be8e2 100644 --- a/persistence/artwork_queue_repository_test.go +++ b/persistence/artwork_queue_repository_test.go @@ -14,6 +14,7 @@ import ( var _ = Describe("ArtworkQueueRepository", func() { var repo model.ArtworkQueueRepository + var ctx context.Context item := func(kind, id string, prio int) model.ArtworkQueueItem { return model.ArtworkQueueItem{ItemKind: kind, ItemID: id, @@ -25,7 +26,7 @@ var _ = Describe("ArtworkQueueRepository", func() { backOff := func(kind, id string, retryAt time.Time) { GinkgoHelper() r := repo.(*artworkQueueRepository) - _, err := r.executeSQL(squirrel.Update(r.tableName). + _, err := r.executeSQL(ctx, squirrel.Update(r.tableName). Set("attempts", squirrel.Expr("attempts + 1")). Set("retry_at", retryAt). Where(squirrel.Eq{"item_kind": kind, "item_id": id, "image_type": model.ImageTypePrimary})) @@ -35,30 +36,31 @@ var _ = Describe("ArtworkQueueRepository", func() { remove := func(kind, id string) { GinkgoHelper() r := repo.(*artworkQueueRepository) - Expect(r.delete(squirrel.Eq{"item_kind": kind, "item_id": id, "image_type": model.ImageTypePrimary})).To(Succeed()) + Expect(r.delete(ctx, squirrel.Eq{"item_kind": kind, "item_id": id, "image_type": model.ImageTypePrimary})).To(Succeed()) } BeforeEach(func() { + ctx = GinkgoT().Context() clearArtworkTables() DeferCleanup(clearArtworkTables) - repo = NewArtworkQueueRepository(context.Background(), GetDBXBuilder()) + repo = NewArtworkQueueRepository(GetDBXBuilder()) }) It("enqueues and dequeues by priority then FIFO", func() { - Expect(repo.Enqueue(item("al", "low", model.ArtworkPriorityBackfill))).To(Succeed()) - Expect(repo.Enqueue(item("ar", "high", model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "low", model.ArtworkPriorityBackfill))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("ar", "high", model.ArtworkPriorityBump))).To(Succeed()) - got, err := repo.DequeueBatch(10) + got, err := repo.DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(got).To(HaveLen(2)) Expect(got[0].ItemID).To(Equal("high")) }) It("Get returns a queued row, including one still backing off", func() { - Expect(repo.Enqueue(item("ar", "g1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("ar", "g1", model.ArtworkPriorityScan))).To(Succeed()) backOff("ar", "g1", time.Now().Add(time.Hour)) - got, err := repo.Get(model.KindArtistArtwork, "g1", model.ImageTypePrimary) + got, err := repo.Get(ctx, model.KindArtistArtwork, "g1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(got.Priority).To(Equal(model.ArtworkPriorityScan)) Expect(got.Attempts).To(Equal(1)) @@ -66,114 +68,114 @@ var _ = Describe("ArtworkQueueRepository", func() { }) It("Get reports ErrNotFound when the item is not queued", func() { - _, err := repo.Get(model.KindArtistArtwork, "nope", model.ImageTypePrimary) + _, err := repo.Get(ctx, model.KindArtistArtwork, "nope", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) It("keeps the higher priority on duplicate enqueue", func() { - Expect(repo.Enqueue(item("al", "a1", model.ArtworkPriorityBump))).To(Succeed()) - Expect(repo.Enqueue(item("al", "a1", model.ArtworkPriorityBackfill))).To(Succeed()) - got, _ := repo.DequeueBatch(10) + Expect(repo.Enqueue(ctx, item("al", "a1", model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "a1", model.ArtworkPriorityBackfill))).To(Succeed()) + got, _ := repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) Expect(got[0].Priority).To(Equal(model.ArtworkPriorityBump)) }) It("EnqueuePreservingBackoff raises priority without resetting a backing-off row's retry_at", func() { - Expect(repo.Enqueue(item("al", "b1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "b1", model.ArtworkPriorityScan))).To(Succeed()) backOff("al", "b1", time.Now().Add(time.Hour)) - Expect(repo.DequeueBatch(10)).To(BeEmpty()) + Expect(repo.DequeueBatch(ctx, 10)).To(BeEmpty()) - Expect(repo.EnqueuePreservingBackoff(item("al", "b1", model.ArtworkPriorityBump))).To(Succeed()) - Expect(repo.DequeueBatch(10)).To(BeEmpty(), "bump must not reset retry_at") + Expect(repo.EnqueuePreservingBackoff(ctx, item("al", "b1", model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.DequeueBatch(ctx, 10)).To(BeEmpty(), "bump must not reset retry_at") // Enqueue (scan/manual), by contrast, resets retry_at and makes it eligible now. - Expect(repo.Enqueue(item("al", "b1", model.ArtworkPriorityScan))).To(Succeed()) - got, _ := repo.DequeueBatch(10) + Expect(repo.Enqueue(ctx, item("al", "b1", model.ArtworkPriorityScan))).To(Succeed()) + got, _ := repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) Expect(got[0].Priority).To(Equal(model.ArtworkPriorityBump), "bump's higher priority is preserved") }) It("EnqueuePreservingBackoff inserts a brand-new row eligible immediately", func() { - Expect(repo.EnqueuePreservingBackoff(item("ar", "n1", model.ArtworkPriorityBump))).To(Succeed()) - got, _ := repo.DequeueBatch(10) + Expect(repo.EnqueuePreservingBackoff(ctx, item("ar", "n1", model.ArtworkPriorityBump))).To(Succeed()) + got, _ := repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) Expect(got[0].ItemID).To(Equal("n1")) }) It("hides failed items until retry_at", func() { - Expect(repo.Enqueue(item("al", "f1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "f1", model.ArtworkPriorityScan))).To(Succeed()) backOff("al", "f1", time.Now().Add(time.Hour)) - got, err := repo.DequeueBatch(10) + got, err := repo.DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(got).To(BeEmpty()) backOff("al", "f1", time.Now().Add(-time.Minute)) - got, _ = repo.DequeueBatch(10) + got, _ = repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) Expect(got[0].Attempts).To(Equal(2)) }) It("MarkFailedIfUnchanged applies backoff only while retry_at is unchanged", func() { - Expect(repo.Enqueue(item("al", "m1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "m1", model.ArtworkPriorityScan))).To(Succeed()) // Anchor retry_at in the past so it can never collide with the re-enqueue's now. backOff("al", "m1", time.Now().Add(-time.Hour)) - got, err := repo.DequeueBatch(10) + got, err := repo.DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(got).To(HaveLen(1)) original := got[0].RetryAt // A concurrent scan re-enqueues, resetting retry_at to now. - Expect(repo.Enqueue(item("al", "m1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "m1", model.ArtworkPriorityScan))).To(Succeed()) future := time.Now().Add(48 * time.Hour) - Expect(repo.MarkFailedIfUnchanged("al", "m1", model.ImageTypePrimary, original, future, "[]")).To(Succeed()) - got, _ = repo.DequeueBatch(10) + Expect(repo.MarkFailedIfUnchanged(ctx, "al", "m1", model.ImageTypePrimary, original, future, "[]")).To(Succeed()) + got, _ = repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1), "the fresh re-enqueue stays immediately eligible") Expect(got[0].Attempts).To(BeZero(), "re-enqueue clears attempts, and the stale failure must not bump them") current := got[0].RetryAt - Expect(repo.MarkFailedIfUnchanged("al", "m1", model.ImageTypePrimary, current, future, `[{"c":"read","o":"error"}]`)).To(Succeed()) - got, _ = repo.DequeueBatch(10) + Expect(repo.MarkFailedIfUnchanged(ctx, "al", "m1", model.ImageTypePrimary, current, future, `[{"c":"read","o":"error"}]`)).To(Succeed()) + got, _ = repo.DequeueBatch(ctx, 10) Expect(got).To(BeEmpty(), "backed-off row is hidden until the future retry_at") - all, _ := repo.Count() + all, _ := repo.Count(ctx) Expect(all).To(Equal(int64(1))) }) It("Enqueue clears a prior lifecycle's failure trace; EnqueuePreservingBackoff keeps it", func() { // Fail an attempt so the queue row carries a failure trace. - Expect(repo.Enqueue(item("al", "t1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "t1", model.ArtworkPriorityScan))).To(Succeed()) backOff("al", "t1", time.Now().Add(-time.Hour)) - got, _ := repo.DequeueBatch(10) + got, _ := repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) future := time.Now().Add(48 * time.Hour) - Expect(repo.MarkFailedIfUnchanged("al", "t1", model.ImageTypePrimary, got[0].RetryAt, future, `[{"c":"read","o":"error"}]`)).To(Succeed()) + Expect(repo.MarkFailedIfUnchanged(ctx, "al", "t1", model.ImageTypePrimary, got[0].RetryAt, future, `[{"c":"read","o":"error"}]`)).To(Succeed()) // A continuation of the same lifecycle must retain the trace. - Expect(repo.EnqueuePreservingBackoff(item("al", "t1", model.ArtworkPriorityBump))).To(Succeed()) - kept, err := repo.Get(model.KindAlbumArtwork, "t1", model.ImageTypePrimary) + Expect(repo.EnqueuePreservingBackoff(ctx, item("al", "t1", model.ArtworkPriorityBump))).To(Succeed()) + kept, err := repo.Get(ctx, model.KindAlbumArtwork, "t1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(kept.Trace).To(Equal(`[{"c":"read","o":"error"}]`)) // A fresh Enqueue resets attempts to 0, so the stale failure trace must be cleared with it. - Expect(repo.Enqueue(item("al", "t1", model.ArtworkPriorityScan))).To(Succeed()) - fresh, err := repo.Get(model.KindAlbumArtwork, "t1", model.ImageTypePrimary) + Expect(repo.Enqueue(ctx, item("al", "t1", model.ArtworkPriorityScan))).To(Succeed()) + fresh, err := repo.Get(ctx, model.KindAlbumArtwork, "t1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(fresh.Attempts).To(BeZero()) Expect(fresh.Trace).To(Equal("[]"), "a fresh lifecycle has no last-attempt trace") }) It("Enqueue restarts the retry budget an existing row had spent", func() { - Expect(repo.Enqueue(item("al", "e1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "e1", model.ArtworkPriorityScan))).To(Succeed()) backOff("al", "e1", time.Now().Add(-time.Hour)) stale := time.Now().Add(-48 * time.Hour) _, err := GetDBXBuilder().NewQuery("UPDATE artwork_queue SET enqueued_at = {:t} WHERE item_id = 'e1'"). Bind(dbx.Params{"t": stale}).Execute() Expect(err).ToNot(HaveOccurred()) - Expect(repo.Enqueue(item("al", "e1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "e1", model.ArtworkPriorityScan))).To(Succeed()) - got, err := repo.DequeueBatch(10) + got, err := repo.DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(got).To(HaveLen(1)) Expect(got[0].EnqueuedAt).To(BeTemporally("~", time.Now(), time.Minute), @@ -182,40 +184,40 @@ var _ = Describe("ArtworkQueueRepository", func() { }) It("deletes on completion and counts", func() { - Expect(repo.Enqueue(item("al", "c1", 0))).To(Succeed()) - n, _ := repo.Count() + Expect(repo.Enqueue(ctx, item("al", "c1", 0))).To(Succeed()) + n, _ := repo.Count(ctx) Expect(n).To(Equal(int64(1))) remove("al", "c1") - n, _ = repo.Count() + n, _ = repo.Count(ctx) Expect(n).To(BeZero()) }) It("DeleteIfUnchanged deletes only while retry_at is unchanged", func() { - Expect(repo.Enqueue(item("al", "d1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "d1", model.ArtworkPriorityScan))).To(Succeed()) // Anchor retry_at in the past so it can never collide with the re-enqueue's now. backOff("al", "d1", time.Now().Add(-time.Hour)) - got, err := repo.DequeueBatch(10) + got, err := repo.DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(got).To(HaveLen(1)) original := got[0].RetryAt // A concurrent scan re-enqueues, resetting retry_at to now. - Expect(repo.Enqueue(item("al", "d1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "d1", model.ArtworkPriorityScan))).To(Succeed()) // Deleting with the stale retry_at is a no-op: the re-enqueued row survives. - Expect(repo.DeleteIfUnchanged("al", "d1", model.ImageTypePrimary, original)).To(Succeed()) - n, _ := repo.Count() + Expect(repo.DeleteIfUnchanged(ctx, "al", "d1", model.ImageTypePrimary, original)).To(Succeed()) + n, _ := repo.Count(ctx) Expect(n).To(Equal(int64(1))) - got, _ = repo.DequeueBatch(10) + got, _ = repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) - Expect(repo.DeleteIfUnchanged("al", "d1", model.ImageTypePrimary, got[0].RetryAt)).To(Succeed()) - n, _ = repo.Count() + Expect(repo.DeleteIfUnchanged(ctx, "al", "d1", model.ImageTypePrimary, got[0].RetryAt)).To(Succeed()) + n, _ = repo.Count(ctx) Expect(n).To(BeZero()) }) It("purges queue rows whose entity no longer exists, per kind", func() { - Expect(repo.Enqueue( + Expect(repo.Enqueue(ctx, item("al", albumSgtPeppers.ID, model.ArtworkPriorityScan), item("al", "no-such-album", model.ArtworkPriorityScan), item("ar", artistKraftwerk.ID, model.ArtworkPriorityScan), @@ -228,25 +230,25 @@ var _ = Describe("ArtworkQueueRepository", func() { item("mf", "no-such-mediafile", model.ArtworkPriorityScan), )).To(Succeed()) - purged, err := repo.PurgeDangling() + purged, err := repo.PurgeDangling(ctx) Expect(err).ToNot(HaveOccurred()) Expect(purged).To(Equal(int64(5))) - got, _ := repo.DequeueBatch(100) + got, _ := repo.DequeueBatch(ctx, 100) ids := slice.Map(got, func(it model.ArtworkQueueItem) string { return it.ItemID }) Expect(ids).To(ConsistOf(albumSgtPeppers.ID, artistKraftwerk.ID, plsBest.ID, radioWithHomePage.ID, songDayInALife.ID)) }) It("enqueues entities that have no item_artwork row at all", func() { - awRepo := NewArtworkRepository(context.Background(), GetDBXBuilder()) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: albumSgtPeppers.ID, ImageType: model.ImageTypePrimary, Hash: "hX", AttemptedAt: time.Now()})).To(Succeed()) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: albumAbbeyRoad.ID, ImageType: model.ImageTypePrimary, Hash: "", AttemptedAt: time.Now()})).To(Succeed()) + awRepo := NewArtworkRepository(GetDBXBuilder()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: albumSgtPeppers.ID, ImageType: model.ImageTypePrimary, Hash: "hX", AttemptedAt: time.Now()})).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: albumAbbeyRoad.ID, ImageType: model.ImageTypePrimary, Hash: "", AttemptedAt: time.Now()})).To(Succeed()) - n, err := repo.EnqueueAllMissing(model.KindAlbumArtwork, model.ArtworkPriorityRecheck) + n, err := repo.EnqueueAllMissing(ctx, model.KindAlbumArtwork, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) Expect(n).To(BeNumerically(">=", 1)) - got, err := repo.DequeueBatch(1000) + got, err := repo.DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) ids := make([]string, 0, len(got)) for _, it := range got { @@ -260,122 +262,122 @@ var _ = Describe("ArtworkQueueRepository", func() { }) It("EnqueueIfMissing skips items that already have an item_artwork row", func() { - awRepo := NewArtworkRepository(context.Background(), GetDBXBuilder()) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "resolved", ImageType: model.ImageTypePrimary, Hash: "hX", AttemptedAt: time.Now()})).To(Succeed()) - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "absent", ImageType: model.ImageTypePrimary, Hash: "", AttemptedAt: time.Now()})).To(Succeed()) + awRepo := NewArtworkRepository(GetDBXBuilder()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "resolved", ImageType: model.ImageTypePrimary, Hash: "hX", AttemptedAt: time.Now()})).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "absent", ImageType: model.ImageTypePrimary, Hash: "", AttemptedAt: time.Now()})).To(Succeed()) - Expect(repo.EnqueueIfMissing( + Expect(repo.EnqueueIfMissing(ctx, item("al", "resolved", model.ArtworkPriorityScan), item("al", "absent", model.ArtworkPriorityScan), item("al", "brandnew", model.ArtworkPriorityScan), )).To(Succeed()) - got, err := repo.DequeueBatch(100) + got, err := repo.DequeueBatch(ctx, 100) Expect(err).ToNot(HaveOccurred()) ids := slice.Map(got, func(it model.ArtworkQueueItem) string { return it.ItemID }) Expect(ids).To(ConsistOf("brandnew"), "only an item with no state row may be enqueued") }) It("EnqueueIfMissing leaves an already-queued row untouched", func() { - Expect(repo.Enqueue(item("al", "queued", model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "queued", model.ArtworkPriorityBump))).To(Succeed()) - Expect(repo.EnqueueIfMissing(item("al", "queued", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.EnqueueIfMissing(ctx, item("al", "queued", model.ArtworkPriorityScan))).To(Succeed()) - got, _ := repo.DequeueBatch(100) + got, _ := repo.DequeueBatch(ctx, 100) Expect(got).To(HaveLen(1)) Expect(got[0].Priority).To(Equal(model.ArtworkPriorityBump), "the existing priority must survive") }) Describe("EnqueueBySource", func() { BeforeEach(func() { - artRepo := NewArtworkRepository(context.Background(), GetDBXBuilder()) + artRepo := NewArtworkRepository(GetDBXBuilder()) for _, ia := range []model.ItemArtwork{ {ItemKind: "ar", ItemID: "ar1", ImageType: model.ImageTypePrimary, Hash: "h1", Source: "external:deezer"}, {ItemKind: "ar", ItemID: "ar2", ImageType: model.ImageTypePrimary, Hash: "h2", Source: "external:lastfm"}, {ItemKind: "ar", ItemID: "ar3", ImageType: model.ImageTypePrimary, Hash: "", Source: ""}, {ItemKind: "al", ItemID: "al1", ImageType: model.ImageTypePrimary, Hash: "h4", Source: "external:deezer"}, } { - Expect(artRepo.PutItemArtwork(&ia)).To(Succeed()) + Expect(artRepo.PutItemArtwork(ctx, &ia)).To(Succeed()) } }) It("enqueues only the matching source within the kind", func() { - n, err := repo.EnqueueBySource(model.KindArtistArtwork, []string{"external:deezer"}, model.ArtworkPriorityRecheck) + n, err := repo.EnqueueBySource(ctx, model.KindArtistArtwork, []string{"external:deezer"}, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(int64(1)), "al1 is a different kind and must not be touched") - got, err := repo.DequeueBatch(10) + got, err := repo.DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(slice.Map(got, func(it model.ArtworkQueueItem) string { return it.ItemID })).To(ConsistOf("ar1")) }) It("treats the empty source as absent", func() { - n, err := repo.EnqueueBySource(model.KindArtistArtwork, []string{""}, model.ArtworkPriorityRecheck) + n, err := repo.EnqueueBySource(ctx, model.KindArtistArtwork, []string{""}, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(int64(1))) - got, _ := repo.DequeueBatch(10) + got, _ := repo.DequeueBatch(ctx, 10) Expect(slice.Map(got, func(it model.ArtworkQueueItem) string { return it.ItemID })).To(ConsistOf("ar3")) }) It("enqueues every source when none is given", func() { - n, err := repo.EnqueueBySource(model.KindArtistArtwork, nil, model.ArtworkPriorityRecheck) + n, err := repo.EnqueueBySource(ctx, model.KindArtistArtwork, nil, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(int64(3))) }) It("leaves the current artwork state in place", func() { - _, err := repo.EnqueueBySource(model.KindArtistArtwork, []string{"external:deezer"}, model.ArtworkPriorityRecheck) + _, err := repo.EnqueueBySource(ctx, model.KindArtistArtwork, []string{"external:deezer"}, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) - artRepo := NewArtworkRepository(context.Background(), GetDBXBuilder()) - ia, err := artRepo.GetItemArtwork(model.KindArtistArtwork, "ar1", model.ImageTypePrimary) + artRepo := NewArtworkRepository(GetDBXBuilder()) + ia, err := artRepo.GetItemArtwork(ctx, model.KindArtistArtwork, "ar1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Hash).To(Equal("h1"), "the current image must survive until it is replaced") Expect(ia.Source).To(Equal("external:deezer")) }) It("does not disturb an already-queued row", func() { - Expect(repo.Enqueue(item("ar", "ar1", model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("ar", "ar1", model.ArtworkPriorityBump))).To(Succeed()) - n, err := repo.EnqueueBySource(model.KindArtistArtwork, []string{"external:deezer"}, model.ArtworkPriorityRecheck) + n, err := repo.EnqueueBySource(ctx, model.KindArtistArtwork, []string{"external:deezer"}, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) Expect(n).To(BeZero()) - got, _ := repo.DequeueBatch(10) + got, _ := repo.DequeueBatch(ctx, 10) Expect(got).To(HaveLen(1)) Expect(got[0].Priority).To(Equal(model.ArtworkPriorityBump)) }) It("counts without enqueueing", func() { - n, err := repo.CountBySource(model.KindArtistArtwork, []string{"external:deezer"}) + n, err := repo.CountBySource(ctx, model.KindArtistArtwork, []string{"external:deezer"}) Expect(err).ToNot(HaveOccurred()) Expect(n).To(Equal(int64(1))) - queued, err := repo.Count() + queued, err := repo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(queued).To(BeZero(), "CountBySource must not enqueue") }) It("counts the absent source and every source", func() { - Expect(repo.CountBySource(model.KindArtistArtwork, []string{""})).To(Equal(int64(1))) - Expect(repo.CountBySource(model.KindArtistArtwork, nil)).To(Equal(int64(3))) + Expect(repo.CountBySource(ctx, model.KindArtistArtwork, []string{""})).To(Equal(int64(1))) + Expect(repo.CountBySource(ctx, model.KindArtistArtwork, nil)).To(Equal(int64(3))) }) It("lists the distinct sources in use by a kind", func() { - Expect(repo.SourcesInUse(model.KindArtistArtwork)).To(ConsistOf("", "external:deezer", "external:lastfm")) - Expect(repo.SourcesInUse(model.KindAlbumArtwork)).To(ConsistOf("external:deezer")) - Expect(repo.SourcesInUse(model.KindRadioArtwork)).To(BeEmpty()) + Expect(repo.SourcesInUse(ctx, model.KindArtistArtwork)).To(ConsistOf("", "external:deezer", "external:lastfm")) + Expect(repo.SourcesInUse(ctx, model.KindAlbumArtwork)).To(ConsistOf("external:deezer")) + Expect(repo.SourcesInUse(ctx, model.KindRadioArtwork)).To(BeEmpty()) }) }) It("does not disturb an already-queued entity when enqueueing missing rows", func() { - Expect(repo.Enqueue(item("al", albumRadioactivity.ID, model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", albumRadioactivity.ID, model.ArtworkPriorityBump))).To(Succeed()) - _, err := repo.EnqueueAllMissing(model.KindAlbumArtwork, model.ArtworkPriorityRecheck) + _, err := repo.EnqueueAllMissing(ctx, model.KindAlbumArtwork, model.ArtworkPriorityRecheck) Expect(err).ToNot(HaveOccurred()) - got, _ := repo.DequeueBatch(1000) + got, _ := repo.DequeueBatch(ctx, 1000) var count int for _, it := range got { if it.ItemID == albumRadioactivity.ID { @@ -388,12 +390,12 @@ var _ = Describe("ArtworkQueueRepository", func() { Describe("status counters", func() { It("groups queue rows by kind and priority", func() { - Expect(repo.Enqueue(item("ar", "a1", model.ArtworkPriorityBackfill))).To(Succeed()) - Expect(repo.Enqueue(item("ar", "a2", model.ArtworkPriorityBackfill))).To(Succeed()) - Expect(repo.Enqueue(item("ar", "a3", model.ArtworkPriorityBump))).To(Succeed()) - Expect(repo.Enqueue(item("al", "b1", model.ArtworkPriorityScan))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("ar", "a1", model.ArtworkPriorityBackfill))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("ar", "a2", model.ArtworkPriorityBackfill))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("ar", "a3", model.ArtworkPriorityBump))).To(Succeed()) + Expect(repo.Enqueue(ctx, item("al", "b1", model.ArtworkPriorityScan))).To(Succeed()) - Expect(repo.CountQueued(nil, nil)).To(ConsistOf( + Expect(repo.CountQueued(ctx, nil, nil)).To(ConsistOf( model.ArtworkQueueStat{ItemKind: "ar", Priority: model.ArtworkPriorityBackfill, Count: 2}, model.ArtworkQueueStat{ItemKind: "ar", Priority: model.ArtworkPriorityBump, Count: 1}, model.ArtworkQueueStat{ItemKind: "al", Priority: model.ArtworkPriorityScan, Count: 1}, @@ -401,46 +403,46 @@ var _ = Describe("ArtworkQueueRepository", func() { }) It("selects only the absent states that gave up, not those a source answered", func() { - awRepo := NewArtworkRepository(context.Background(), GetDBXBuilder()) + awRepo := NewArtworkRepository(GetDBXBuilder()) for _, ia := range []model.ItemArtwork{ {ItemKind: "ar", ItemID: "gaveup", ImageType: model.ImageTypePrimary, LastFailure: "[]"}, {ItemKind: "ar", ItemID: "toldno", ImageType: model.ImageTypePrimary}, {ItemKind: "ar", ItemID: "hasart", ImageType: model.ImageTypePrimary, Hash: "hX", LastFailure: "[]"}, } { - Expect(awRepo.PutItemArtwork(&ia)).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &ia)).To(Succeed()) } - Expect(repo.CountBySource(model.KindArtistArtwork, []string{model.ArtworkSourceFailed})).To(Equal(int64(1)), + Expect(repo.CountBySource(ctx, model.KindArtistArtwork, []string{model.ArtworkSourceFailed})).To(Equal(int64(1)), "an item still serving art is not absent, however its last attempt went") // A later success rewrites the row, clearing the record. - Expect(awRepo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "gaveup", + Expect(awRepo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "gaveup", ImageType: model.ImageTypePrimary, Hash: "hZ"})).To(Succeed()) - Expect(repo.CountBySource(model.KindArtistArtwork, []string{model.ArtworkSourceFailed})).To(Equal(int64(0))) + Expect(repo.CountBySource(ctx, model.KindArtistArtwork, []string{model.ArtworkSourceFailed})).To(Equal(int64(0))) }) It("unions the failed pseudo-source with a real one, so absent plus failed is just absent", func() { - awRepo := NewArtworkRepository(context.Background(), GetDBXBuilder()) + awRepo := NewArtworkRepository(GetDBXBuilder()) for _, ia := range []model.ItemArtwork{ {ItemKind: "ar", ItemID: "gaveup", ImageType: model.ImageTypePrimary, LastFailure: "[]"}, {ItemKind: "ar", ItemID: "toldno", ImageType: model.ImageTypePrimary}, {ItemKind: "ar", ItemID: "folder", ImageType: model.ImageTypePrimary, Hash: "hX", Source: "folder"}, } { - Expect(awRepo.PutItemArtwork(&ia)).To(Succeed()) + Expect(awRepo.PutItemArtwork(ctx, &ia)).To(Succeed()) } failedAndAbsent := []string{model.ArtworkSourceFailed, ""} - Expect(repo.CountBySource(model.KindArtistArtwork, failedAndAbsent)).To(Equal(int64(2)), + Expect(repo.CountBySource(ctx, model.KindArtistArtwork, failedAndAbsent)).To(Equal(int64(2)), "failed is a subset of absent, so asking for both is asking for absent") - Expect(repo.CountBySource(model.KindArtistArtwork, []string{model.ArtworkSourceFailed, "folder"})). + Expect(repo.CountBySource(ctx, model.KindArtistArtwork, []string{model.ArtworkSourceFailed, "folder"})). To(Equal(int64(2)), "a pseudo-source and a stored source combine as a union") }) It("reports a kind with nothing failed as zero", func() { - Expect(repo.CountBySource(model.KindRadioArtwork, []string{model.ArtworkSourceFailed})).To(Equal(int64(0))) + Expect(repo.CountBySource(ctx, model.KindRadioArtwork, []string{model.ArtworkSourceFailed})).To(Equal(int64(0))) }) It("reports an empty queue as no rows", func() { - Expect(repo.CountQueued(nil, nil)).To(BeEmpty()) + Expect(repo.CountQueued(ctx, nil, nil)).To(BeEmpty()) }) }) @@ -448,13 +450,13 @@ var _ = Describe("ArtworkQueueRepository", func() { Describe("PurgeQueued", func() { queuedIDs := func() []string { GinkgoHelper() - got, err := repo.DequeueBatch(100) + got, err := repo.DequeueBatch(ctx, 100) Expect(err).ToNot(HaveOccurred()) return slice.Map(got, func(it model.ArtworkQueueItem) string { return it.ItemID }) } BeforeEach(func() { - Expect(repo.Enqueue( + Expect(repo.Enqueue(ctx, item("ar", "ar-backfill", model.ArtworkPriorityBackfill), item("ar", "ar-bump", model.ArtworkPriorityBump), item("al", "al-backfill", model.ArtworkPriorityBackfill), @@ -466,7 +468,7 @@ var _ = Describe("ArtworkQueueRepository", func() { // every selection must count exactly what it deletes. DescribeTable("selects the same rows to count and to delete", func(kinds []model.Kind, priorities []int, deleted int, remaining []string) { - counted, err := repo.CountQueued(kinds, priorities) + counted, err := repo.CountQueued(ctx, kinds, priorities) Expect(err).ToNot(HaveOccurred()) var total int64 for _, s := range counted { @@ -474,7 +476,7 @@ var _ = Describe("ArtworkQueueRepository", func() { } Expect(total).To(BeNumerically("==", deleted), "the preview must match the delete") - Expect(repo.PurgeQueued(kinds, priorities)).To(BeNumerically("==", deleted)) + Expect(repo.PurgeQueued(ctx, kinds, priorities)).To(BeNumerically("==", deleted)) Expect(queuedIDs()).To(ConsistOf(remaining)) }, Entry("only the given kinds", []model.Kind{model.KindArtistArtwork}, nil, @@ -496,8 +498,8 @@ var _ = Describe("ArtworkQueueRepository", func() { It("deletes a row that is still backing off", func() { backOff("ar", "ar-bump", time.Now().Add(time.Hour)) - Expect(repo.PurgeQueued([]model.Kind{model.KindArtistArtwork}, nil)).To(BeNumerically("==", 2)) - Expect(repo.Get(model.KindArtistArtwork, "ar-bump", model.ImageTypePrimary)). + Expect(repo.PurgeQueued(ctx, []model.Kind{model.KindArtistArtwork}, nil)).To(BeNumerically("==", 2)) + Expect(repo.Get(ctx, model.KindArtistArtwork, "ar-bump", model.ImageTypePrimary)). Error().To(MatchError(model.ErrNotFound)) }) diff --git a/persistence/artwork_repository.go b/persistence/artwork_repository.go index 89eb1d415..bd3a3d87c 100644 --- a/persistence/artwork_repository.go +++ b/persistence/artwork_repository.go @@ -21,27 +21,25 @@ type artworkRepository struct { items sqlRepository } -func NewArtworkRepository(ctx context.Context, db dbx.Builder) model.ArtworkRepository { +func NewArtworkRepository(db dbx.Builder) model.ArtworkRepository { r := &artworkRepository{} - r.ctx = ctx r.db = db r.tableName = "artwork" - r.items.ctx = ctx r.items.db = db r.items.tableName = itemArtworkTable return r } -func (r *artworkRepository) GetImage(hash string) (*model.Artwork, error) { +func (r *artworkRepository) GetImage(ctx context.Context, hash string) (*model.Artwork, error) { sel := Select("*").From(r.tableName).Where(Eq{"hash": hash}) var res model.Artwork - if err := r.queryOne(sel, &res); err != nil { + if err := r.queryOne(ctx, sel, &res); err != nil { return nil, err } return &res, nil } -func (r *artworkRepository) PutImage(a *model.Artwork) error { +func (r *artworkRepository) PutImage(ctx context.Context, a *model.Artwork) error { // created_at is the last-acquisition-write time the prune grace window keys on. a.CreatedAt = time.Now() values, err := toSQLArgs(*a) @@ -53,17 +51,17 @@ func (r *artworkRepository) PutImage(a *model.Artwork) error { height=excluded.height, size_bytes=excluded.size_bytes, blur_hash=excluded.blur_hash, thumb_hash=excluded.thumb_hash, dominant_color=excluded.dominant_color, created_at=excluded.created_at`) - _, err = r.executeSQL(ins) + _, err = r.executeSQL(ctx, ins) return err } -func (r *artworkRepository) GetMimeByHash() (map[string]string, error) { +func (r *artworkRepository) GetMimeByHash(ctx context.Context) (map[string]string, error) { sel := Select("hash", "mime").From(r.tableName) var rows []struct { Hash string Mime string } - if err := r.queryAll(sel, &rows); err != nil { + if err := r.queryAll(ctx, sel, &rows); err != nil { return nil, err } res := make(map[string]string, len(rows)) @@ -73,12 +71,12 @@ func (r *artworkRepository) GetMimeByHash() (map[string]string, error) { return res, nil } -func (r *artworkRepository) PurgeOrphans(createdBefore time.Time) (int64, error) { +func (r *artworkRepository) PurgeOrphans(ctx context.Context, createdBefore time.Time) (int64, error) { del := Delete(r.tableName).Where(And{ Lt{"created_at": createdBefore}, Expr("hash NOT IN (SELECT hash FROM " + itemArtworkTable + " WHERE hash <> '')"), }) - return r.executeSQL(del) + return r.executeSQL(ctx, del) } // artworkOwnerTables maps an artwork kind to the table that owns the entity. @@ -91,14 +89,14 @@ var artworkOwnerTables = map[model.Kind]string{ } // purgeDangling deletes rows in r's table whose owning entity is gone, one statement per kind. -func purgeDangling(r sqlRepository) (int64, error) { +func purgeDangling(ctx context.Context, r sqlRepository) (int64, error) { var total int64 for kind, entityTable := range artworkOwnerTables { del := Delete(r.tableName).Where(And{ Eq{"item_kind": kind.Prefix()}, Expr("item_id NOT IN (SELECT id FROM " + entityTable + ")"), }) - c, err := r.executeSQL(del) + c, err := r.executeSQL(ctx, del) if err != nil { return total, err } @@ -107,21 +105,21 @@ func purgeDangling(r sqlRepository) (int64, error) { return total, nil } -func (r *artworkRepository) PurgeDanglingItems() (int64, error) { - return purgeDangling(r.items) +func (r *artworkRepository) PurgeDanglingItems(ctx context.Context) (int64, error) { + return purgeDangling(ctx, r.items) } -func (r *artworkRepository) GetItemArtwork(kind model.Kind, id, imageType string) (*model.ItemArtwork, error) { +func (r *artworkRepository) GetItemArtwork(ctx context.Context, kind model.Kind, id, imageType string) (*model.ItemArtwork, error) { sel := Select("*").From(itemArtworkTable). Where(Eq{"item_kind": kind.Prefix(), "item_id": id, "image_type": imageType}) var res model.ItemArtwork - if err := r.items.queryOne(sel, &res); err != nil { + if err := r.items.queryOne(ctx, sel, &res); err != nil { return nil, err } return &res, nil } -func (r *artworkRepository) PutItemArtwork(ia *model.ItemArtwork) error { +func (r *artworkRepository) PutItemArtwork(ctx context.Context, ia *model.ItemArtwork) error { ia.ImageType = cmp.Or(ia.ImageType, model.ImageTypePrimary) ia.UpdatedAt = time.Now() // PutItemArtwork records the outcome of an attempt, so an unset attempted_at is now. @@ -136,29 +134,29 @@ func (r *artworkRepository) PutItemArtwork(ia *model.ItemArtwork) error { hash=excluded.hash, source=excluded.source, source_path=excluded.source_path, ref_mtime=excluded.ref_mtime, trace=excluded.trace, last_failure=excluded.last_failure, attempted_at=excluded.attempted_at, updated_at=excluded.updated_at`) - _, err = r.items.executeSQL(ins) + _, err = r.items.executeSQL(ctx, ins) return err } // PutLastFailure records why an item exhausted its retry budget. It only updates an existing row: // inserting one would write an empty hash, which the rest of the system reads as a settled absent. -func (r *artworkRepository) PutLastFailure(kind model.Kind, id, imageType, trace string) error { +func (r *artworkRepository) PutLastFailure(ctx context.Context, kind model.Kind, id, imageType, trace string) error { upd := Update(itemArtworkTable).Set("last_failure", trace). Where(Eq{"item_kind": kind.Prefix(), "item_id": id, "image_type": imageType}) - _, err := r.items.executeSQL(upd) + _, err := r.items.executeSQL(ctx, upd) return err } -func (r *artworkRepository) DeleteForItems(kind model.Kind, ids []string) error { +func (r *artworkRepository) DeleteForItems(ctx context.Context, kind model.Kind, ids []string) error { for chunk := range slices.Chunk(ids, artworkBatchSize) { - if err := r.items.delete(Eq{"item_kind": kind.Prefix(), "item_id": chunk}); err != nil { + if err := r.items.delete(ctx, Eq{"item_kind": kind.Prefix(), "item_id": chunk}); err != nil { return err } } return nil } -func (r *artworkRepository) GetInfoForItems(kind model.Kind, ids []string) (map[string]model.ItemArtworkInfo, error) { +func (r *artworkRepository) GetInfoForItems(ctx context.Context, kind model.Kind, ids []string) (map[string]model.ItemArtworkInfo, error) { res := map[string]model.ItemArtworkInfo{} for chunk := range slices.Chunk(ids, artworkBatchSize) { sel := Select("ia.item_id", "ia.hash", "COALESCE(a.blur_hash, '') as blur_hash", @@ -173,7 +171,7 @@ func (r *artworkRepository) GetInfoForItems(kind model.Kind, ids []string) (map[ Eq{"ia.item_id": chunk}, }) var rows []model.ItemArtworkInfo - if err := r.items.queryAll(sel, &rows); err != nil { + if err := r.items.queryAll(ctx, sel, &rows); err != nil { return nil, err } for _, row := range rows { diff --git a/persistence/artwork_repository_test.go b/persistence/artwork_repository_test.go index a9687f76b..338a7c3ab 100644 --- a/persistence/artwork_repository_test.go +++ b/persistence/artwork_repository_test.go @@ -20,55 +20,57 @@ func clearArtworkTables() { } var _ = Describe("ArtworkRepository", func() { + var ctx context.Context var repo model.ArtworkRepository BeforeEach(func() { + ctx = GinkgoT().Context() clearArtworkTables() DeferCleanup(clearArtworkTables) - repo = NewArtworkRepository(context.Background(), GetDBXBuilder()) + repo = NewArtworkRepository(GetDBXBuilder()) }) Context("resolution traces", func() { const traceJSON = `[{"c":"cover.*","o":"hit"}]` It("round-trips the trace with the state row", func() { - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "t1", + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "t1", ImageType: model.ImageTypePrimary, Hash: "h1", Trace: traceJSON})).To(Succeed()) - got, err := repo.GetItemArtwork(model.KindAlbumArtwork, "t1", model.ImageTypePrimary) + got, err := repo.GetItemArtwork(ctx, model.KindAlbumArtwork, "t1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(got.Trace).To(Equal(traceJSON)) Expect(got.LastFailure).To(BeEmpty()) }) It("replaces the trace when the item is resolved again", func() { - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "t2", + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "t2", ImageType: model.ImageTypePrimary, Trace: traceJSON})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "t2", + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "t2", ImageType: model.ImageTypePrimary, Trace: `[{"c":"embedded","o":"hit"}]`})).To(Succeed()) - got, _ := repo.GetItemArtwork(model.KindAlbumArtwork, "t2", model.ImageTypePrimary) + got, _ := repo.GetItemArtwork(ctx, model.KindAlbumArtwork, "t2", model.ImageTypePrimary) Expect(got.Trace).To(Equal(`[{"c":"embedded","o":"hit"}]`)) }) It("records a last failure on an existing row", func() { - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "t3", + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "t3", ImageType: model.ImageTypePrimary, Hash: "h3"})).To(Succeed()) - Expect(repo.PutLastFailure(model.KindAlbumArtwork, "t3", model.ImageTypePrimary, + Expect(repo.PutLastFailure(ctx, model.KindAlbumArtwork, "t3", model.ImageTypePrimary, `[{"c":"decode","o":"error"}]`)).To(Succeed()) - got, _ := repo.GetItemArtwork(model.KindAlbumArtwork, "t3", model.ImageTypePrimary) + got, _ := repo.GetItemArtwork(ctx, model.KindAlbumArtwork, "t3", model.ImageTypePrimary) Expect(got.LastFailure).To(Equal(`[{"c":"decode","o":"error"}]`)) Expect(got.Hash).To(Equal("h3"), "recording a failure must not disturb the served artwork") }) // Inserting here would write hash='', which every reader treats as a settled absent. It("never creates a row for an item that has no state", func() { - Expect(repo.PutLastFailure(model.KindAlbumArtwork, "ghost", model.ImageTypePrimary, + Expect(repo.PutLastFailure(ctx, model.KindAlbumArtwork, "ghost", model.ImageTypePrimary, `[{"c":"decode","o":"error"}]`)).To(Succeed()) - _, err := repo.GetItemArtwork(model.KindAlbumArtwork, "ghost", model.ImageTypePrimary) + _, err := repo.GetItemArtwork(ctx, model.KindAlbumArtwork, "ghost", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) }) @@ -76,9 +78,9 @@ var _ = Describe("ArtworkRepository", func() { Context("image identity", func() { It("stores and retrieves an artwork by hash", func() { a := &model.Artwork{Hash: "abc123", Mime: "image/jpeg", Width: 500, Height: 500, SizeBytes: 1234, BlurHash: "LKO2?U%2Tw=w"} - Expect(repo.PutImage(a)).To(Succeed()) + Expect(repo.PutImage(ctx, a)).To(Succeed()) - got, err := repo.GetImage("abc123") + got, err := repo.GetImage(ctx, "abc123") Expect(err).ToNot(HaveOccurred()) Expect(got.Mime).To(Equal("image/jpeg")) Expect(got.BlurHash).To(Equal("LKO2?U%2Tw=w")) @@ -88,9 +90,9 @@ var _ = Describe("ArtworkRepository", func() { It("round-trips the thumbhash alongside the blurhash", func() { a := &model.Artwork{Hash: "both1", Mime: "image/jpeg", BlurHash: "LKO2?U%2Tw=w", ThumbHash: "1QcSHQRnh493V4dIh4eXh1h4kJUI", DominantColor: "#336699"} - Expect(repo.PutImage(a)).To(Succeed()) + Expect(repo.PutImage(ctx, a)).To(Succeed()) - got, err := repo.GetImage("both1") + got, err := repo.GetImage(ctx, "both1") Expect(err).ToNot(HaveOccurred()) Expect(got.BlurHash).To(Equal("LKO2?U%2Tw=w")) Expect(got.ThumbHash).To(Equal("1QcSHQRnh493V4dIh4eXh1h4kJUI")) @@ -98,10 +100,10 @@ var _ = Describe("ArtworkRepository", func() { }) It("overwrites the thumbhash on re-acquisition", func() { - Expect(repo.PutImage(&model.Artwork{Hash: "th2", Mime: "image/png", ThumbHash: "first", DominantColor: "#111111"})).To(Succeed()) - Expect(repo.PutImage(&model.Artwork{Hash: "th2", Mime: "image/png", ThumbHash: "second", DominantColor: "#222222"})).To(Succeed()) + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "th2", Mime: "image/png", ThumbHash: "first", DominantColor: "#111111"})).To(Succeed()) + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "th2", Mime: "image/png", ThumbHash: "second", DominantColor: "#222222"})).To(Succeed()) - got, err := repo.GetImage("th2") + got, err := repo.GetImage(ctx, "th2") Expect(err).ToNot(HaveOccurred()) Expect(got.ThumbHash).To(Equal("second")) Expect(got.DominantColor).To(Equal("#222222")) @@ -109,76 +111,76 @@ var _ = Describe("ArtworkRepository", func() { It("is idempotent on Put (upsert by hash)", func() { a := &model.Artwork{Hash: "dup1", Mime: "image/png"} - Expect(repo.PutImage(a)).To(Succeed()) + Expect(repo.PutImage(ctx, a)).To(Succeed()) a.BlurHash = "XYZ" - Expect(repo.PutImage(a)).To(Succeed()) - got, _ := repo.GetImage("dup1") + Expect(repo.PutImage(ctx, a)).To(Succeed()) + got, _ := repo.GetImage(ctx, "dup1") Expect(got.BlurHash).To(Equal("XYZ")) }) It("refreshes created_at when reacquiring an existing hash", func() { - Expect(repo.PutImage(&model.Artwork{Hash: "reacq", Mime: "image/jpeg"})).To(Succeed()) + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "reacq", Mime: "image/jpeg"})).To(Succeed()) _, err := GetDBXBuilder().NewQuery("UPDATE artwork SET created_at={:t} WHERE hash='reacq'"). Bind(dbx.Params{"t": "2000-01-01 00:00:00"}).Execute() Expect(err).ToNot(HaveOccurred()) - Expect(repo.PutImage(&model.Artwork{Hash: "reacq", Mime: "image/png"})).To(Succeed()) + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "reacq", Mime: "image/png"})).To(Succeed()) - got, err := repo.GetImage("reacq") + got, err := repo.GetImage(ctx, "reacq") Expect(err).ToNot(HaveOccurred()) Expect(got.CreatedAt).To(BeTemporally(">", time.Date(2020, 1, 1, 0, 0, 0, 0, time.UTC))) }) It("returns ErrNotFound for a missing hash", func() { - _, err := repo.GetImage("nope") + _, err := repo.GetImage(ctx, "nope") Expect(err).To(MatchError(model.ErrNotFound)) }) It("returns every stored hash with its current mime", func() { - Expect(repo.PutImage(&model.Artwork{Hash: "all1", Mime: "image/jpeg"})).To(Succeed()) - Expect(repo.PutImage(&model.Artwork{Hash: "all2", Mime: "image/png"})).To(Succeed()) - mimes, err := repo.GetMimeByHash() + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "all1", Mime: "image/jpeg"})).To(Succeed()) + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "all2", Mime: "image/png"})).To(Succeed()) + mimes, err := repo.GetMimeByHash(ctx) Expect(err).ToNot(HaveOccurred()) Expect(mimes).To(HaveKeyWithValue("all1", "image/jpeg")) Expect(mimes).To(HaveKeyWithValue("all2", "image/png")) }) It("deletes only unreferenced rows older than the cutoff, reporting the count", func() { - Expect(repo.PutImage(&model.Artwork{Hash: "d1", Mime: "image/jpeg"})).To(Succeed()) - Expect(repo.PutImage(&model.Artwork{Hash: "dref", Mime: "image/jpeg"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "a1", + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "d1", Mime: "image/jpeg"})).To(Succeed()) + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "dref", Mime: "image/jpeg"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "a1", ImageType: model.ImageTypePrimary, Hash: "dref", Source: "folder"})).To(Succeed()) - Expect(repo.PurgeOrphans(time.Now().Add(time.Minute))).To(BeNumerically("==", 1)) + Expect(repo.PurgeOrphans(ctx, time.Now().Add(time.Minute))).To(BeNumerically("==", 1)) - _, err := repo.GetImage("d1") + _, err := repo.GetImage(ctx, "d1") Expect(err).To(MatchError(model.ErrNotFound)) - _, err = repo.GetImage("dref") + _, err = repo.GetImage(ctx, "dref") Expect(err).ToNot(HaveOccurred()) }) It("spares an unreferenced row younger than the cutoff", func() { - Expect(repo.PutImage(&model.Artwork{Hash: "young", Mime: "image/jpeg"})).To(Succeed()) - Expect(repo.PurgeOrphans(time.Now().Add(-time.Hour))).To(BeNumerically("==", 0)) - _, err := repo.GetImage("young") + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "young", Mime: "image/jpeg"})).To(Succeed()) + Expect(repo.PurgeOrphans(ctx, time.Now().Add(-time.Hour))).To(BeNumerically("==", 0)) + _, err := repo.GetImage(ctx, "young") Expect(err).ToNot(HaveOccurred()) }) }) Context("dangling state cleanup", func() { It("purges item_artwork rows per kind whose entity no longer exists, summing counts", func() { - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: albumSgtPeppers.ID, ImageType: model.ImageTypePrimary, Hash: "keepAl"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "no-such-album", ImageType: model.ImageTypePrimary, Hash: "danglingAl"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: artistKraftwerk.ID, ImageType: model.ImageTypePrimary, Hash: "keepAr"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "no-such-artist", ImageType: model.ImageTypePrimary, Hash: "danglingAr"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "pl", ItemID: plsBest.ID, ImageType: model.ImageTypePrimary, Hash: "keepPl"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "pl", ItemID: "no-such-playlist", ImageType: model.ImageTypePrimary, Hash: "danglingPl"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ra", ItemID: radioWithHomePage.ID, ImageType: model.ImageTypePrimary, Hash: "keepRa"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ra", ItemID: "no-such-radio", ImageType: model.ImageTypePrimary, Hash: "danglingRa"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "mf", ItemID: songDayInALife.ID, ImageType: model.ImageTypePrimary, Hash: "keepMf"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "mf", ItemID: "no-such-mediafile", ImageType: model.ImageTypePrimary, Hash: "danglingMf"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: albumSgtPeppers.ID, ImageType: model.ImageTypePrimary, Hash: "keepAl"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "no-such-album", ImageType: model.ImageTypePrimary, Hash: "danglingAl"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: artistKraftwerk.ID, ImageType: model.ImageTypePrimary, Hash: "keepAr"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "no-such-artist", ImageType: model.ImageTypePrimary, Hash: "danglingAr"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "pl", ItemID: plsBest.ID, ImageType: model.ImageTypePrimary, Hash: "keepPl"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "pl", ItemID: "no-such-playlist", ImageType: model.ImageTypePrimary, Hash: "danglingPl"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ra", ItemID: radioWithHomePage.ID, ImageType: model.ImageTypePrimary, Hash: "keepRa"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ra", ItemID: "no-such-radio", ImageType: model.ImageTypePrimary, Hash: "danglingRa"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "mf", ItemID: songDayInALife.ID, ImageType: model.ImageTypePrimary, Hash: "keepMf"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "mf", ItemID: "no-such-mediafile", ImageType: model.ImageTypePrimary, Hash: "danglingMf"})).To(Succeed()) - purged, err := repo.PurgeDanglingItems() + purged, err := repo.PurgeDanglingItems(ctx) Expect(err).ToNot(HaveOccurred()) Expect(purged).To(Equal(int64(5))) @@ -190,7 +192,7 @@ var _ = Describe("ArtworkRepository", func() { {ItemKind: "mf", ItemID: songDayInALife.ID}, } { k, _ := model.ParseKind(kept.ItemKind) - _, err := repo.GetItemArtwork(k, kept.ItemID, model.ImageTypePrimary) + _, err := repo.GetItemArtwork(ctx, k, kept.ItemID, model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) } for _, gone := range []model.ItemArtwork{ @@ -201,7 +203,7 @@ var _ = Describe("ArtworkRepository", func() { {ItemKind: "mf", ItemID: "no-such-mediafile"}, } { k, _ := model.ParseKind(gone.ItemKind) - _, err := repo.GetItemArtwork(k, gone.ItemID, model.ImageTypePrimary) + _, err := repo.GetItemArtwork(ctx, k, gone.ItemID, model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) } }) @@ -211,13 +213,13 @@ var _ = Describe("ArtworkRepository", func() { It("upserts and reads state, including per-item provenance", func() { ia := &model.ItemArtwork{ItemKind: "al", ItemID: "al1", ImageType: model.ImageTypePrimary, Hash: "h1", Source: "folder", SourcePath: "/music/a/cover.jpg", RefMtime: 111, AttemptedAt: time.Now()} - Expect(repo.PutItemArtwork(ia)).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, ia)).To(Succeed()) ia.Source = "embedded" ia.SourcePath = "/music/a/track.mp3" ia.RefMtime = 222 - Expect(repo.PutItemArtwork(ia)).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, ia)).To(Succeed()) - got, err := repo.GetItemArtwork(model.KindAlbumArtwork, "al1", model.ImageTypePrimary) + got, err := repo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(got.Source).To(Equal("embedded")) Expect(got.SourcePath).To(Equal("/music/a/track.mp3")) @@ -227,28 +229,28 @@ var _ = Describe("ArtworkRepository", func() { It("defaults attempted_at to now when unset", func() { before := time.Now().Add(-time.Second) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "noattempt", + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "noattempt", ImageType: model.ImageTypePrimary, Hash: ""})).To(Succeed()) - got, err := repo.GetItemArtwork(model.KindArtistArtwork, "noattempt", model.ImageTypePrimary) + got, err := repo.GetItemArtwork(ctx, model.KindArtistArtwork, "noattempt", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(got.AttemptedAt).To(BeTemporally(">", before)) }) It("represents known-absent as empty hash", func() { - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "ar", ItemID: "ar1", + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "ar", ItemID: "ar1", ImageType: model.ImageTypePrimary, Hash: "", AttemptedAt: time.Now()})).To(Succeed()) - got, err := repo.GetItemArtwork(model.KindArtistArtwork, "ar1", model.ImageTypePrimary) + got, err := repo.GetItemArtwork(ctx, model.KindArtistArtwork, "ar1", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(got.Hash).To(BeEmpty()) }) It("hydrates a page in one batch, including blurhash, dimensions and absence", func() { - Expect(repo.PutImage(&model.Artwork{Hash: "h9", Mime: "image/jpeg", BlurHash: "BH9", + Expect(repo.PutImage(ctx, &model.Artwork{Hash: "h9", Mime: "image/jpeg", BlurHash: "BH9", DominantColor: "#abcdef", Width: 1200, Height: 800})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "x1", ImageType: model.ImageTypePrimary, Hash: "h9", Source: "folder"})).To(Succeed()) - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "al", ItemID: "x2", ImageType: model.ImageTypePrimary, Hash: "", Source: ""})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "x1", ImageType: model.ImageTypePrimary, Hash: "h9", Source: "folder"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "al", ItemID: "x2", ImageType: model.ImageTypePrimary, Hash: "", Source: ""})).To(Succeed()) - info, err := repo.GetInfoForItems(model.KindAlbumArtwork, []string{"x1", "x2", "x3"}) + info, err := repo.GetInfoForItems(ctx, model.KindAlbumArtwork, []string{"x1", "x2", "x3"}) Expect(err).ToNot(HaveOccurred()) Expect(info).To(HaveLen(2)) Expect(info["x1"].Hash).To(Equal("h9")) @@ -263,9 +265,9 @@ var _ = Describe("ArtworkRepository", func() { }) It("deletes all rows for a single item", func() { - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "pl", ItemID: "p1", ImageType: model.ImageTypePrimary, Hash: "h1"})).To(Succeed()) - Expect(repo.DeleteForItems(model.KindPlaylistArtwork, []string{"p1"})).To(Succeed()) - _, err := repo.GetItemArtwork(model.KindPlaylistArtwork, "p1", model.ImageTypePrimary) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "pl", ItemID: "p1", ImageType: model.ImageTypePrimary, Hash: "h1"})).To(Succeed()) + Expect(repo.DeleteForItems(ctx, model.KindPlaylistArtwork, []string{"p1"})).To(Succeed()) + _, err := repo.GetItemArtwork(ctx, model.KindPlaylistArtwork, "p1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -275,17 +277,17 @@ var _ = Describe("ArtworkRepository", func() { for i := range n { id := fmt.Sprintf("mf-%d", i) ids[i] = id - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "mf", ItemID: id, ImageType: model.ImageTypePrimary, Hash: "h1"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "mf", ItemID: id, ImageType: model.ImageTypePrimary, Hash: "h1"})).To(Succeed()) } - Expect(repo.PutItemArtwork(&model.ItemArtwork{ItemKind: "mf", ItemID: "keep", ImageType: model.ImageTypePrimary, Hash: "h1"})).To(Succeed()) + Expect(repo.PutItemArtwork(ctx, &model.ItemArtwork{ItemKind: "mf", ItemID: "keep", ImageType: model.ImageTypePrimary, Hash: "h1"})).To(Succeed()) - Expect(repo.DeleteForItems(model.KindMediaFileArtwork, ids)).To(Succeed()) + Expect(repo.DeleteForItems(ctx, model.KindMediaFileArtwork, ids)).To(Succeed()) for _, id := range ids { - _, err := repo.GetItemArtwork(model.KindMediaFileArtwork, id, model.ImageTypePrimary) + _, err := repo.GetItemArtwork(ctx, model.KindMediaFileArtwork, id, model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) } - kept, err := repo.GetItemArtwork(model.KindMediaFileArtwork, "keep", model.ImageTypePrimary) + kept, err := repo.GetItemArtwork(ctx, model.KindMediaFileArtwork, "keep", model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(kept.ItemID).To(Equal("keep")) }) diff --git a/persistence/criteria_sql_benchmark_test.go b/persistence/criteria_sql_benchmark_test.go index 523a5d825..30c1ba100 100644 --- a/persistence/criteria_sql_benchmark_test.go +++ b/persistence/criteria_sql_benchmark_test.go @@ -189,11 +189,11 @@ func setupBenchData(b *testing.B, ctx context.Context, conn *dbx.DB, user model. sqlDB := db.Db() - ur := NewUserRepository(ctx, conn) - if err := ur.Put(&user); err != nil { + ur := NewUserRepository(conn) + if err := ur.Put(ctx, &user); err != nil { b.Fatal(err) } - if err := ur.SetUserLibraries(user.ID, []int{1}); err != nil { + if err := ur.SetUserLibraries(ctx, user.ID, []int{1}); err != nil { b.Fatal(err) } diff --git a/persistence/e2e/e2e_suite_test.go b/persistence/e2e/e2e_suite_test.go index 2b617f5b0..a757901d3 100644 --- a/persistence/e2e/e2e_suite_test.go +++ b/persistence/e2e/e2e_suite_test.go @@ -143,7 +143,7 @@ func buildTestFS() { } func findMediaFileByTitle(title string) string { - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"media_file.title": title}, }) Expect(err).ToNot(HaveOccurred()) @@ -178,10 +178,10 @@ func evaluateRuleOrderedAs(owner model.User, jsonRule string) []string { OwnerID: owner.ID, Rules: &rules, } - err = ds.Playlist(userCtx).Put(pls) + err = ds.Playlist().Put(userCtx, pls) Expect(err).ToNot(HaveOccurred()) - loaded, err := ds.Playlist(userCtx).GetWithTracks(pls.ID, true, false) + loaded, err := ds.Playlist().GetWithTracks(userCtx, pls.ID, true, false) Expect(err).ToNot(HaveOccurred()) titles := make([]string, len(loaded.Tracks)) @@ -201,7 +201,7 @@ func createPlaylist(owner model.User, public bool, titles ...string) string { mfID := findMediaFileByTitle(title) pls.AddMediaFilesByID([]string{mfID}) } - Expect(ds.Playlist(ctx).Put(pls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, pls)).To(Succeed()) return pls.ID } @@ -230,7 +230,7 @@ func createSmartPlaylist(owner model.User, public bool, jsonRule string) string Public: public, Rules: &rules, } - Expect(ds.Playlist(ctx).Put(pls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, pls)).To(Succeed()) return pls.ID } @@ -252,22 +252,22 @@ var _ = BeforeSuite(func() { userWithPass := adminUser userWithPass.NewPassword = "password" - Expect(initDS.User(ctx).Put(&userWithPass)).To(Succeed()) + Expect(initDS.User().Put(ctx, &userWithPass)).To(Succeed()) regularUserWithPass := regularUser regularUserWithPass.NewPassword = "password" - Expect(initDS.User(ctx).Put(®ularUserWithPass)).To(Succeed()) + Expect(initDS.User().Put(ctx, ®ularUserWithPass)).To(Succeed()) lib = model.Library{ID: 1, Name: "Music Library", Path: "fake:///music"} - Expect(initDS.Library(ctx).Put(&lib)).To(Succeed()) - Expect(initDS.User(ctx).SetUserLibraries(adminUser.ID, []int{lib.ID})).To(Succeed()) - Expect(initDS.User(ctx).SetUserLibraries(regularUser.ID, []int{lib.ID})).To(Succeed()) + Expect(initDS.Library().Put(ctx, &lib)).To(Succeed()) + Expect(initDS.User().SetUserLibraries(ctx, adminUser.ID, []int{lib.ID})).To(Succeed()) + Expect(initDS.User().SetUserLibraries(ctx, regularUser.ID, []int{lib.ID})).To(Succeed()) - loadedUser, err := initDS.User(ctx).FindByUsername(adminUser.UserName) + loadedUser, err := initDS.User().FindByUsername(ctx, adminUser.UserName) Expect(err).ToNot(HaveOccurred()) adminUser.Libraries = loadedUser.Libraries - loadedOther, err := initDS.User(ctx).FindByUsername(regularUser.UserName) + loadedOther, err := initDS.User().FindByUsername(ctx, regularUser.UserName) Expect(err).ToNot(HaveOccurred()) regularUser.Libraries = loadedOther.Libraries @@ -282,14 +282,14 @@ var _ = BeforeSuite(func() { ds = &tests.MockDataStore{RealDS: persistence.New(db.Db())} comeTogetherID := findMediaFileByTitle("Come Together") - Expect(ds.MediaFile(ctx).SetStar(true, comeTogetherID)).To(Succeed()) - Expect(ds.MediaFile(ctx).SetStar(true, findMediaFileByTitle("So What"))).To(Succeed()) - Expect(ds.MediaFile(ctx).SetRating(3, findMediaFileByTitle("Stairway To Heaven"))).To(Succeed()) - Expect(ds.MediaFile(ctx).SetRating(5, findMediaFileByTitle("Bohemian Rhapsody"))).To(Succeed()) + Expect(ds.MediaFile().SetStar(ctx, true, comeTogetherID)).To(Succeed()) + Expect(ds.MediaFile().SetStar(ctx, true, findMediaFileByTitle("So What"))).To(Succeed()) + Expect(ds.MediaFile().SetRating(ctx, 3, findMediaFileByTitle("Stairway To Heaven"))).To(Succeed()) + Expect(ds.MediaFile().SetRating(ctx, 5, findMediaFileByTitle("Bohemian Rhapsody"))).To(Succeed()) for range 10 { - Expect(ds.MediaFile(ctx).IncPlayCount(comeTogetherID, time.Now())).To(Succeed()) + Expect(ds.MediaFile().IncPlayCount(ctx, comeTogetherID, time.Now())).To(Succeed()) } - Expect(ds.MediaFile(ctx).IncPlayCount(findMediaFileByTitle("Black Dog"), time.Now())).To(Succeed()) + Expect(ds.MediaFile().IncPlayCount(ctx, findMediaFileByTitle("Black Dog"), time.Now())).To(Succeed()) rows, err := db.Db().Query("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' AND name NOT LIKE '%_fts' AND name NOT LIKE '%_fts_%'") Expect(err).ToNot(HaveOccurred()) diff --git a/persistence/e2e/smartplaylist_test.go b/persistence/e2e/smartplaylist_test.go index b96605d09..51442c822 100644 --- a/persistence/e2e/smartplaylist_test.go +++ b/persistence/e2e/smartplaylist_test.go @@ -317,12 +317,12 @@ var _ = Describe("Smart Playlists", func() { smartBID := createPrivateSmartPlaylist(adminUser, `{"all":[{"is":{"genre":"Jazz"}}]}`) smartAID := createPublicSmartPlaylist(regularUser, `{"all":[{"inPlaylist":{"id":"`+smartBID+`"}}]}`) - loadedA, err := ds.Playlist(ctx).GetWithTracks(smartAID, true, false) + loadedA, err := ds.Playlist().GetWithTracks(ctx, smartAID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(loadedA.Tracks).To(BeEmpty()) Expect(loadedA.EvaluatedAt).To(BeNil()) - loadedB, err := ds.Playlist(ctx).Get(smartBID) + loadedB, err := ds.Playlist().Get(ctx, smartBID) Expect(err).ToNot(HaveOccurred()) Expect(loadedB.EvaluatedAt).To(BeNil()) }) diff --git a/persistence/folder_repository.go b/persistence/folder_repository.go index a1136ad8d..b4f7c9069 100644 --- a/persistence/folder_repository.go +++ b/persistence/folder_repository.go @@ -64,55 +64,54 @@ func (fs dbFolders) toModels() []model.Folder { return slice.Map(fs, func(f dbFolder) model.Folder { return *f.Folder }) } -func newFolderRepository(ctx context.Context, db dbx.Builder) model.FolderRepository { +func newFolderRepository(db dbx.Builder) model.FolderRepository { r := &folderRepository{} - r.ctx = ctx r.db = db r.tableName = "folder" return r } -func (r folderRepository) selectFolder(options ...model.QueryOptions) SelectBuilder { - sql := r.newSelect(options...).Columns("folder.*", "library.path as library_path"). +func (r folderRepository) selectFolder(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + sql := r.newSelect(ctx, options...).Columns("folder.*", "library.path as library_path"). Join("library on library.id = folder.library_id") - return r.applyLibraryFilter(sql) + return r.applyLibraryFilter(ctx, sql) } -func (r folderRepository) Get(id string) (*model.Folder, error) { - sq := r.selectFolder().Where(Eq{"folder.id": id}) +func (r folderRepository) Get(ctx context.Context, id string) (*model.Folder, error) { + sq := r.selectFolder(ctx).Where(Eq{"folder.id": id}) var res dbFolder - err := r.queryOne(sq, &res) + err := r.queryOne(ctx, sq, &res) return res.Folder, err } -func (r folderRepository) GetByPath(lib model.Library, path string) (*model.Folder, error) { +func (r folderRepository) GetByPath(ctx context.Context, lib model.Library, path string) (*model.Folder, error) { id := model.NewFolder(lib, path).ID - return r.Get(id) + return r.Get(ctx, id) } -func (r folderRepository) GetAll(opt ...model.QueryOptions) ([]model.Folder, error) { - sq := r.selectFolder(opt...) +func (r folderRepository) GetAll(ctx context.Context, opt ...model.QueryOptions) ([]model.Folder, error) { + sq := r.selectFolder(ctx, opt...) var res dbFolders - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) return res.toModels(), err } -func (r folderRepository) CountAll(opt ...model.QueryOptions) (int64, error) { - query := r.newSelect(opt...).Columns("count(*)") - query = r.applyLibraryFilter(query) - return r.count(query) +func (r folderRepository) CountAll(ctx context.Context, opt ...model.QueryOptions) (int64, error) { + query := r.newSelect(ctx, opt...).Columns("count(*)") + query = r.applyLibraryFilter(ctx, query) + return r.count(ctx, query) } -func (r folderRepository) GetFolderUpdateInfo(lib model.Library, targetPaths ...string) (map[string]model.FolderUpdateInfo, error) { +func (r folderRepository) GetFolderUpdateInfo(ctx context.Context, lib model.Library, targetPaths ...string) (map[string]model.FolderUpdateInfo, error) { // If no specific paths, return all folders in the library if len(targetPaths) == 0 { - return r.getFolderUpdateInfoAll(lib) + return r.getFolderUpdateInfoAll(ctx, lib) } // Check if any path is root (return all folders) for _, targetPath := range targetPaths { if targetPath == "" || targetPath == "." { - return r.getFolderUpdateInfoAll(lib) + return r.getFolderUpdateInfoAll(ctx, lib) } } @@ -122,7 +121,7 @@ func (r folderRepository) GetFolderUpdateInfo(lib model.Library, targetPaths ... result := make(map[string]model.FolderUpdateInfo) for batch := range slices.Chunk(targetPaths, batchSize) { - batchResult, err := r.getFolderUpdateInfoBatch(lib, batch) + batchResult, err := r.getFolderUpdateInfoBatch(ctx, lib, batch) if err != nil { return nil, err } @@ -133,16 +132,16 @@ func (r folderRepository) GetFolderUpdateInfo(lib model.Library, targetPaths ... } // getFolderUpdateInfoAll returns update info for all non-missing folders in the library -func (r folderRepository) getFolderUpdateInfoAll(lib model.Library) (map[string]model.FolderUpdateInfo, error) { +func (r folderRepository) getFolderUpdateInfoAll(ctx context.Context, lib model.Library) (map[string]model.FolderUpdateInfo, error) { where := And{ Eq{"library_id": lib.ID}, Eq{"missing": false}, } - return r.queryFolderUpdateInfo(where) + return r.queryFolderUpdateInfo(ctx, where) } // getFolderUpdateInfoBatch returns update info for a batch of target paths and their descendants -func (r folderRepository) getFolderUpdateInfoBatch(lib model.Library, targetPaths []string) (map[string]model.FolderUpdateInfo, error) { +func (r folderRepository) getFolderUpdateInfoBatch(ctx context.Context, lib model.Library, targetPaths []string) (map[string]model.FolderUpdateInfo, error) { where := And{ Eq{"library_id": lib.ID}, Eq{"missing": false}, @@ -172,12 +171,12 @@ func (r folderRepository) getFolderUpdateInfoBatch(lib model.Library, targetPath where = append(where, pathConditions) } - return r.queryFolderUpdateInfo(where) + return r.queryFolderUpdateInfo(ctx, where) } // queryFolderUpdateInfo executes the query and returns the result map -func (r folderRepository) queryFolderUpdateInfo(where And) (map[string]model.FolderUpdateInfo, error) { - sq := r.newSelect().Columns("id", "updated_at", "hash", "image_files", "images_updated_at").Where(where) +func (r folderRepository) queryFolderUpdateInfo(ctx context.Context, where And) (map[string]model.FolderUpdateInfo, error) { + sq := r.newSelect(ctx).Columns("id", "updated_at", "hash", "image_files", "images_updated_at").Where(where) var res []struct { ID string UpdatedAt time.Time @@ -185,7 +184,7 @@ func (r folderRepository) queryFolderUpdateInfo(where And) (map[string]model.Fol ImageFiles string ImagesUpdatedAt time.Time } - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) if err != nil { return nil, err } @@ -230,12 +229,12 @@ func folderSubtreeFilter(lib model.Library, paths []string) Sqlizer { // (including parent itself) contains audio files and is not one of the given // folder IDs. LIKE wildcards in the parent path are escaped, so it is always // matched as a literal prefix. -func (r folderRepository) HasAudioOutsideFolders(parent model.Folder, excludeFolderIDs []string) (bool, error) { +func (r folderRepository) HasAudioOutsideFolders(ctx context.Context, parent model.Folder, excludeFolderIDs []string) (bool, error) { if parent.NumAudioFiles > 0 { return true, nil } parentPath := strings.TrimPrefix(path.Join(parent.Path, parent.Name), "/") - return r.exists(And{ + return r.exists(ctx, And{ Eq{"library_id": parent.LibraryID, "missing": false}, Gt{"num_audio_files": 0}, NotEq{"id": excludeFolderIDs}, @@ -253,20 +252,20 @@ func escapeLikePrefix(s string) string { return strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`).Replace(s) } -func (r folderRepository) Put(f *model.Folder) error { +func (r folderRepository) Put(ctx context.Context, f *model.Folder) error { dbf := dbFolder{Folder: f} - _, err := r.put(dbf.ID, &dbf) + _, err := r.put(ctx, dbf.ID, &dbf) return err } -func (r folderRepository) MarkMissing(missing bool, ids ...string) error { - log.Debug(r.ctx, "Marking folders as missing", "ids", ids, "missing", missing) +func (r folderRepository) MarkMissing(ctx context.Context, missing bool, ids ...string) error { + log.Debug(ctx, "Marking folders as missing", "ids", ids, "missing", missing) for chunk := range slices.Chunk(ids, 200) { sq := Update(r.tableName). Set("missing", missing). Set("updated_at", time.Now()). Where(Eq{"id": chunk}) - _, err := r.executeSQL(sq) + _, err := r.executeSQL(ctx, sq) if err != nil { return err } @@ -274,25 +273,25 @@ func (r folderRepository) MarkMissing(missing bool, ids ...string) error { return nil } -func (r folderRepository) GetTouchedWithPlaylists() (model.FolderCursor, error) { - query := r.selectFolder().Where(And{ +func (r folderRepository) GetTouchedWithPlaylists(ctx context.Context) (model.FolderCursor, error) { + query := r.selectFolder(ctx).Where(And{ Eq{"missing": false}, Gt{"num_playlists": 0}, ConcatExpr("folder.updated_at > library.last_scan_at"), }) - cursor, err := queryWithStableResults[dbFolder](r.sqlRepository, query) + cursor, err := queryWithStableResults[dbFolder](ctx, r.sqlRepository, query) if err != nil { return nil, err } return wrapFolderCursor(cursor), nil } -func (r folderRepository) GetAllWithPlaylists() (model.FolderCursor, error) { - query := r.selectFolder().Where(And{ +func (r folderRepository) GetAllWithPlaylists(ctx context.Context) (model.FolderCursor, error) { + query := r.selectFolder(ctx).Where(And{ Eq{"missing": false}, Gt{"num_playlists": 0}, }) - cursor, err := queryWithStableResults[dbFolder](r.sqlRepository, query) + cursor, err := queryWithStableResults[dbFolder](ctx, r.sqlRepository, query) if err != nil { return nil, err } @@ -303,7 +302,7 @@ func wrapFolderCursor(cursor iter.Seq2[dbFolder, error]) model.FolderCursor { return model.FolderCursor(wrapCursor(cursor, func(f dbFolder) *model.Folder { return f.Folder })) } -func (r folderRepository) purgeEmpty(libraryIDs ...int) error { +func (r folderRepository) purgeEmpty(ctx context.Context, libraryIDs ...int) error { sq := Delete(r.tableName).Where(And{ Eq{"num_audio_files": 0}, Eq{"num_playlists": 0}, @@ -315,12 +314,12 @@ func (r folderRepository) purgeEmpty(libraryIDs ...int) error { if len(libraryIDs) > 0 { sq = sq.Where(Eq{"library_id": libraryIDs}) } - c, err := r.executeSQL(sq) + c, err := r.executeSQL(ctx, sq) if err != nil { return fmt.Errorf("purging empty folders: %w", err) } if c > 0 { - log.Debug(r.ctx, "Purging empty folders", "totalDeleted", c) + log.Debug(ctx, "Purging empty folders", "totalDeleted", c) } return nil } diff --git a/persistence/folder_repository_test.go b/persistence/folder_repository_test.go index b7bc52751..4b1c54a3d 100644 --- a/persistence/folder_repository_test.go +++ b/persistence/folder_repository_test.go @@ -23,17 +23,17 @@ var _ = Describe("FolderRepository", func() { BeforeEach(func() { ctx = request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) conn = GetDBXBuilder() - repo = newFolderRepository(ctx, conn) + repo = newFolderRepository(conn) // Use existing library ID 1 from test fixtures - libRepo := NewLibraryRepository(ctx, conn) - lib, err := libRepo.Get(1) + libRepo := NewLibraryRepository(conn) + lib, err := libRepo.Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) testLib = *lib // Create a second library with its own folder to verify isolation otherLib = model.Library{Name: "Other Library", Path: "/other/path"} - Expect(libRepo.Put(&otherLib)).To(Succeed()) + Expect(libRepo.Put(ctx, &otherLib)).To(Succeed()) }) AfterEach(func() { @@ -48,7 +48,7 @@ var _ = Describe("FolderRepository", func() { matching := func(paths ...string) []string { GinkgoHelper() - folders, err := repo.GetAll(model.QueryOptions{Filters: folderSubtreeFilter(testLib, paths)}) + folders, err := repo.GetAll(ctx, model.QueryOptions{Filters: folderSubtreeFilter(testLib, paths)}) Expect(err).ToNot(HaveOccurred()) return slice.Map(folders, func(f model.Folder) string { return f.ID }) } @@ -59,7 +59,7 @@ var _ = Describe("FolderRepository", func() { grandchild = model.NewFolder(testLib, "TestSubtree/Child/Grandchild") other = model.NewFolder(testLib, "TestSubtreeOther") for _, f := range []*model.Folder{parent, child, grandchild, other} { - Expect(repo.Put(f)).To(Succeed()) + Expect(repo.Put(ctx, f)).To(Succeed()) } DeferCleanup(func() { _, _ = conn.NewQuery("DELETE FROM folder WHERE name LIKE 'TestSubtree%' OR path LIKE 'TestSubtree%'").Execute() @@ -86,17 +86,17 @@ var _ = Describe("FolderRepository", func() { folder1 := model.NewFolder(testLib, "TestGetLastUpdates/Folder1") folder2 := model.NewFolder(testLib, "TestGetLastUpdates/Folder2") - err := repo.Put(folder1) + err := repo.Put(ctx, folder1) Expect(err).ToNot(HaveOccurred()) - err = repo.Put(folder2) + err = repo.Put(ctx, folder2) Expect(err).ToNot(HaveOccurred()) otherFolder := model.NewFolder(otherLib, "TestOtherLib/Folder") - err = repo.Put(otherFolder) + err = repo.Put(ctx, otherFolder) Expect(err).ToNot(HaveOccurred()) // Query all folders (no target paths) - should only return folders from testLib - results, err := repo.GetFolderUpdateInfo(testLib) + results, err := repo.GetFolderUpdateInfo(ctx, testLib) Expect(err).ToNot(HaveOccurred()) // Should include folders from testLib Expect(results).To(HaveKey(folder1.ID)) @@ -113,15 +113,15 @@ var _ = Describe("FolderRepository", func() { folder2 := model.NewFolder(testLib, "TestSpecific/Jazz") folder3 := model.NewFolder(testLib, "TestSpecific/Classical") - err := repo.Put(folder1) + err := repo.Put(ctx, folder1) Expect(err).ToNot(HaveOccurred()) - err = repo.Put(folder2) + err = repo.Put(ctx, folder2) Expect(err).ToNot(HaveOccurred()) - err = repo.Put(folder3) + err = repo.Put(ctx, folder3) Expect(err).ToNot(HaveOccurred()) // Query specific paths - results, err := repo.GetFolderUpdateInfo(testLib, "TestSpecific/Rock", "TestSpecific/Classical") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "TestSpecific/Rock", "TestSpecific/Classical") Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(2)) @@ -142,12 +142,12 @@ var _ = Describe("FolderRepository", func() { child2 := model.NewFolder(testLib, "TestParent/Music/Jazz") otherParent := model.NewFolder(testLib, "TestParent2/Music/Jazz") - Expect(repo.Put(parent)).To(Succeed()) - Expect(repo.Put(child1)).To(Succeed()) - Expect(repo.Put(child2)).To(Succeed()) + Expect(repo.Put(ctx, parent)).To(Succeed()) + Expect(repo.Put(ctx, child1)).To(Succeed()) + Expect(repo.Put(ctx, child2)).To(Succeed()) // Query the parent folder - should return parent and all children - results, err := repo.GetFolderUpdateInfo(testLib, "TestParent/Music") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "TestParent/Music") Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(3)) Expect(results).To(HaveKey(parent.ID)) @@ -161,18 +161,18 @@ var _ = Describe("FolderRepository", func() { parent := model.NewFolder(testLib, "TestIsolation/Parent") child := model.NewFolder(testLib, "TestIsolation/Parent/Child") - Expect(repo.Put(parent)).To(Succeed()) - Expect(repo.Put(child)).To(Succeed()) + Expect(repo.Put(ctx, parent)).To(Succeed()) + Expect(repo.Put(ctx, child)).To(Succeed()) // Create similar path in other library otherParent := model.NewFolder(otherLib, "TestIsolation/Parent") otherChild := model.NewFolder(otherLib, "TestIsolation/Parent/Child") - Expect(repo.Put(otherParent)).To(Succeed()) - Expect(repo.Put(otherChild)).To(Succeed()) + Expect(repo.Put(ctx, otherParent)).To(Succeed()) + Expect(repo.Put(ctx, otherChild)).To(Succeed()) // Query should only return folders from testLib - results, err := repo.GetFolderUpdateInfo(testLib, "TestIsolation/Parent") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "TestIsolation/Parent") Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(2)) Expect(results).To(HaveKey(parent.ID)) @@ -188,12 +188,12 @@ var _ = Describe("FolderRepository", func() { child2 := model.NewFolder(testLib, "TestMissingChild/Parent/Child2") child2.Missing = true - Expect(repo.Put(parent)).To(Succeed()) - Expect(repo.Put(child1)).To(Succeed()) - Expect(repo.Put(child2)).To(Succeed()) + Expect(repo.Put(ctx, parent)).To(Succeed()) + Expect(repo.Put(ctx, child1)).To(Succeed()) + Expect(repo.Put(ctx, child2)).To(Succeed()) // Query parent - should only return parent and non-missing child - results, err := repo.GetFolderUpdateInfo(testLib, "TestMissingChild/Parent") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "TestMissingChild/Parent") Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(2)) Expect(results).To(HaveKey(parent.ID)) @@ -206,11 +206,11 @@ var _ = Describe("FolderRepository", func() { existingParent := model.NewFolder(testLib, "TestMixed/Exists") existingChild := model.NewFolder(testLib, "TestMixed/Exists/Child") - Expect(repo.Put(existingParent)).To(Succeed()) - Expect(repo.Put(existingChild)).To(Succeed()) + Expect(repo.Put(ctx, existingParent)).To(Succeed()) + Expect(repo.Put(ctx, existingChild)).To(Succeed()) // Query both existing and non-existing paths - results, err := repo.GetFolderUpdateInfo(testLib, "TestMixed/Exists", "TestMixed/DoesNotExist") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "TestMixed/Exists", "TestMixed/DoesNotExist") Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(2)) Expect(results).To(HaveKey(existingParent.ID)) @@ -221,7 +221,7 @@ var _ = Describe("FolderRepository", func() { // Test querying for root folder without creating it (fixtures should have one) rootFolderID := model.FolderID(testLib, ".") - results, err := repo.GetFolderUpdateInfo(testLib, "") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "") Expect(err).ToNot(HaveOccurred()) // Should return the root folder if it exists if len(results) > 0 { @@ -230,7 +230,7 @@ var _ = Describe("FolderRepository", func() { }) It("returns empty map for non-existent folders", func() { - results, err := repo.GetFolderUpdateInfo(testLib, "NonExistent/Path") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "NonExistent/Path") Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -239,10 +239,10 @@ var _ = Describe("FolderRepository", func() { // Create a folder and mark it as missing folder := model.NewFolder(testLib, "TestMissing/Folder") folder.Missing = true - err := repo.Put(folder) + err := repo.Put(ctx, folder) Expect(err).ToNot(HaveOccurred()) - results, err := repo.GetFolderUpdateInfo(testLib, "TestMissing/Folder") + results, err := repo.GetFolderUpdateInfo(ctx, testLib, "TestMissing/Folder") Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -262,51 +262,51 @@ var _ = Describe("FolderRepository", func() { disc2 = model.NewFolder(testLib, "TestHasAudio/Album/CD2") disc2.NumAudioFiles = 5 for _, f := range []*model.Folder{albumRoot, disc1, disc2} { - Expect(repo.Put(f)).To(Succeed()) + Expect(repo.Put(ctx, f)).To(Succeed()) } }) It("returns false when all audio under the parent belongs to the given folders", func() { - Expect(repo.HasAudioOutsideFolders(*albumRoot, []string{disc1.ID, disc2.ID})).To(BeFalse()) + Expect(repo.HasAudioOutsideFolders(ctx, *albumRoot, []string{disc1.ID, disc2.ID})).To(BeFalse()) }) It("returns true when another folder under the parent has audio", func() { bonus := model.NewFolder(testLib, "TestHasAudio/Album/Bonus") bonus.NumAudioFiles = 1 - Expect(repo.Put(bonus)).To(Succeed()) + Expect(repo.Put(ctx, bonus)).To(Succeed()) - Expect(repo.HasAudioOutsideFolders(*albumRoot, []string{disc1.ID, disc2.ID})).To(BeTrue()) + Expect(repo.HasAudioOutsideFolders(ctx, *albumRoot, []string{disc1.ID, disc2.ID})).To(BeTrue()) }) It("returns true when the parent itself contains audio files", func() { albumRoot.NumAudioFiles = 2 - Expect(repo.HasAudioOutsideFolders(*albumRoot, []string{disc1.ID, disc2.ID})).To(BeTrue()) + Expect(repo.HasAudioOutsideFolders(ctx, *albumRoot, []string{disc1.ID, disc2.ID})).To(BeTrue()) }) It("ignores audio outside the parent's subtree", func() { other := model.NewFolder(testLib, "TestHasAudio/Other Album") other.NumAudioFiles = 10 - Expect(repo.Put(other)).To(Succeed()) + Expect(repo.Put(ctx, other)).To(Succeed()) - Expect(repo.HasAudioOutsideFolders(*albumRoot, []string{disc1.ID, disc2.ID})).To(BeFalse()) + Expect(repo.HasAudioOutsideFolders(ctx, *albumRoot, []string{disc1.ID, disc2.ID})).To(BeFalse()) }) It("ignores missing folders", func() { gone := model.NewFolder(testLib, "TestHasAudio/Album/Gone") gone.NumAudioFiles = 3 gone.Missing = true - Expect(repo.Put(gone)).To(Succeed()) + Expect(repo.Put(ctx, gone)).To(Succeed()) - Expect(repo.HasAudioOutsideFolders(*albumRoot, []string{disc1.ID, disc2.ID})).To(BeFalse()) + Expect(repo.HasAudioOutsideFolders(ctx, *albumRoot, []string{disc1.ID, disc2.ID})).To(BeFalse()) }) It("does not treat LIKE wildcards in the parent path as patterns", func() { // "TestHas_udio" would LIKE-match "TestHasAudio" if "_" were not escaped wildcardRoot := model.NewFolder(testLib, "TestHas_udio/Album") - Expect(repo.Put(wildcardRoot)).To(Succeed()) + Expect(repo.Put(ctx, wildcardRoot)).To(Succeed()) - Expect(repo.HasAudioOutsideFolders(*wildcardRoot, []string{"none"})).To(BeFalse()) + Expect(repo.HasAudioOutsideFolders(ctx, *wildcardRoot, []string{"none"})).To(BeFalse()) }) }) @@ -367,9 +367,9 @@ var _ = Describe("FolderRepository", func() { missingWithPls.NumPlaylists = 1 missingWithPls.Missing = true - Expect(repo.Put(withPls)).To(Succeed()) - Expect(repo.Put(noPls)).To(Succeed()) - Expect(repo.Put(missingWithPls)).To(Succeed()) + Expect(repo.Put(ctx, withPls)).To(Succeed()) + Expect(repo.Put(ctx, noPls)).To(Succeed()) + Expect(repo.Put(ctx, missingWithPls)).To(Succeed()) // Force the folder's updated_at to the past so GetTouchedWithPlaylists // (which gates on updated_at > last_scan_at) would NOT return it. @@ -378,7 +378,7 @@ var _ = Describe("FolderRepository", func() { Expect(err).ToNot(HaveOccurred()) var ids []string - cursor, err := repo.GetAllWithPlaylists() + cursor, err := repo.GetAllWithPlaylists(ctx) Expect(err).ToNot(HaveOccurred()) for f, err := range cursor { Expect(err).ToNot(HaveOccurred()) diff --git a/persistence/genre_repository.go b/persistence/genre_repository.go index 0bb22c21b..d88cb8672 100644 --- a/persistence/genre_repository.go +++ b/persistence/genre_repository.go @@ -13,43 +13,39 @@ type genreRepository struct { *baseTagRepository } -func NewGenreRepository(ctx context.Context, db dbx.Builder) model.GenreRepository { +func NewGenreRepository(db dbx.Builder) model.GenreRepository { return &genreRepository{ - baseTagRepository: newBaseTagRepository(ctx, db, new(model.TagGenre)), + baseTagRepository: newBaseTagRepository(db, new(model.TagGenre)), } } -func (r *genreRepository) selectGenre(opt ...model.QueryOptions) SelectBuilder { - return r.newSelect(opt...).Columns("tag.tag_value as name") +func (r *genreRepository) selectGenre(ctx context.Context, opt ...model.QueryOptions) SelectBuilder { + return r.newSelect(ctx, opt...).Columns("tag.tag_value as name") } -func (r *genreRepository) GetAll(opt ...model.QueryOptions) (model.Genres, error) { - sq := r.selectGenre(opt...) +func (r *genreRepository) GetAll(ctx context.Context, opt ...model.QueryOptions) (model.Genres, error) { + sq := r.selectGenre(ctx, opt...) res := model.Genres{} - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) return res, err } -func (r *genreRepository) Get(id string) (*model.Genre, error) { - sel := r.selectGenre().Where(Eq{"tag.id": id}) +func (r *genreRepository) Get(ctx context.Context, id string) (*model.Genre, error) { + sel := r.selectGenre(ctx).Where(Eq{"tag.id": id}) var res model.Genre - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -// Override ResourceRepository methods to return Genre objects instead of Tag objects +// Override the base tag REST methods to return Genre objects instead of Tag objects -func (r *genreRepository) Read(id string) (any, error) { - return r.Get(id) +func (r *genreRepository) Read(ctx context.Context, id string) (*model.Genre, error) { + return r.Get(ctx, id) } -func (r *genreRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) -} - -func (r *genreRepository) NewInstance() any { - return &model.Genre{} +func (r *genreRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Genre, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } var _ model.GenreRepository = (*genreRepository)(nil) -var _ model.ResourceRepository = (*genreRepository)(nil) +var _ rest.Repository[model.Genre] = (*genreRepository)(nil) diff --git a/persistence/genre_repository_test.go b/persistence/genre_repository_test.go index e3779725c..52bebf446 100644 --- a/persistence/genre_repository_test.go +++ b/persistence/genre_repository_test.go @@ -16,16 +16,16 @@ import ( var _ = Describe("GenreRepository", func() { var repo model.GenreRepository - var restRepo model.ResourceRepository + var restRepo rest.Repository[model.Genre] var tagRepo model.TagRepository var ctx context.Context BeforeEach(func() { ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "userid", UserName: "johndoe", IsAdmin: true}) - genreRepo := NewGenreRepository(ctx, GetDBXBuilder()) + genreRepo := NewGenreRepository(GetDBXBuilder()) repo = genreRepo - restRepo = genreRepo.(model.ResourceRepository) - tagRepo = NewTagRepository(ctx, GetDBXBuilder()) + restRepo = genreRepo + tagRepo = NewTagRepository(GetDBXBuilder()) // Clear any existing tags to ensure test isolation db := GetDBXBuilder() @@ -43,7 +43,7 @@ var _ = Describe("GenreRepository", func() { return model.Tag{ID: id.NewTagID(name, value), TagName: model.TagName(name), TagValue: value} } - err = tagRepo.Add(1, + err = tagRepo.Add(ctx, 1, newTag("genre", "rock"), newTag("genre", "pop"), newTag("genre", "jazz"), @@ -65,7 +65,7 @@ var _ = Describe("GenreRepository", func() { Describe("GetAll", func() { It("should return all genres", func() { - genres, err := repo.GetAll() + genres, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(genres).To(HaveLen(12)) @@ -83,7 +83,7 @@ var _ = Describe("GenreRepository", func() { It("should support query options", func() { // Test with limiting results - genres, err := repo.GetAll(model.QueryOptions{Max: 1}) + genres, err := repo.GetAll(ctx, model.QueryOptions{Max: 1}) Expect(err).ToNot(HaveOccurred()) Expect(genres).To(HaveLen(1)) }) @@ -93,7 +93,7 @@ var _ = Describe("GenreRepository", func() { _, err := GetDBXBuilder().NewQuery("DELETE FROM tag WHERE tag_name = 'genre'").Execute() Expect(err).ToNot(HaveOccurred()) - genres, err := repo.GetAll() + genres, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(genres).To(BeEmpty()) }) @@ -103,7 +103,7 @@ var _ = Describe("GenreRepository", func() { options := model.QueryOptions{ Filters: squirrel.Like{"tag_value": "%rock%"}, // Direct field access } - genres, err := repo.GetAll(options) + genres, err := repo.GetAll(ctx, options) Expect(err).ToNot(HaveOccurred()) Expect(genres).To(HaveLen(2)) // Should match "rock" and "Alternative Rock" @@ -119,7 +119,7 @@ var _ = Describe("GenreRepository", func() { Filters: squirrel.Like{"tag_value": "%e%"}, // Should match genres containing "e" Sort: "name", } - genres, err := repo.GetAll(options) + genres, err := repo.GetAll(ctx, options) Expect(err).ToNot(HaveOccurred()) Expect(genres).To(HaveLen(7)) @@ -135,7 +135,7 @@ var _ = Describe("GenreRepository", func() { Sort: "name", Order: "desc", } - genres, err := repo.GetAll(options) + genres, err := repo.GetAll(ctx, options) Expect(err).ToNot(HaveOccurred()) Expect(genres).To(HaveLen(7)) @@ -148,7 +148,7 @@ var _ = Describe("GenreRepository", func() { Describe("Count", func() { It("should return correct count of genres", func() { - count, err := restRepo.Count() + count, err := restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(12))) // We have 12 genre tags }) @@ -158,7 +158,7 @@ var _ = Describe("GenreRepository", func() { _, err := GetDBXBuilder().NewQuery("DELETE FROM tag WHERE tag_name = 'genre'").Execute() Expect(err).ToNot(HaveOccurred()) - count, err := restRepo.Count() + count, err := restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeZero()) }) @@ -170,10 +170,10 @@ var _ = Describe("GenreRepository", func() { TagName: "mood", TagValue: "energetic", } - err := tagRepo.Add(1, nonGenreTag) + err := tagRepo.Add(ctx, 1, nonGenreTag) Expect(err).ToNot(HaveOccurred()) - count, err := restRepo.Count() + count, err := restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) // Count should not include the mood tag Expect(count).To(Equal(int64(12))) // Should still be 12 genre tags @@ -184,7 +184,7 @@ var _ = Describe("GenreRepository", func() { options := rest.QueryOptions{ Filters: map[string]any{"name": "%rock%"}, } - count, err := restRepo.Count(options) + count, err := restRepo.Count(ctx, options) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeNumerically("==", 2)) }) @@ -194,30 +194,28 @@ var _ = Describe("GenreRepository", func() { It("should return existing genre", func() { // Use one of the existing genres from our consolidated dataset genreID := id.NewTagID("genre", "rock") - result, err := restRepo.Read(genreID) + genre, err := restRepo.Read(ctx, genreID) Expect(err).ToNot(HaveOccurred()) - genre := result.(*model.Genre) Expect(genre.ID).To(Equal(genreID)) Expect(genre.Name).To(Equal("rock")) }) It("should return error for non-existent genre", func() { - _, err := restRepo.Read("non-existent-id") + _, err := restRepo.Read(ctx, "non-existent-id") Expect(err).To(HaveOccurred()) }) It("should not return non-genre tags", func() { moodID := id.NewTagID("mood", "happy") // This exists as a mood tag, not genre - _, err := restRepo.Read(moodID) + _, err := restRepo.Read(ctx, moodID) Expect(err).To(HaveOccurred()) // Should not find it as a genre }) }) Describe("ReadAll", func() { It("should return all genres through ReadAll", func() { - result, err := restRepo.ReadAll() + genres, err := restRepo.ReadAll(ctx) Expect(err).ToNot(HaveOccurred()) - genres := result.(model.Genres) Expect(genres).To(HaveLen(12)) // We have 12 genre tags genreNames := make([]string, len(genres)) @@ -231,7 +229,7 @@ var _ = Describe("GenreRepository", func() { }) It("should support rest query options", func() { - result, err := restRepo.ReadAll() + result, err := restRepo.ReadAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(result).ToNot(BeNil()) }) @@ -240,13 +238,15 @@ var _ = Describe("GenreRepository", func() { Describe("Library Filtering", func() { Context("Headless Processes (No User Context)", func() { var headlessRepo model.GenreRepository - var headlessRestRepo model.ResourceRepository + var headlessRestRepo rest.Repository[model.Genre] + var headlessCtx context.Context BeforeEach(func() { + headlessCtx = GinkgoT().Context() // Create a repository with no user context (headless) - headlessGenreRepo := NewGenreRepository(context.Background(), GetDBXBuilder()) + headlessGenreRepo := NewGenreRepository(GetDBXBuilder()) headlessRepo = headlessGenreRepo - headlessRestRepo = headlessGenreRepo.(model.ResourceRepository) + headlessRestRepo = headlessGenreRepo // Add genres to different libraries db := GetDBXBuilder() @@ -258,13 +258,13 @@ var _ = Describe("GenreRepository", func() { return model.Tag{ID: id.NewTagID(name, value), TagName: model.TagName(name), TagValue: value} } - err = tagRepo.Add(2, newTag("genre", "jazz")) + err = tagRepo.Add(ctx, 2, newTag("genre", "jazz")) Expect(err).ToNot(HaveOccurred()) }) It("should see all genres from all libraries when no user is in context", func() { // Headless processes should see all genres regardless of library - genres, err := headlessRepo.GetAll() + genres, err := headlessRepo.GetAll(headlessCtx) Expect(err).ToNot(HaveOccurred()) // Should see genres from all libraries @@ -279,7 +279,7 @@ var _ = Describe("GenreRepository", func() { }) It("should count all genres from all libraries when no user is in context", func() { - count, err := headlessRestRepo.Count() + count, err := headlessRestRepo.Count(headlessCtx) Expect(err).ToNot(HaveOccurred()) // Should count all genres from all libraries @@ -288,12 +288,11 @@ var _ = Describe("GenreRepository", func() { It("should allow headless processes to apply explicit library_id filters", func() { // Filter by specific library - genres, err := headlessRestRepo.ReadAll(rest.QueryOptions{ + genreList, err := headlessRestRepo.ReadAll(headlessCtx, rest.QueryOptions{ Filters: map[string]any{"library_id": 2}, }) Expect(err).ToNot(HaveOccurred()) - genreList := genres.(model.Genres) // Should see only genres from library 2 Expect(genreList).To(HaveLen(1)) Expect(genreList[0].Name).To(Equal("jazz")) @@ -301,29 +300,15 @@ var _ = Describe("GenreRepository", func() { It("should get individual genres when no user is in context", func() { // Get all genres first to find an ID - genres, err := headlessRepo.GetAll() + genres, err := headlessRepo.GetAll(headlessCtx) Expect(err).ToNot(HaveOccurred()) Expect(genres).ToNot(BeEmpty()) // Headless process should be able to get the genre - genre, err := headlessRestRepo.Read(genres[0].ID) + genre, err := headlessRestRepo.Read(headlessCtx, genres[0].ID) Expect(err).ToNot(HaveOccurred()) Expect(genre).ToNot(BeNil()) }) }) }) - - Describe("EntityName", func() { - It("should return correct entity name", func() { - name := restRepo.EntityName() - Expect(name).To(Equal("tag")) // Genre repository uses tag table - }) - }) - - Describe("NewInstance", func() { - It("should return new genre instance", func() { - instance := restRepo.NewInstance() - Expect(instance).To(BeAssignableToTypeOf(&model.Genre{})) - }) - }) }) diff --git a/persistence/item_tags_test.go b/persistence/item_tags_test.go index a2317a1a5..1106c7e9c 100644 --- a/persistence/item_tags_test.go +++ b/persistence/item_tags_test.go @@ -1,6 +1,8 @@ package persistence import ( + "context" + "github.com/deluan/rest" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" @@ -14,17 +16,18 @@ var _ = Describe("item genre tag indexes", func() { var mr model.MediaFileRepository var ar model.AlbumRepository var rock, jazz model.Tag + var ctx context.Context BeforeEach(func() { - ctx := request.WithUser(GinkgoT().Context(), model.User{ID: "userid"}) + ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "userid"}) conn = GetDBXBuilder() - mr = NewMediaFileRepository(ctx, conn) - ar = NewAlbumRepository(ctx, conn) + mr = NewMediaFileRepository(conn) + ar = NewAlbumRepository(conn) // Test-only genre values, so they can't collide with the golden fixtures. rock = model.NewTag(model.TagGenre, "GenreIdxRock") jazz = model.NewTag(model.TagGenre, "GenreIdxJazz") // The join tables FK to tag(id); the scanner adds tags before saving items. - Expect(NewTagRepository(ctx, conn).Add(1, rock, jazz)).To(Succeed()) + Expect(NewTagRepository(conn).Add(ctx, 1, rock, jazz)).To(Succeed()) // The suite shares one golden DB with no per-test restore, so undo the rows we add // (media_file/album deletes cascade to the *_tags join rows; tag deletes cascade too). DeferCleanup(func() { @@ -53,25 +56,25 @@ var _ = Describe("item genre tag indexes", func() { It("writes a media_file_tags row for each genre when the track is saved", func() { mf := model.MediaFile{ID: "mf-g1", LibraryID: 1, Path: "/m/g1.mp3", Title: "G1", Tags: model.Tags{model.TagGenre: []string{rock.TagValue, jazz.TagValue}}} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) Expect(tagIDsFor("media_file_tags", "media_file_id", "mf-g1")).To(ConsistOf(rock.ID, jazz.ID)) }) It("replaces the rows when the genres change", func() { mf := model.MediaFile{ID: "mf-g2", LibraryID: 1, Path: "/m/g2.mp3", Title: "G2", Tags: model.Tags{model.TagGenre: []string{rock.TagValue}}} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) mf.Tags = model.Tags{model.TagGenre: []string{jazz.TagValue}} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) Expect(tagIDsFor("media_file_tags", "media_file_id", "mf-g2")).To(ConsistOf(jazz.ID)) }) It("clears the rows when all genres are removed", func() { mf := model.MediaFile{ID: "mf-g3", LibraryID: 1, Path: "/m/g3.mp3", Title: "G3", Tags: model.Tags{model.TagGenre: []string{rock.TagValue}}} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) mf.Tags = model.Tags{} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) Expect(tagIDsFor("media_file_tags", "media_file_id", "mf-g3")).To(BeEmpty()) }) }) @@ -80,7 +83,7 @@ var _ = Describe("item genre tag indexes", func() { It("writes an album_tags row for each genre when the album is saved", func() { al := model.Album{ID: "al-g1", LibraryID: 1, Name: "AG1", Tags: model.Tags{model.TagGenre: []string{rock.TagValue, jazz.TagValue}}} - Expect(ar.Put(&al)).To(Succeed()) + Expect(ar.Put(ctx, &al)).To(Succeed()) Expect(tagIDsFor("album_tags", "album_id", "al-g1")).To(ConsistOf(rock.ID, jazz.ID)) }) }) @@ -90,19 +93,19 @@ var _ = Describe("item genre tag indexes", func() { It("filters media files by genre_id", func() { mf := model.MediaFile{ID: "mf-nat1", LibraryID: 1, Path: "/m/nat1.mp3", Title: "Nat1", Tags: model.Tags{model.TagGenre: []string{rock.TagValue}}} - Expect(mr.Put(&mf)).To(Succeed()) - res, err := mr.(model.ResourceRepository).ReadAll(rest.QueryOptions{Filters: map[string]any{"genre_id": rock.ID}}) + Expect(mr.Put(ctx, &mf)).To(Succeed()) + res, err := mr.ReadAll(ctx, rest.QueryOptions{Filters: map[string]any{"genre_id": rock.ID}}) Expect(err).ToNot(HaveOccurred()) - Expect(res.(model.MediaFiles)).To(ContainElement(HaveField("ID", "mf-nat1"))) + Expect(res).To(ContainElement(HaveField("ID", "mf-nat1"))) }) It("filters albums by genre_id", func() { al := model.Album{ID: "al-nat1", LibraryID: 1, Name: "ANat1", Tags: model.Tags{model.TagGenre: []string{rock.TagValue}}} - Expect(ar.Put(&al)).To(Succeed()) - res, err := ar.(model.ResourceRepository).ReadAll(rest.QueryOptions{Filters: map[string]any{"genre_id": rock.ID}}) + Expect(ar.Put(ctx, &al)).To(Succeed()) + res, err := ar.ReadAll(ctx, rest.QueryOptions{Filters: map[string]any{"genre_id": rock.ID}}) Expect(err).ToNot(HaveOccurred()) - Expect(res.(model.Albums)).To(ContainElement(HaveField("ID", "al-nat1"))) + Expect(res).To(ContainElement(HaveField("ID", "al-nat1"))) }) }) }) diff --git a/persistence/library_repository.go b/persistence/library_repository.go index 2e8feea7a..bf6b8995e 100644 --- a/persistence/library_repository.go +++ b/persistence/library_repository.go @@ -25,22 +25,21 @@ var ( libLock sync.RWMutex ) -func NewLibraryRepository(ctx context.Context, db dbx.Builder) model.LibraryRepository { +func NewLibraryRepository(db dbx.Builder) model.LibraryRepository { r := &libraryRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Library{}, nil) return r } -func (r *libraryRepository) Get(id int) (*model.Library, error) { - sq := r.newSelect().Columns("*").Where(Eq{"id": id}) +func (r *libraryRepository) Get(ctx context.Context, id int) (*model.Library, error) { + sq := r.newSelect(ctx).Columns("*").Where(Eq{"id": id}) var res model.Library - err := r.queryOne(sq, &res) + err := r.queryOne(ctx, sq, &res) return &res, err } -func (r *libraryRepository) GetPath(id int) (string, error) { +func (r *libraryRepository) GetPath(ctx context.Context, id int) (string, error) { l := func() string { libLock.RLock() defer libLock.RUnlock() @@ -55,9 +54,9 @@ func (r *libraryRepository) GetPath(id int) (string, error) { libLock.Lock() defer libLock.Unlock() - libs, err := r.GetAll() + libs, err := r.GetAll(ctx) if err != nil { - log.Error(r.ctx, "Error loading libraries from DB", err) + log.Error(ctx, "Error loading libraries from DB", err) return "", err } for _, l := range libs { @@ -70,9 +69,9 @@ func (r *libraryRepository) GetPath(id int) (string, error) { } } -func (r *libraryRepository) Put(l *model.Library, colsToUpdate ...string) error { +func (r *libraryRepository) Put(ctx context.Context, l *model.Library, colsToUpdate ...string) error { if l.ID == model.DefaultLibraryID { - currentLib, err := r.Get(1) + currentLib, err := r.Get(ctx, 1) // if we are creating it, it's ok. if err == nil { // it exists, so we are updating it if currentLib.Path != l.Path { @@ -97,7 +96,7 @@ func (r *libraryRepository) Put(l *model.Library, colsToUpdate ...string) error }, colsToUpdate...) cols["updated_at"] = l.UpdatedAt sq := Update(r.tableName).SetMap(cols).Where(Eq{"id": l.ID}) - rowsAffected, updateErr := r.executeSQL(sq) + rowsAffected, updateErr := r.executeSQL(ctx, sq) if updateErr != nil { return updateErr } @@ -122,7 +121,7 @@ CROSS JOIN library l WHERE u.is_admin = true ON CONFLICT (user_id, library_id) DO NOTHING;`, ) - if _, err = r.executeSQL(sql); err != nil { + if _, err = r.executeSQL(ctx, sql); err != nil { return fmt.Errorf("failed to assign library to admin users: %w", err) } @@ -134,12 +133,12 @@ ON CONFLICT (user_id, library_id) DO NOTHING;`, // TODO Remove this method when we have a proper UI to add libraries // This is a temporary method to store the music folder path from the config in the DB -func (r *libraryRepository) StoreMusicFolder() error { +func (r *libraryRepository) StoreMusicFolder(ctx context.Context) error { sq := Update(r.tableName).Set("path", conf.Server.MusicFolder). Set("updated_at", time.Now()). Where(Eq{"id": model.DefaultLibraryID}). Where(NotEq{"path": conf.Server.MusicFolder}) - rowsAffected, err := r.executeSQL(sq) + rowsAffected, err := r.executeSQL(ctx, sq) if err == nil && rowsAffected > 0 { libLock.Lock() defer libLock.Unlock() @@ -148,77 +147,77 @@ func (r *libraryRepository) StoreMusicFolder() error { return err } -func (r *libraryRepository) AddArtist(id int, artistID string) error { +func (r *libraryRepository) AddArtist(ctx context.Context, id int, artistID string) error { sq := Insert("library_artist").Columns("library_id", "artist_id").Values(id, artistID). Suffix(`on conflict(library_id, artist_id) do nothing`) - _, err := r.executeSQL(sq) + _, err := r.executeSQL(ctx, sq) if err != nil { return err } return nil } -func (r *libraryRepository) ScanBegin(id int, fullScan bool) error { +func (r *libraryRepository) ScanBegin(ctx context.Context, id int, fullScan bool) error { sq := Update(r.tableName). Set("last_scan_started_at", time.Now()). Set("full_scan_in_progress", fullScan). Where(Eq{"id": id}) - _, err := r.executeSQL(sq) + _, err := r.executeSQL(ctx, sq) return err } -func (r *libraryRepository) ScanEnd(id int) error { +func (r *libraryRepository) ScanEnd(ctx context.Context, id int) error { sq := Update(r.tableName). Set("last_scan_at", time.Now()). Set("full_scan_in_progress", false). Set("last_scan_started_at", time.Time{}). Where(Eq{"id": id}) - _, err := r.executeSQL(sq) + _, err := r.executeSQL(ctx, sq) return err } -func (r *libraryRepository) ScanInProgress() (bool, error) { - query := r.newSelect().Where(NotEq{"last_scan_started_at": time.Time{}}) - count, err := r.count(query) +func (r *libraryRepository) ScanInProgress(ctx context.Context) (bool, error) { + query := r.newSelect(ctx).Where(NotEq{"last_scan_started_at": time.Time{}}) + count, err := r.count(ctx, query) return count > 0, err } -func (r *libraryRepository) RefreshStats(id int) error { +func (r *libraryRepository) RefreshStats(ctx context.Context, id int) error { var songsRes, albumsRes, artistsRes, foldersRes, filesRes, missingRes struct{ Count int64 } var sizeRes struct{ Sum int64 } var durationRes struct{ Sum float64 } err := run.Parallel( func() error { - return r.queryOne(Select("count(*) as count").From("media_file").Where(Eq{"library_id": id, "missing": false}), &songsRes) + return r.queryOne(ctx, Select("count(*) as count").From("media_file").Where(Eq{"library_id": id, "missing": false}), &songsRes) }, func() error { - return r.queryOne(Select("count(*) as count").From("album").Where(Eq{"library_id": id, "missing": false}), &albumsRes) + return r.queryOne(ctx, Select("count(*) as count").From("album").Where(Eq{"library_id": id, "missing": false}), &albumsRes) }, func() error { - return r.queryOne(Select("count(*) as count").From("library_artist la"). + return r.queryOne(ctx, Select("count(*) as count").From("library_artist la"). Join("artist a on la.artist_id = a.id"). Where(Eq{"la.library_id": id, "a.missing": false}), &artistsRes) }, func() error { - return r.queryOne(Select("count(*) as count").From("folder"). + return r.queryOne(ctx, Select("count(*) as count").From("folder"). Where(And{ Eq{"library_id": id, "missing": false}, Gt{"num_audio_files": 0}, }), &foldersRes) }, func() error { - return r.queryOne(Select("ifnull(sum(num_audio_files + num_playlists + json_array_length(image_files)),0) as count"). + return r.queryOne(ctx, Select("ifnull(sum(num_audio_files + num_playlists + json_array_length(image_files)),0) as count"). From("folder").Where(Eq{"library_id": id, "missing": false}), &filesRes) }, func() error { - return r.queryOne(Select("count(*) as count").From("media_file").Where(Eq{"library_id": id, "missing": true}), &missingRes) + return r.queryOne(ctx, Select("count(*) as count").From("media_file").Where(Eq{"library_id": id, "missing": true}), &missingRes) }, func() error { - return r.queryOne(Select("ifnull(sum(size),0) as sum").From("album").Where(Eq{"library_id": id, "missing": false}), &sizeRes) + return r.queryOne(ctx, Select("ifnull(sum(size),0) as sum").From("album").Where(Eq{"library_id": id, "missing": false}), &sizeRes) }, func() error { - return r.queryOne(Select("ifnull(sum(duration),0) as sum").From("album").Where(Eq{"library_id": id, "missing": false}), &durationRes) + return r.queryOne(ctx, Select("ifnull(sum(duration),0) as sum").From("album").Where(Eq{"library_id": id, "missing": false}), &durationRes) }, )() if err != nil { @@ -236,25 +235,25 @@ func (r *libraryRepository) RefreshStats(id int) error { Set("total_duration", durationRes.Sum). Set("updated_at", time.Now()). Where(Eq{"id": id}) - _, err = r.executeSQL(sq) + _, err = r.executeSQL(ctx, sq) return err } -func (r *libraryRepository) Delete(id int) error { - if !loggedUser(r.ctx).IsAdmin { +func (r *libraryRepository) Delete(ctx context.Context, id int) error { + if !loggedUser(ctx).IsAdmin { return model.ErrNotAuthorized } if id == 1 { return fmt.Errorf("%w: library with ID 1 cannot be deleted", model.ErrValidation) } - err := r.delete(Eq{"id": id}) + err := r.delete(ctx, Eq{"id": id}) if err != nil { return err } // The cascade above can drop an artist's last library_artist row; reconcile any such orphans. - if err := NewArtistRepository(r.ctx, r.db).(*artistRepository).markOrphansMissing(); err != nil { + if err := NewArtistRepository(r.db).(*artistRepository).markOrphansMissing(ctx); err != nil { return fmt.Errorf("marking orphaned artists missing after deleting library %d: %w", id, err) } @@ -265,26 +264,26 @@ func (r *libraryRepository) Delete(id int) error { // Clean up orphaned plugin references for the deleted library if err := cleanupPluginLibraryReferences(r.db, id); err != nil { - log.Error(r.ctx, "Failed to cleanup plugin library references", "libraryID", id, err) + log.Error(ctx, "Failed to cleanup plugin library references", "libraryID", id, err) } return nil } -func (r *libraryRepository) GetAll(ops ...model.QueryOptions) (model.Libraries, error) { - sq := r.newSelect(ops...).Columns("*") +func (r *libraryRepository) GetAll(ctx context.Context, ops ...model.QueryOptions) (model.Libraries, error) { + sq := r.newSelect(ctx, ops...).Columns("*") res := model.Libraries{} - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) return res, err } -func (r *libraryRepository) CountAll(ops ...model.QueryOptions) (int64, error) { - sq := r.newSelect(ops...) - return r.count(sq) +func (r *libraryRepository) CountAll(ctx context.Context, ops ...model.QueryOptions) (int64, error) { + sq := r.newSelect(ctx, ops...) + return r.count(ctx, sq) } // User-library association methods -func (r *libraryRepository) GetUsersWithLibraryAccess(libraryID int) (model.Users, error) { +func (r *libraryRepository) GetUsersWithLibraryAccess(ctx context.Context, libraryID int) (model.Users, error) { sel := Select("u.*"). From("user u"). Join("user_library ul ON u.id = ul.user_id"). @@ -292,57 +291,28 @@ func (r *libraryRepository) GetUsersWithLibraryAccess(libraryID int) (model.User OrderBy("u.name") var res model.Users - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) return res, err } // REST interface methods -func (r *libraryRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *libraryRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *libraryRepository) Read(id string) (any, error) { +func (r *libraryRepository) Read(ctx context.Context, id string) (*model.Library, error) { idInt, err := strconv.Atoi(id) if err != nil { - log.Trace(r.ctx, "invalid library id: %s", id, err) + log.Trace(ctx, "invalid library id: %s", id, err) return nil, rest.ErrNotFound } - return r.Get(idInt) + return r.Get(ctx, idInt) } -func (r *libraryRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) -} - -func (r *libraryRepository) EntityName() string { - return "library" -} - -func (r *libraryRepository) NewInstance() any { - return &model.Library{} -} - -func (r *libraryRepository) Save(entity any) (string, error) { - lib := entity.(*model.Library) - lib.ID = 0 // Reset ID to ensure we create a new library - err := r.Put(lib) - if err != nil { - return "", err - } - return strconv.Itoa(lib.ID), nil -} - -func (r *libraryRepository) Update(id string, entity any, cols ...string) error { - lib := entity.(*model.Library) - idInt, err := strconv.Atoi(id) - if err != nil { - return fmt.Errorf("invalid library ID: %s", id) - } - - lib.ID = idInt - return r.Put(lib, cols...) +func (r *libraryRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Library, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } var _ model.LibraryRepository = (*libraryRepository)(nil) -var _ rest.Repository = (*libraryRepository)(nil) +var _ rest.Repository[model.Library] = (*libraryRepository)(nil) diff --git a/persistence/library_repository_test.go b/persistence/library_repository_test.go index 6aede8c4e..0ff470861 100644 --- a/persistence/library_repository_test.go +++ b/persistence/library_repository_test.go @@ -23,7 +23,7 @@ var _ = Describe("LibraryRepository", func() { BeforeEach(func() { ctx = request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) conn = GetDBXBuilder() - repo = NewLibraryRepository(ctx, conn) + repo = NewLibraryRepository(conn) }) AfterEach(func() { @@ -40,14 +40,14 @@ var _ = Describe("LibraryRepository", func() { Path: "/music/test", } - err := repo.Put(lib) + err := repo.Put(ctx, lib) Expect(err).ToNot(HaveOccurred()) Expect(lib.ID).To(BeNumerically(">", 0)) Expect(lib.CreatedAt).ToNot(BeZero()) Expect(lib.UpdatedAt).ToNot(BeZero()) // Verify it was inserted - savedLib, err := repo.Get(lib.ID) + savedLib, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(savedLib.Name).To(Equal("Test Library")) Expect(savedLib.Path).To(Equal("/music/test")) @@ -62,11 +62,11 @@ var _ = Describe("LibraryRepository", func() { RemotePath: "/remote/original", DefaultNewUsers: true, } - Expect(repo.Put(lib)).To(Succeed()) + Expect(repo.Put(ctx, lib)).To(Succeed()) - Expect(repo.Put(&model.Library{ID: lib.ID, Name: "Renamed", Path: lib.Path}, "name", "path")).To(Succeed()) + Expect(repo.Put(ctx, &model.Library{ID: lib.ID, Name: "Renamed", Path: lib.Path}, "name", "path")).To(Succeed()) - saved, err := repo.Get(lib.ID) + saved, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(saved.Name).To(Equal("Renamed")) Expect(saved.RemotePath).To(Equal("/remote/original")) @@ -82,7 +82,7 @@ var _ = Describe("LibraryRepository", func() { Name: "Original Library", Path: "/music/original", } - err := repo.Put(lib) + err := repo.Put(ctx, lib) Expect(err).ToNot(HaveOccurred()) originalID := lib.ID @@ -96,7 +96,7 @@ var _ = Describe("LibraryRepository", func() { // Now update it lib.Name = "Updated Library" lib.Path = "/music/updated" - err = repo.Put(lib) + err = repo.Put(ctx, lib) Expect(err).ToNot(HaveOccurred()) // Verify it was updated, not inserted @@ -105,7 +105,7 @@ var _ = Describe("LibraryRepository", func() { Expect(lib.UpdatedAt).To(BeTemporally(">", originalCreatedAt)) // Verify the changes were saved - savedLib, err := repo.Get(lib.ID) + savedLib, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(savedLib.Name).To(Equal("Updated Library")) Expect(savedLib.Path).To(Equal("/music/updated")) @@ -121,18 +121,18 @@ var _ = Describe("LibraryRepository", func() { } // Ensure the record doesn't exist - _, err := repo.Get(999) + _, err := repo.Get(ctx, 999) Expect(err).To(HaveOccurred()) // Put should insert it - err = repo.Put(lib) + err = repo.Put(ctx, lib) Expect(err).ToNot(HaveOccurred()) Expect(lib.ID).To(Equal(999)) Expect(lib.CreatedAt).ToNot(BeZero()) Expect(lib.UpdatedAt).ToNot(BeZero()) // Verify it was inserted with the correct ID - savedLib, err := repo.Get(999) + savedLib, err := repo.Get(ctx, 999) Expect(err).ToNot(HaveOccurred()) Expect(savedLib.ID).To(Equal(999)) Expect(savedLib.Name).To(Equal("New Library with ID")) @@ -146,7 +146,7 @@ var _ = Describe("LibraryRepository", func() { BeforeEach(func() { var err error - libBefore, err = repo.Get(model.DefaultLibraryID) + libBefore, err = repo.Get(ctx, model.DefaultLibraryID) Expect(err).ToNot(HaveOccurred()) DeferCleanup(configtest.SetupConfig()) @@ -162,9 +162,9 @@ var _ = Describe("LibraryRepository", func() { It("skips updating the default library when the configured path is unchanged", func() { conf.Server.MusicFolder = libBefore.Path - Expect(repo.StoreMusicFolder()).To(Succeed()) + Expect(repo.StoreMusicFolder(ctx)).To(Succeed()) - libAfter, err := repo.Get(model.DefaultLibraryID) + libAfter, err := repo.Get(ctx, model.DefaultLibraryID) Expect(err).ToNot(HaveOccurred()) Expect(libAfter.Path).To(Equal(libBefore.Path)) Expect(libAfter.UpdatedAt).To(Equal(libBefore.UpdatedAt)) @@ -172,9 +172,9 @@ var _ = Describe("LibraryRepository", func() { It("updates the default library only when the configured path changes", func() { conf.Server.MusicFolder = libBefore.Path + "-updated" - Expect(repo.StoreMusicFolder()).To(Succeed()) + Expect(repo.StoreMusicFolder(ctx)).To(Succeed()) - libAfter, err := repo.Get(model.DefaultLibraryID) + libAfter, err := repo.Get(ctx, model.DefaultLibraryID) Expect(err).ToNot(HaveOccurred()) Expect(libAfter.Path).To(Equal(conf.Server.MusicFolder)) Expect(libAfter.UpdatedAt).ToNot(Equal(libBefore.UpdatedAt)) @@ -182,10 +182,10 @@ var _ = Describe("LibraryRepository", func() { }) It("refreshes stats", func() { - libBefore, err := repo.Get(1) + libBefore, err := repo.Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) - Expect(repo.RefreshStats(1)).To(Succeed()) - libAfter, err := repo.Get(1) + Expect(repo.RefreshStats(ctx, 1)).To(Succeed()) + libAfter, err := repo.Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) Expect(libAfter.UpdatedAt).To(BeTemporally(">", libBefore.UpdatedAt)) @@ -221,16 +221,16 @@ var _ = Describe("LibraryRepository", func() { Name: "Test Scan Library", Path: "/music/test-scan", } - err := repo.Put(lib) + err := repo.Put(ctx, lib) Expect(err).ToNot(HaveOccurred()) }) DescribeTable("ScanBegin", func(fullScan bool, expectedFullScanInProgress bool) { - err := repo.ScanBegin(lib.ID, fullScan) + err := repo.ScanBegin(ctx, lib.ID, fullScan) Expect(err).ToNot(HaveOccurred()) - updatedLib, err := repo.Get(lib.ID) + updatedLib, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(updatedLib.LastScanStartedAt).ToNot(BeZero()) Expect(updatedLib.FullScanInProgress).To(Equal(expectedFullScanInProgress)) @@ -241,15 +241,15 @@ var _ = Describe("LibraryRepository", func() { Context("ScanEnd", func() { BeforeEach(func() { - err := repo.ScanBegin(lib.ID, true) + err := repo.ScanBegin(ctx, lib.ID, true) Expect(err).ToNot(HaveOccurred()) }) It("sets LastScanAt and clears FullScanInProgress and LastScanStartedAt", func() { - err := repo.ScanEnd(lib.ID) + err := repo.ScanEnd(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) - updatedLib, err := repo.Get(lib.ID) + updatedLib, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(updatedLib.LastScanAt).ToNot(BeZero()) Expect(updatedLib.FullScanInProgress).To(BeFalse()) @@ -257,13 +257,13 @@ var _ = Describe("LibraryRepository", func() { }) It("sets LastScanAt to be after LastScanStartedAt", func() { - libBefore, err := repo.Get(lib.ID) + libBefore, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) - err = repo.ScanEnd(lib.ID) + err = repo.ScanEnd(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) - libAfter, err := repo.Get(lib.ID) + libAfter, err := repo.Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(libAfter.LastScanAt).To(BeTemporally(">=", libBefore.LastScanStartedAt)) }) @@ -273,6 +273,7 @@ var _ = Describe("LibraryRepository", func() { Describe("Delete", func() { var adminRepo model.LibraryRepository var artistRepo model.ArtistRepository + var adminCtx context.Context artistMissing := func(id string) bool { var missing bool @@ -283,32 +284,32 @@ var _ = Describe("LibraryRepository", func() { } BeforeEach(func() { - adminCtx := request.WithUser(log.NewContext(context.TODO()), adminUser) - adminRepo = NewLibraryRepository(adminCtx, conn) - artistRepo = NewArtistRepository(adminCtx, conn) + adminCtx = request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) + adminRepo = NewLibraryRepository(conn) + artistRepo = NewArtistRepository(conn) }) It("marks artists orphaned by the delete as missing", func() { lib := model.Library{Name: "Doomed Library", Path: "/doomed"} - Expect(adminRepo.Put(&lib)).To(Succeed()) + Expect(adminRepo.Put(adminCtx, &lib)).To(Succeed()) orphanArtist := model.Artist{ID: "delete-orphan", Name: "Orphan To Be"} sharedArtist := model.Artist{ID: "delete-shared", Name: "Shared Artist"} - Expect(artistRepo.Put(&orphanArtist)).To(Succeed()) - Expect(artistRepo.Put(&sharedArtist)).To(Succeed()) - Expect(adminRepo.AddArtist(lib.ID, orphanArtist.ID)).To(Succeed()) - Expect(adminRepo.AddArtist(lib.ID, sharedArtist.ID)).To(Succeed()) - Expect(adminRepo.AddArtist(1, sharedArtist.ID)).To(Succeed()) + Expect(artistRepo.Put(adminCtx, &orphanArtist)).To(Succeed()) + Expect(artistRepo.Put(adminCtx, &sharedArtist)).To(Succeed()) + Expect(adminRepo.AddArtist(adminCtx, lib.ID, orphanArtist.ID)).To(Succeed()) + Expect(adminRepo.AddArtist(adminCtx, lib.ID, sharedArtist.ID)).To(Succeed()) + Expect(adminRepo.AddArtist(adminCtx, 1, sharedArtist.ID)).To(Succeed()) DeferCleanup(func() { if raw, ok := artistRepo.(*artistRepository); ok { - _, _ = raw.executeSQL(squirrel.Delete("artist"). + _, _ = raw.executeSQL(adminCtx, squirrel.Delete("artist"). Where(squirrel.Eq{"id": []string{orphanArtist.ID, sharedArtist.ID}})) } }) Expect(artistMissing(orphanArtist.ID)).To(BeFalse()) - Expect(adminRepo.Delete(lib.ID)).To(Succeed()) + Expect(adminRepo.Delete(adminCtx, lib.ID)).To(Succeed()) Expect(artistMissing(orphanArtist.ID)).To(BeTrue(), "orphaned artist should be marked missing") Expect(artistMissing(sharedArtist.ID)).To(BeFalse(), "artist still in another library must stay visible") diff --git a/persistence/mediafile_repository.go b/persistence/mediafile_repository.go index 835c39b1e..b9a5c3374 100644 --- a/persistence/mediafile_repository.go +++ b/persistence/mediafile_repository.go @@ -85,9 +85,8 @@ func (m dbMediaFiles) toModels() model.MediaFiles { return slice.Map(m, func(mf dbMediaFile) model.MediaFile { return *mf.MediaFile }) } -func NewMediaFileRepository(ctx context.Context, db dbx.Builder) model.MediaFileRepository { +func NewMediaFileRepository(db dbx.Builder) model.MediaFileRepository { r := &mediaFileRepository{} - r.ctx = ctx r.db = db r.tableName = "media_file" r.registerModel(&model.MediaFile{}, mediaFileFilter()) @@ -147,25 +146,25 @@ func mediaFileRecentlyAddedSort() string { return "media_file.created_at, media_file.id" } -func (r *mediaFileRepository) CountAll(options ...model.QueryOptions) (int64, error) { - query := r.newSelect() - query = r.applyLibraryFilter(query) +func (r *mediaFileRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + query := r.newSelect(ctx) + query = r.applyLibraryFilter(ctx, query) // The annotation join is expensive with count(distinct) and pointless unless a filter uses it. if filtersNeedAnnotation(r.applyFilters(query, options...)) { - query = r.withAnnotation(query, "media_file.id") + query = r.withAnnotation(ctx, query, "media_file.id") } - return r.count(query, options...) + return r.count(ctx, query, options...) } -func (r *mediaFileRepository) CountBySuffix(options ...model.QueryOptions) (map[string]int64, error) { - sel := r.newSelect(options...). +func (r *mediaFileRepository) CountBySuffix(ctx context.Context, options ...model.QueryOptions) (map[string]int64, error) { + sel := r.newSelect(ctx, options...). Columns("lower(suffix) as suffix", "count(*) as count"). GroupBy("lower(suffix)") var res []struct { Suffix string Count int64 } - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } @@ -176,42 +175,42 @@ func (r *mediaFileRepository) CountBySuffix(options ...model.QueryOptions) (map[ return counts, nil } -func (r *mediaFileRepository) Exists(id string) (bool, error) { +func (r *mediaFileRepository) Exists(ctx context.Context, id string) (bool, error) { // The exists() helper applies no library filter, so it would report rows the caller cannot see. - c, err := r.count(r.applyLibraryFilter(r.newSelect().Where(Eq{"media_file.id": id}))) + c, err := r.count(ctx, r.applyLibraryFilter(ctx, r.newSelect(ctx).Where(Eq{"media_file.id": id}))) return c > 0, err } -func (r *mediaFileRepository) Put(m *model.MediaFile) error { +func (r *mediaFileRepository) Put(ctx context.Context, m *model.MediaFile) error { if m.CreatedAt.IsZero() { m.CreatedAt = time.Now() } - id, err := r.putByMatch(Eq{"path": m.Path, "library_id": m.LibraryID}, m.ID, &dbMediaFile{MediaFile: m}) + id, err := r.putByMatch(ctx, Eq{"path": m.Path, "library_id": m.LibraryID}, m.ID, &dbMediaFile{MediaFile: m}) if err != nil { return err } m.ID = id - if err := r.updateParticipants(m.ID, m.Participants); err != nil { + if err := r.updateParticipants(ctx, m.ID, m.Participants); err != nil { return err } - return r.updateTags(m.ID, m.Tags) + return r.updateTags(ctx, m.ID, m.Tags) } -func (r *mediaFileRepository) UpdateProbeData(id string, data string) error { - _, err := r.executeSQL(Update(r.tableName).Set("probe_data", data).Where(Eq{"id": id})) +func (r *mediaFileRepository) UpdateProbeData(ctx context.Context, id string, data string) error { + _, err := r.executeSQL(ctx, Update(r.tableName).Set("probe_data", data).Where(Eq{"id": id})) return err } -func (r *mediaFileRepository) selectMediaFile(options ...model.QueryOptions) SelectBuilder { - sql := r.newSelect(options...).Columns("media_file.*", "library.path as library_path", "library.name as library_name"). +func (r *mediaFileRepository) selectMediaFile(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + sql := r.newSelect(ctx, options...).Columns("media_file.*", "library.path as library_path", "library.name as library_name"). LeftJoin("library on media_file.library_id = library.id") - sql = r.withAnnotation(sql, "media_file.id") - sql = r.withBookmark(sql, "media_file.id") - return r.applyLibraryFilter(sql) + sql = r.withAnnotation(ctx, sql, "media_file.id") + sql = r.withBookmark(ctx, sql, "media_file.id") + return r.applyLibraryFilter(ctx, sql) } -func (r *mediaFileRepository) Get(id string) (*model.MediaFile, error) { - res, err := r.GetAll(model.QueryOptions{Filters: Eq{"media_file.id": id}}) +func (r *mediaFileRepository) Get(ctx context.Context, id string) (*model.MediaFile, error) { + res, err := r.GetAll(ctx, model.QueryOptions{Filters: Eq{"media_file.id": id}}) if err != nil { return nil, err } @@ -221,34 +220,34 @@ func (r *mediaFileRepository) Get(id string) (*model.MediaFile, error) { return &res[0], nil } -func (r *mediaFileRepository) GetWithParticipants(id string) (*model.MediaFile, error) { - m, err := r.Get(id) +func (r *mediaFileRepository) GetWithParticipants(ctx context.Context, id string) (*model.MediaFile, error) { + m, err := r.Get(ctx, id) if err != nil { return nil, err } - m.Participants, err = r.getParticipants(m) + m.Participants, err = r.getParticipants(ctx, m) return m, err } -func (r *mediaFileRepository) GetAll(options ...model.QueryOptions) (model.MediaFiles, error) { - sq := r.selectMediaFile(options...) +func (r *mediaFileRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { + sq := r.selectMediaFile(ctx, options...) var res dbMediaFiles - err := r.queryAll(sq, &res, options...) + err := r.queryAll(ctx, sq, &res, options...) if err != nil { return nil, err } mfs := res.toModels() - r.hydrateArtwork(mfs) + r.hydrateArtwork(ctx, mfs) return mfs, nil } -func (r *mediaFileRepository) hydrateArtwork(mfs model.MediaFiles) { - hydrateMediaFileArtwork(r.ctx, r.db, mfs) +func (r *mediaFileRepository) hydrateArtwork(ctx context.Context, mfs model.MediaFiles) { + hydrateMediaFileArtwork(ctx, r.db, mfs) } // GetRandom uses two passes so the random sort runs over a narrow rowid index instead of the // wide media_file row: pick random rowids first, then hydrate only those. -func (r *mediaFileRepository) GetRandom(options ...model.QueryOptions) (model.MediaFiles, error) { +func (r *mediaFileRepository) GetRandom(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { var opt model.QueryOptions if len(options) > 0 { opt = options[0] @@ -256,14 +255,14 @@ func (r *mediaFileRepository) GetRandom(options ...model.QueryOptions) (model.Me rowidQuery := Select("media_file.rowid").From(r.tableName) rowidQuery = r.applyFilters(rowidQuery, model.QueryOptions{Filters: opt.Filters}) - rowidQuery = r.applyLibraryFilter(rowidQuery) + rowidQuery = r.applyLibraryFilter(ctx, rowidQuery) rowidQuery = rowidQuery.OrderBy("random()") if opt.Max > 0 { rowidQuery = rowidQuery.Limit(uint64(opt.Max)) } var rowids []int64 - if err := r.queryAllSlice(rowidQuery, &rowids); err != nil { + if err := r.queryAllSlice(ctx, rowidQuery, &rowids); err != nil { return nil, err } if len(rowids) == 0 { @@ -272,17 +271,17 @@ func (r *mediaFileRepository) GetRandom(options ...model.QueryOptions) (model.Me // Re-shuffle in Phase 2: `WHERE rowid IN (...)` returns rows in ascending rowid order, not // the random order from Phase 1. Sorting only the (<=Max) hydrated rows is negligible. - sq := r.selectMediaFile().Where(Eq{"media_file.rowid": rowids}).OrderBy("random()") + sq := r.selectMediaFile(ctx).Where(Eq{"media_file.rowid": rowids}).OrderBy("random()") var res dbMediaFiles - if err := r.queryAll(sq, &res); err != nil { + if err := r.queryAll(ctx, sq, &res); err != nil { return nil, err } mfs := res.toModels() - r.hydrateArtwork(mfs) + r.hydrateArtwork(ctx, mfs) return mfs, nil } -func (r *mediaFileRepository) GetAllByTags(tag model.TagName, values []string, options ...model.QueryOptions) (model.MediaFiles, error) { +func (r *mediaFileRepository) GetAllByTags(ctx context.Context, tag model.TagName, values []string, options ...model.QueryOptions) (model.MediaFiles, error) { placeholders := make([]string, len(values)) args := make([]any, len(values)) for i, v := range values { @@ -304,12 +303,12 @@ func (r *mediaFileRepository) GetAllByTags(tag model.TagName, values []string, o } else { opts.Filters = tagFilter } - return r.GetAll(opts) + return r.GetAll(ctx, opts) } -func (r *mediaFileRepository) GetCursor(options ...model.QueryOptions) (model.MediaFileCursor, error) { - sq := r.selectMediaFile(options...) - cursor, err := queryWithStableResults[dbMediaFile](r.sqlRepository, sq) +func (r *mediaFileRepository) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.MediaFileCursor, error) { + sq := r.selectMediaFile(ctx, options...) + cursor, err := queryWithStableResults[dbMediaFile](ctx, r.sqlRepository, sq) if err != nil { return nil, err } @@ -317,17 +316,17 @@ func (r *mediaFileRepository) GetCursor(options ...model.QueryOptions) (model.Me } // getAllIDs returns the IDs of GetAll's row set, skipping its wide column projection. -func (r *mediaFileRepository) getAllIDs(options ...model.QueryOptions) ([]string, error) { - sq := r.applyLibraryFilter(r.newSelect(options...).Columns("media_file.id")) +func (r *mediaFileRepository) getAllIDs(ctx context.Context, options ...model.QueryOptions) ([]string, error) { + sq := r.applyLibraryFilter(ctx, r.newSelect(ctx, options...).Columns("media_file.id")) if filtersNeedAnnotation(sq) { - sq = r.withAnnotation(sq, "media_file.id") + sq = r.withAnnotation(ctx, sq, "media_file.id") } ids := []string{} - err := r.queryAllSlice(sq, &ids) + err := r.queryAllSlice(ctx, sq, &ids) return ids, err } -func (r *mediaFileRepository) GetAlbumIDsByFolder(lib model.Library, folderIDs ...string) ([]string, error) { +func (r *mediaFileRepository) GetAlbumIDsByFolder(ctx context.Context, lib model.Library, folderIDs ...string) ([]string, error) { ids := []string{} for chunk := range slices.Chunk(folderIDs, 200) { // A folder's own cover also covers albums whose tracks sit in its disc subfolders. @@ -339,7 +338,7 @@ func (r *mediaFileRepository) GetAlbumIDsByFolder(lib model.Library, folderIDs . sq := Select("distinct album_id").From("media_file"). Where(And{Eq{"missing": false}, ConcatExpr("folder_id IN (", inFolders, ")")}) var chunkIDs []string - if err := r.queryAllSlice(sq, &chunkIDs); err != nil { + if err := r.queryAllSlice(ctx, sq, &chunkIDs); err != nil { return nil, err } ids = append(ids, chunkIDs...) @@ -348,14 +347,14 @@ func (r *mediaFileRepository) GetAlbumIDsByFolder(lib model.Library, folderIDs . } // GetCursorWithArtwork streams the same rows as GetCursor, hydrated, via an id pre-pass. -func (r *mediaFileRepository) GetCursorWithArtwork(options ...model.QueryOptions) (model.MediaFileCursor, error) { - ids, err := r.getAllIDs(options...) +func (r *mediaFileRepository) GetCursorWithArtwork(ctx context.Context, options ...model.QueryOptions) (model.MediaFileCursor, error) { + ids, err := r.getAllIDs(ctx, options...) if err != nil { return nil, err } opts := chunkOptions(options, "media_file.id") return model.MediaFileCursor(streamByIDs(ids, func(chunk []string) (model.MediaFiles, error) { - return r.GetAll(opts(chunk)) + return r.GetAll(ctx, opts(chunk)) })), nil } @@ -363,7 +362,7 @@ func (r *mediaFileRepository) GetCursorWithArtwork(options ...model.QueryOptions // The paths can be library-qualified (format: "libraryID:path") or unqualified ("path"). // Library-qualified paths search within the specified library, while unqualified paths // search across all libraries for backward compatibility. -func (r *mediaFileRepository) FindByPaths(paths []string) (model.MediaFiles, error) { +func (r *mediaFileRepository) FindByPaths(ctx context.Context, paths []string) (model.MediaFiles, error) { // One IN list per library instead of one OR term per path: SQLite abandons the // path index at just two OR-ed equality terms and scans the whole table. byLibrary := map[int][]string{} @@ -395,57 +394,57 @@ func (r *mediaFileRepository) FindByPaths(paths []string) (model.MediaFiles, err return model.MediaFiles{}, nil } - sel := r.applyLibraryFilter(r.newSelect().Columns("*").Where(query)) + sel := r.applyLibraryFilter(ctx, r.newSelect(ctx).Columns("*").Where(query)) var res dbMediaFiles - if err := r.queryAll(sel, &res); err != nil { + if err := r.queryAll(ctx, sel, &res); err != nil { return nil, err } return res.toModels(), nil } -func (r *mediaFileRepository) Delete(id string) error { - return r.delete(Eq{"id": id}) +func (r *mediaFileRepository) Delete(ctx context.Context, id string) error { + return r.delete(ctx, Eq{"id": id}) } -func (r *mediaFileRepository) ReassignReferences(prevID, newID string) error { - if err := r.ReassignAnnotation(prevID, newID); err != nil { +func (r *mediaFileRepository) ReassignReferences(ctx context.Context, prevID, newID string) error { + if err := r.ReassignAnnotation(ctx, prevID, newID); err != nil { return fmt.Errorf("reassigning annotations: %w", err) } - if err := r.reassignBookmark(prevID, newID); err != nil { + if err := r.reassignBookmark(ctx, prevID, newID); err != nil { return fmt.Errorf("reassigning bookmarks: %w", err) } upd := Update("playlist_tracks").Set("media_file_id", newID).Where(Eq{"media_file_id": prevID}) - if _, err := r.executeSQL(upd); err != nil { + if _, err := r.executeSQL(ctx, upd); err != nil { return fmt.Errorf("reassigning playlist tracks: %w", err) } upd = Update("scrobbles").Set("media_file_id", newID).Where(Eq{"media_file_id": prevID}) - if _, err := r.executeSQL(upd); err != nil { + if _, err := r.executeSQL(ctx, upd); err != nil { return fmt.Errorf("reassigning scrobbles: %w", err) } // OR IGNORE: scrobble_buffer is unique on (user_id, service, media_file_id, play_time) buf := Expr("update or ignore scrobble_buffer set media_file_id = ? where media_file_id = ?", newID, prevID) - if _, err := r.executeSQL(buf); err != nil { + if _, err := r.executeSQL(ctx, buf); err != nil { return fmt.Errorf("reassigning buffered scrobbles: %w", err) } return nil } -func (r *mediaFileRepository) DeleteAllMissing() (int64, error) { - user := loggedUser(r.ctx) +func (r *mediaFileRepository) DeleteAllMissing(ctx context.Context) (int64, error) { + user := loggedUser(ctx) if !user.IsAdmin { return 0, rest.ErrPermissionDenied } del := Delete(r.tableName).Where(Eq{"missing": true}) - return r.executeSQL(del) + return r.executeSQL(ctx, del) } -func (r *mediaFileRepository) DeleteMissing(ids []string) error { - user := loggedUser(r.ctx) +func (r *mediaFileRepository) DeleteMissing(ctx context.Context, ids []string) error { + user := loggedUser(ctx) if !user.IsAdmin { return rest.ErrPermissionDenied } - return r.delete( + return r.delete(ctx, And{ Eq{"missing": true}, Eq{"id": ids}, @@ -453,24 +452,24 @@ func (r *mediaFileRepository) DeleteMissing(ids []string) error { ) } -func (r *mediaFileRepository) MarkMissing(missing bool, mfs ...*model.MediaFile) error { +func (r *mediaFileRepository) MarkMissing(ctx context.Context, missing bool, mfs ...*model.MediaFile) error { ids := slice.SeqFunc(mfs, func(m *model.MediaFile) string { return m.ID }) for chunk := range slice.CollectChunks(ids, 200) { upd := Update(r.tableName). Set("missing", missing). Set("updated_at", time.Now()). Where(Eq{"id": chunk}) - c, err := r.executeSQL(upd) + c, err := r.executeSQL(ctx, upd) if err != nil || c == 0 { - log.Error(r.ctx, "Error setting mediafile missing flag", "ids", chunk, err) + log.Error(ctx, "Error setting mediafile missing flag", "ids", chunk, err) return err } - log.Debug(r.ctx, "Marked missing mediafiles", "total", c, "ids", chunk) + log.Debug(ctx, "Marked missing mediafiles", "total", c, "ids", chunk) } return nil } -func (r *mediaFileRepository) MarkMissingByFolder(missing bool, folderIDs ...string) error { +func (r *mediaFileRepository) MarkMissingByFolder(ctx context.Context, missing bool, folderIDs ...string) error { for chunk := range slices.Chunk(folderIDs, 200) { upd := Update(r.tableName). Set("missing", missing). @@ -479,12 +478,12 @@ func (r *mediaFileRepository) MarkMissingByFolder(missing bool, folderIDs ...str Eq{"folder_id": chunk}, Eq{"missing": !missing}, }) - c, err := r.executeSQL(upd) + c, err := r.executeSQL(ctx, upd) if err != nil { - log.Error(r.ctx, "Error setting mediafile missing flag", "folderIDs", chunk, err) + log.Error(ctx, "Error setting mediafile missing flag", "folderIDs", chunk, err) return err } - log.Debug(r.ctx, "Marked missing mediafiles from missing folders", "total", c, "folders", chunk) + log.Debug(ctx, "Marked missing mediafiles from missing folders", "total", c, "folders", chunk) } return nil } @@ -492,8 +491,8 @@ func (r *mediaFileRepository) MarkMissingByFolder(missing bool, folderIDs ...str // GetMissingAndMatching returns all mediafiles that are missing and their potential matches (comparing PIDs) // that were added/updated after the last scan started. The result is ordered by PID. // It does not need to load bookmarks, annotations and participants, as they are not used by the scanner. -func (r *mediaFileRepository) GetMissingAndMatching(libId int) (model.MediaFileCursor, error) { - subQ := r.newSelect().Columns("pid"). +func (r *mediaFileRepository) GetMissingAndMatching(ctx context.Context, libId int) (model.MediaFileCursor, error) { + subQ := r.newSelect(ctx).Columns("pid"). Where(And{ Eq{"media_file.missing": true}, Eq{"library_id": libId}, @@ -502,7 +501,7 @@ func (r *mediaFileRepository) GetMissingAndMatching(libId int) (model.MediaFileC if err != nil { return nil, err } - sel := r.newSelect().Columns("media_file.*", "library.path as library_path", "library.name as library_name"). + sel := r.newSelect(ctx).Columns("media_file.*", "library.path as library_path", "library.name as library_name"). LeftJoin("library on media_file.library_id = library.id"). Where("pid in ("+subQText+")", subQArgs...). Where(Or{ @@ -510,7 +509,7 @@ func (r *mediaFileRepository) GetMissingAndMatching(libId int) (model.MediaFileC ConcatExpr("media_file.created_at > library.last_scan_started_at"), }). OrderBy("pid") - cursor, err := queryWithStableResults[dbMediaFile](r.sqlRepository, sel) + cursor, err := queryWithStableResults[dbMediaFile](ctx, r.sqlRepository, sel) if err != nil { return nil, err } @@ -523,8 +522,8 @@ func wrapMediaFileCursor(cursor iter.Seq2[dbMediaFile, error]) model.MediaFileCu // FindRecentFilesByMBZTrackID finds recently added files by MusicBrainz Track ID in other libraries // It uses a lightweight query without annotation/bookmark joins since those are not needed for matching -func (r *mediaFileRepository) FindRecentFilesByMBZTrackID(missing model.MediaFile, since time.Time) (model.MediaFiles, error) { - sel := r.newSelect().Columns("media_file.*", "library.path as library_path", "library.name as library_name"). +func (r *mediaFileRepository) FindRecentFilesByMBZTrackID(ctx context.Context, missing model.MediaFile, since time.Time) (model.MediaFiles, error) { + sel := r.newSelect(ctx).Columns("media_file.*", "library.path as library_path", "library.name as library_name"). LeftJoin("library on media_file.library_id = library.id"). Where(And{ NotEq{"media_file.library_id": missing.LibraryID}, @@ -536,7 +535,7 @@ func (r *mediaFileRepository) FindRecentFilesByMBZTrackID(missing model.MediaFil }).OrderBy("media_file.created_at DESC") var res dbMediaFiles - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } @@ -545,8 +544,8 @@ func (r *mediaFileRepository) FindRecentFilesByMBZTrackID(missing model.MediaFil // FindRecentFilesByProperties finds recently added files by intrinsic properties in other libraries // It uses a lightweight query without annotation/bookmark joins since those are not needed for matching -func (r *mediaFileRepository) FindRecentFilesByProperties(missing model.MediaFile, since time.Time) (model.MediaFiles, error) { - sel := r.newSelect().Columns("media_file.*", "library.path as library_path", "library.name as library_name"). +func (r *mediaFileRepository) FindRecentFilesByProperties(ctx context.Context, missing model.MediaFile, since time.Time) (model.MediaFiles, error) { + sel := r.newSelect(ctx).Columns("media_file.*", "library.path as library_path", "library.name as library_name"). LeftJoin("library on media_file.library_id = library.id"). Where(And{ NotEq{"media_file.library_id": missing.LibraryID}, @@ -562,7 +561,7 @@ func (r *mediaFileRepository) FindRecentFilesByProperties(missing model.MediaFil }).OrderBy("media_file.created_at DESC") var res dbMediaFiles - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } @@ -575,8 +574,8 @@ var mediaFileSearchConfig = searchConfig{ MBIDFields: []string{"mbz_recording_id", "mbz_release_track_id"}, } -func (r *mediaFileRepository) MatchesCriteria(id string, c criteria.Criteria) (bool, error) { - usr := loggedUser(r.ctx) +func (r *mediaFileRepository) MatchesCriteria(ctx context.Context, id string, c criteria.Criteria) (bool, error) { + usr := loggedUser(ctx) rulesSQL := newSmartPlaylistCriteria(c, withSmartPlaylistOwner(*usr)) cond, err := rulesSQL.where() if err != nil { @@ -586,46 +585,38 @@ func (r *mediaFileRepository) MatchesCriteria(id string, c criteria.Criteria) (b sq = rulesSQL.applyExpressionJoins(sq, usr.ID) sq = sq.Where(And{Eq{"media_file.id": id}, cond}) var res struct{ Count int64 } - if err := r.queryOne(sq, &res); err != nil { + if err := r.queryOne(ctx, sq, &res); err != nil { return false, err } return res.Count > 0, nil } -func (r *mediaFileRepository) Search(q string, options ...model.QueryOptions) (model.MediaFiles, error) { +func (r *mediaFileRepository) Search(ctx context.Context, q string, options ...model.QueryOptions) (model.MediaFiles, error) { var opts model.QueryOptions if len(options) > 0 { opts = options[0] } var res dbMediaFiles - err := r.doSearch(r.selectMediaFile(options...), q, &res, mediaFileSearchConfig, opts) + err := r.doSearch(ctx, r.selectMediaFile(ctx, options...), q, &res, mediaFileSearchConfig, opts) if err != nil { return nil, fmt.Errorf("searching media_file %q: %w", q, err) } mfs := res.toModels() - r.hydrateArtwork(mfs) + r.hydrateArtwork(ctx, mfs) return mfs, nil } -func (r *mediaFileRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *mediaFileRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *mediaFileRepository) Read(id string) (any, error) { - return r.Get(id) +func (r *mediaFileRepository) Read(ctx context.Context, id string) (*model.MediaFile, error) { + return r.Get(ctx, id) } -func (r *mediaFileRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) -} - -func (r *mediaFileRepository) EntityName() string { - return "mediafile" -} - -func (r *mediaFileRepository) NewInstance() any { - return &model.MediaFile{} +func (r *mediaFileRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.MediaFile, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } var _ model.MediaFileRepository = (*mediaFileRepository)(nil) -var _ model.ResourceRepository = (*mediaFileRepository)(nil) +var _ rest.Repository[model.MediaFile] = (*mediaFileRepository)(nil) diff --git a/persistence/mediafile_repository_test.go b/persistence/mediafile_repository_test.go index d81b5e9ee..14590090c 100644 --- a/persistence/mediafile_repository_test.go +++ b/persistence/mediafile_repository_test.go @@ -24,11 +24,11 @@ import ( var _ = Describe("MediaRepository", func() { var mr model.MediaFileRepository + var ctx context.Context BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, model.User{ID: "userid"}) - mr = NewMediaFileRepository(ctx, GetDBXBuilder()) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid"}) + mr = NewMediaFileRepository(GetDBXBuilder()) }) Describe("GetAlbumIDsByFolder", func() { @@ -37,22 +37,22 @@ var _ = Describe("MediaRepository", func() { BeforeEach(func() { ctx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) - libPtr, err := NewLibraryRepository(ctx, GetDBXBuilder()).Get(1) + libPtr, err := NewLibraryRepository(GetDBXBuilder()).Get(ctx, 1) Expect(err).ToNot(HaveOccurred()) lib = *libPtr - folderRepo := newFolderRepository(ctx, GetDBXBuilder()) + folderRepo := newFolderRepository(GetDBXBuilder()) albumRoot = model.NewFolder(lib, "ByFolder/Album") disc1 = model.NewFolder(lib, "ByFolder/Album/CD1") sibling = model.NewFolder(lib, "ByFolder/Other") for _, f := range []*model.Folder{albumRoot, disc1, sibling} { - Expect(folderRepo.Put(f)).To(Succeed()) + Expect(folderRepo.Put(ctx, f)).To(Succeed()) } // Tracks live in the disc subfolder; the sibling album is the negative control. - Expect(mr.Put(&model.MediaFile{ID: "fol-mf-1", LibraryID: 1, AlbumID: "fol-al-1", FolderID: disc1.ID, Path: "t/1.mp3"})).To(Succeed()) - Expect(mr.Put(&model.MediaFile{ID: "fol-mf-2", LibraryID: 1, AlbumID: "fol-al-1", FolderID: disc1.ID, Path: "t/2.mp3"})).To(Succeed()) - Expect(mr.Put(&model.MediaFile{ID: "fol-mf-3", LibraryID: 1, AlbumID: "fol-al-2", FolderID: sibling.ID, Path: "t/3.mp3"})).To(Succeed()) - Expect(mr.Put(&model.MediaFile{ID: "fol-mf-4", LibraryID: 1, AlbumID: "fol-al-3", FolderID: disc1.ID, Path: "t/4.mp3", Missing: true})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{ID: "fol-mf-1", LibraryID: 1, AlbumID: "fol-al-1", FolderID: disc1.ID, Path: "t/1.mp3"})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{ID: "fol-mf-2", LibraryID: 1, AlbumID: "fol-al-1", FolderID: disc1.ID, Path: "t/2.mp3"})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{ID: "fol-mf-3", LibraryID: 1, AlbumID: "fol-al-2", FolderID: sibling.ID, Path: "t/3.mp3"})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{ID: "fol-mf-4", LibraryID: 1, AlbumID: "fol-al-3", FolderID: disc1.ID, Path: "t/4.mp3", Missing: true})).To(Succeed()) DeferCleanup(func() { _, _ = GetDBXBuilder().NewQuery("DELETE FROM media_file WHERE id LIKE 'fol-mf-%'").Execute() _, _ = GetDBXBuilder().NewQuery("DELETE FROM folder WHERE path LIKE 'ByFolder%' OR name = 'ByFolder'").Execute() @@ -60,20 +60,20 @@ var _ = Describe("MediaRepository", func() { }) It("returns the distinct album IDs of non-missing tracks in the folder", func() { - ids, err := mr.GetAlbumIDsByFolder(lib, disc1.ID) + ids, err := mr.GetAlbumIDsByFolder(ctx, lib, disc1.ID) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(ConsistOf("fol-al-1")) }) It("also matches albums whose tracks are in a direct child of the folder", func() { // A cover in the album root must reach the album whose tracks sit in CD1 - ids, err := mr.GetAlbumIDsByFolder(lib, albumRoot.ID) + ids, err := mr.GetAlbumIDsByFolder(ctx, lib, albumRoot.ID) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(ConsistOf("fol-al-1")) }) It("does not match albums outside the folder", func() { - ids, err := mr.GetAlbumIDsByFolder(lib, albumRoot.ID) + ids, err := mr.GetAlbumIDsByFolder(ctx, lib, albumRoot.ID) Expect(err).ToNot(HaveOccurred()) Expect(ids).ToNot(ContainElement("fol-al-2")) }) @@ -82,46 +82,47 @@ var _ = Describe("MediaRepository", func() { Describe("GetCursor", func() { It("yields the same media files as GetAll", func() { opts := model.QueryOptions{Sort: "title"} - want, err := mr.GetAll(opts) + want, err := mr.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(mr.GetCursor(opts))).To(Equal([]model.MediaFile(want))) + Expect(collectCursor(mr.GetCursor(ctx, opts))).To(Equal([]model.MediaFile(want))) }) It("honors Max/Offset like GetAll", func() { opts := model.QueryOptions{Sort: "title", Max: 2, Offset: 1} - want, err := mr.GetAll(opts) + want, err := mr.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(mr.GetCursor(opts))).To(Equal([]model.MediaFile(want))) + Expect(collectCursor(mr.GetCursor(ctx, opts))).To(Equal([]model.MediaFile(want))) }) }) It("gets mediafile from the DB", func() { - actual, err := mr.Get("1004") + actual, err := mr.Get(ctx, "1004") Expect(err).ToNot(HaveOccurred()) actual.CreatedAt = time.Time{} Expect(actual).To(Equal(&songAntenna)) }) It("returns ErrNotFound", func() { - _, err := mr.Get("56") + _, err := mr.Get(ctx, "56") Expect(err).To(MatchError(model.ErrNotFound)) }) It("counts the number of mediafiles in the DB", func() { - Expect(mr.CountAll()).To(Equal(int64(13))) + Expect(mr.CountAll(ctx)).To(Equal(int64(13))) }) Describe("CountAll annotation-join gating", func() { var adminRepo model.MediaFileRepository + var adminCtx context.Context BeforeEach(func() { - adminCtx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid", IsAdmin: true}) - adminRepo = NewMediaFileRepository(adminCtx, GetDBXBuilder()) + adminCtx = request.WithUser(ctx, model.User{ID: "userid", IsAdmin: true}) + adminRepo = NewMediaFileRepository(GetDBXBuilder()) }) It("counts starred songs when an annotation filter is present", func() { // Come Together (id 1002) is starred for the admin user in the seed data - count, err := adminRepo.CountAll(model.QueryOptions{ + count, err := adminRepo.CountAll(adminCtx, model.QueryOptions{ Filters: annotationBoolFilter("starred")("starred", "true"), }) Expect(err).ToNot(HaveOccurred()) @@ -129,7 +130,7 @@ var _ = Describe("MediaRepository", func() { }) It("counts with starred=false without a 'no such column' error (join kept)", func() { - count, err := adminRepo.CountAll(model.QueryOptions{ + count, err := adminRepo.CountAll(adminCtx, model.QueryOptions{ Filters: annotationBoolFilter("starred")("starred", "false"), }) Expect(err).ToNot(HaveOccurred()) @@ -138,7 +139,7 @@ var _ = Describe("MediaRepository", func() { }) It("counts unfiltered with the join dropped", func() { - Expect(adminRepo.CountAll()).To(Equal(int64(13))) + Expect(adminRepo.CountAll(adminCtx)).To(Equal(int64(13))) }) }) @@ -151,21 +152,21 @@ var _ = Describe("MediaRepository", func() { flacFile2 = model.MediaFile{ID: "suffix-flac2", LibraryID: 1, Suffix: "flac", Path: "test/file2.flac"} flacUpperFile = model.MediaFile{ID: "suffix-FLAC", LibraryID: 1, Suffix: "FLAC", Path: "test/file.FLAC"} - Expect(mr.Put(&mp3File)).To(Succeed()) - Expect(mr.Put(&flacFile1)).To(Succeed()) - Expect(mr.Put(&flacFile2)).To(Succeed()) - Expect(mr.Put(&flacUpperFile)).To(Succeed()) + Expect(mr.Put(ctx, &mp3File)).To(Succeed()) + Expect(mr.Put(ctx, &flacFile1)).To(Succeed()) + Expect(mr.Put(ctx, &flacFile2)).To(Succeed()) + Expect(mr.Put(ctx, &flacUpperFile)).To(Succeed()) }) AfterEach(func() { - _ = mr.Delete(mp3File.ID) - _ = mr.Delete(flacFile1.ID) - _ = mr.Delete(flacFile2.ID) - _ = mr.Delete(flacUpperFile.ID) + _ = mr.Delete(ctx, mp3File.ID) + _ = mr.Delete(ctx, flacFile1.ID) + _ = mr.Delete(ctx, flacFile2.ID) + _ = mr.Delete(ctx, flacUpperFile.ID) }) It("counts media files grouped by suffix with lowercase normalization", func() { - counts, err := mr.CountBySuffix() + counts, err := mr.CountBySuffix(ctx) Expect(err).ToNot(HaveOccurred()) // Should have lowercase keys only @@ -182,7 +183,7 @@ var _ = Describe("MediaRepository", func() { It("returns songs ordered by lyrics with a specific title/artist", func() { // attempt to mimic filters.SongsByArtistTitleWithLyricsFirst, except we want all items - results, err := mr.GetAll(model.QueryOptions{ + results, err := mr.GetAll(ctx, model.QueryOptions{ Sort: "lyrics, updated_at", Order: "desc", Filters: squirrel.And{ @@ -206,14 +207,14 @@ var _ = Describe("MediaRepository", func() { Describe("GetRandom", func() { It("returns the requested number of distinct, fully-hydrated media files", func() { - results, err := mr.GetRandom(model.QueryOptions{Max: 5}) + results, err := mr.GetRandom(ctx, model.QueryOptions{Max: 5}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(5)) // Each returned row must match its GetAll counterpart exactly — proves Phase 2 // hydrates full rows (not bare rowids) — and ids must be distinct. byID := map[string]model.MediaFile{} - all, err := mr.GetAll() + all, err := mr.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) for _, mf := range all { byID[mf.ID] = mf @@ -229,13 +230,13 @@ var _ = Describe("MediaRepository", func() { }) It("returns all matching files when Max exceeds the total", func() { - results, err := mr.GetRandom(model.QueryOptions{Max: 1000}) + results, err := mr.GetRandom(ctx, model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(13)) }) It("honors filters", func() { - results, err := mr.GetRandom(model.QueryOptions{ + results, err := mr.GetRandom(ctx, model.QueryOptions{ Max: 10, Filters: squirrel.Eq{"media_file.title": "Antenna"}, }) @@ -248,7 +249,7 @@ var _ = Describe("MediaRepository", func() { It("returns varying results across calls", func() { // Retry a few times: two random draws of 5 from 13 rows differ with near-certainty. - first, err := mr.GetRandom(model.QueryOptions{Max: 5}) + first, err := mr.GetRandom(ctx, model.QueryOptions{Max: 5}) Expect(err).ToNot(HaveOccurred()) firstIDs := func() []string { ids := make([]string, len(first)) @@ -259,7 +260,7 @@ var _ = Describe("MediaRepository", func() { }() differed := false for range 10 { - next, err := mr.GetRandom(model.QueryOptions{Max: 5}) + next, err := mr.GetRandom(ctx, model.QueryOptions{Max: 5}) Expect(err).ToNot(HaveOccurred()) nextIDs := make([]string, len(next)) for i, mf := range next { @@ -276,7 +277,7 @@ var _ = Describe("MediaRepository", func() { It("randomizes order even when Max exceeds the total", func() { // Same set of rows every time (all 13), but the order must still be shuffled — // guards against Phase 2's `rowid IN (...)` returning rows in rowid order. - first, err := mr.GetRandom(model.QueryOptions{Max: 100}) + first, err := mr.GetRandom(ctx, model.QueryOptions{Max: 100}) Expect(err).ToNot(HaveOccurred()) Expect(first).To(HaveLen(13)) firstIDs := make([]string, len(first)) @@ -285,7 +286,7 @@ var _ = Describe("MediaRepository", func() { } differed := false for range 10 { - next, err := mr.GetRandom(model.QueryOptions{Max: 100}) + next, err := mr.GetRandom(ctx, model.QueryOptions{Max: 100}) Expect(err).ToNot(HaveOccurred()) nextIDs := make([]string, len(next)) for i, mf := range next { @@ -304,13 +305,13 @@ var _ = Describe("MediaRepository", func() { It("sets CreatedAt to now when inserting a new file with zero CreatedAt", func() { before := time.Now().Add(-time.Second) newFile := model.MediaFile{ID: id.NewRandom(), LibraryID: 1, Path: "test/created-at-zero.mp3"} - Expect(mr.Put(&newFile)).To(Succeed()) + Expect(mr.Put(ctx, &newFile)).To(Succeed()) - retrieved, err := mr.Get(newFile.ID) + retrieved, err := mr.Get(ctx, newFile.ID) Expect(err).ToNot(HaveOccurred()) Expect(retrieved.CreatedAt).To(BeTemporally(">", before)) - _ = mr.Delete(newFile.ID) + _ = mr.Delete(ctx, newFile.ID) }) It("preserves CreatedAt when inserting a new file with non-zero CreatedAt", func() { @@ -321,13 +322,13 @@ var _ = Describe("MediaRepository", func() { Path: "test/created-at-preserved.mp3", CreatedAt: originalTime, } - Expect(mr.Put(&newFile)).To(Succeed()) + Expect(mr.Put(ctx, &newFile)).To(Succeed()) - retrieved, err := mr.Get(newFile.ID) + retrieved, err := mr.Get(ctx, newFile.ID) Expect(err).ToNot(HaveOccurred()) Expect(retrieved.CreatedAt).To(BeTemporally("~", originalTime, time.Second)) - _ = mr.Delete(newFile.ID) + _ = mr.Delete(ctx, newFile.ID) }) It("does not reset CreatedAt when updating an existing file", func() { @@ -340,7 +341,7 @@ var _ = Describe("MediaRepository", func() { Title: "Original Title", CreatedAt: originalTime, } - Expect(mr.Put(&newFile)).To(Succeed()) + Expect(mr.Put(ctx, &newFile)).To(Succeed()) // Update the file with a new title but zero CreatedAt updatedFile := model.MediaFile{ @@ -350,66 +351,66 @@ var _ = Describe("MediaRepository", func() { Title: "Updated Title", // CreatedAt is zero - should NOT overwrite the stored value } - Expect(mr.Put(&updatedFile)).To(Succeed()) + Expect(mr.Put(ctx, &updatedFile)).To(Succeed()) - retrieved, err := mr.Get(fileID) + retrieved, err := mr.Get(ctx, fileID) Expect(err).ToNot(HaveOccurred()) Expect(retrieved.Title).To(Equal("Updated Title")) // CreatedAt should still be the original time (not reset) Expect(retrieved.CreatedAt).To(BeTemporally("~", originalTime, time.Second)) - _ = mr.Delete(fileID) + _ = mr.Delete(ctx, fileID) }) }) It("checks existence of mediafiles in the DB", func() { - Expect(mr.Exists(songAntenna.ID)).To(BeTrue()) - Expect(mr.Exists("666")).To(BeFalse()) + Expect(mr.Exists(ctx, songAntenna.ID)).To(BeTrue()) + Expect(mr.Exists(ctx, "666")).To(BeFalse()) }) It("delete tracks by id", func() { newID := id.NewRandom() - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: newID})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: newID})).To(Succeed()) - Expect(mr.Delete(newID)).To(Succeed()) + Expect(mr.Delete(ctx, newID)).To(Succeed()) - _, err := mr.Get(newID) + _, err := mr.Get(ctx, newID) Expect(err).To(MatchError(model.ErrNotFound)) }) It("deletes all missing files", func() { new1 := model.MediaFile{ID: id.NewRandom(), LibraryID: 1} new2 := model.MediaFile{ID: id.NewRandom(), LibraryID: 1} - Expect(mr.Put(&new1)).To(Succeed()) - Expect(mr.Put(&new2)).To(Succeed()) - Expect(mr.MarkMissing(true, &new1, &new2)).To(Succeed()) + Expect(mr.Put(ctx, &new1)).To(Succeed()) + Expect(mr.Put(ctx, &new2)).To(Succeed()) + Expect(mr.MarkMissing(ctx, true, &new1, &new2)).To(Succeed()) adminCtx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid", IsAdmin: true}) - adminRepo := NewMediaFileRepository(adminCtx, GetDBXBuilder()) + adminRepo := NewMediaFileRepository(GetDBXBuilder()) // Ensure the files are marked as missing and we have 2 of them - count, err := adminRepo.CountAll(model.QueryOptions{Filters: squirrel.Eq{"missing": true}}) + count, err := adminRepo.CountAll(adminCtx, model.QueryOptions{Filters: squirrel.Eq{"missing": true}}) Expect(count).To(BeNumerically("==", 2)) Expect(err).ToNot(HaveOccurred()) - count, err = adminRepo.DeleteAllMissing() + count, err = adminRepo.DeleteAllMissing(adminCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeNumerically("==", 2)) - _, err = mr.Get(new1.ID) + _, err = mr.Get(ctx, new1.ID) Expect(err).To(MatchError(model.ErrNotFound)) - _, err = mr.Get(new2.ID) + _, err = mr.Get(ctx, new2.ID) Expect(err).To(MatchError(model.ErrNotFound)) }) Context("Annotations", func() { It("increments play count when the tracks does not have annotations", func() { id := "incplay.firsttime" - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: id})).To(BeNil()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: id})).To(BeNil()) playDate := time.Now() - Expect(mr.IncPlayCount(id, playDate)).To(BeNil()) + Expect(mr.IncPlayCount(ctx, id, playDate)).To(BeNil()) - mf, err := mr.Get(id) + mf, err := mr.Get(ctx, id) Expect(err).To(BeNil()) Expect(mf.PlayDate.Unix()).To(Equal(playDate.Unix())) @@ -425,85 +426,85 @@ var _ = Describe("MediaRepository", func() { It("returns 0 when no ratings exist", func() { newID := id.NewRandom() - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: newID, Path: "test/no-rating.mp3"})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: newID, Path: "test/no-rating.mp3"})).To(Succeed()) - mf, err := mr.Get(newID) + mf, err := mr.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(mf.AverageRating).To(Equal(0.0)) - _, _ = raw.executeSQL(squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) }) It("returns the user's rating as average when only one user rated", func() { newID := id.NewRandom() - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: newID, Path: "test/single-rating.mp3"})).To(Succeed()) - Expect(mr.SetRating(5, newID)).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: newID, Path: "test/single-rating.mp3"})).To(Succeed()) + Expect(mr.SetRating(ctx, 5, newID)).To(Succeed()) - mf, err := mr.Get(newID) + mf, err := mr.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(mf.AverageRating).To(Equal(5.0)) - _, _ = raw.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = raw.executeSQL(squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) }) It("calculates average across multiple users", func() { newID := id.NewRandom() - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: newID, Path: "test/multi-rating.mp3"})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: newID, Path: "test/multi-rating.mp3"})).To(Succeed()) - Expect(mr.SetRating(3, newID)).To(Succeed()) + Expect(mr.SetRating(ctx, 3, newID)).To(Succeed()) user2Ctx := request.WithUser(GinkgoT().Context(), regularUser) - user2Repo := NewMediaFileRepository(user2Ctx, GetDBXBuilder()) - Expect(user2Repo.SetRating(5, newID)).To(Succeed()) + user2Repo := NewMediaFileRepository(GetDBXBuilder()) + Expect(user2Repo.SetRating(user2Ctx, 5, newID)).To(Succeed()) - mf, err := mr.Get(newID) + mf, err := mr.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(mf.AverageRating).To(Equal(4.0)) - _, _ = raw.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = raw.executeSQL(squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) }) It("excludes zero ratings from average calculation", func() { newID := id.NewRandom() - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: newID, Path: "test/zero-excluded.mp3"})).To(Succeed()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: newID, Path: "test/zero-excluded.mp3"})).To(Succeed()) - Expect(mr.SetRating(4, newID)).To(Succeed()) + Expect(mr.SetRating(ctx, 4, newID)).To(Succeed()) user2Ctx := request.WithUser(GinkgoT().Context(), regularUser) - user2Repo := NewMediaFileRepository(user2Ctx, GetDBXBuilder()) - Expect(user2Repo.SetRating(0, newID)).To(Succeed()) + user2Repo := NewMediaFileRepository(GetDBXBuilder()) + Expect(user2Repo.SetRating(user2Ctx, 0, newID)).To(Succeed()) - mf, err := mr.Get(newID) + mf, err := mr.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(mf.AverageRating).To(Equal(4.0)) - _, _ = raw.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) - _, _ = raw.executeSQL(squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": newID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete("media_file").Where(squirrel.Eq{"id": newID})) }) }) It("preserves play date if and only if provided date is older", func() { id := "incplay.playdate" - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: id})).To(BeNil()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: id})).To(BeNil()) playDate := time.Now() - Expect(mr.IncPlayCount(id, playDate)).To(BeNil()) - mf, err := mr.Get(id) + Expect(mr.IncPlayCount(ctx, id, playDate)).To(BeNil()) + mf, err := mr.Get(ctx, id) Expect(err).To(BeNil()) Expect(mf.PlayDate.Unix()).To(Equal(playDate.Unix())) Expect(mf.PlayCount).To(Equal(int64(1))) playDateLate := playDate.AddDate(0, 0, 1) - Expect(mr.IncPlayCount(id, playDateLate)).To(BeNil()) - mf, err = mr.Get(id) + Expect(mr.IncPlayCount(ctx, id, playDateLate)).To(BeNil()) + mf, err = mr.Get(ctx, id) Expect(err).To(BeNil()) Expect(mf.PlayDate.Unix()).To(Equal(playDateLate.Unix())) Expect(mf.PlayCount).To(Equal(int64(2))) playDateEarly := playDate.AddDate(0, 0, -1) - Expect(mr.IncPlayCount(id, playDateEarly)).To(BeNil()) - mf, err = mr.Get(id) + Expect(mr.IncPlayCount(ctx, id, playDateEarly)).To(BeNil()) + mf, err = mr.Get(ctx, id) Expect(err).To(BeNil()) Expect(mf.PlayDate.Unix()).To(Equal(playDateLate.Unix())) Expect(mf.PlayCount).To(Equal(int64(3))) @@ -511,12 +512,12 @@ var _ = Describe("MediaRepository", func() { It("increments play count on newly starred items", func() { id := "star.incplay" - Expect(mr.Put(&model.MediaFile{LibraryID: 1, ID: id})).To(BeNil()) - Expect(mr.SetStar(true, id)).To(BeNil()) + Expect(mr.Put(ctx, &model.MediaFile{LibraryID: 1, ID: id})).To(BeNil()) + Expect(mr.SetStar(ctx, true, id)).To(BeNil()) playDate := time.Now() - Expect(mr.IncPlayCount(id, playDate)).To(BeNil()) + Expect(mr.IncPlayCount(ctx, id, playDate)).To(BeNil()) - mf, err := mr.Get(id) + mf, err := mr.Get(ctx, id) Expect(err).To(BeNil()) Expect(mf.PlayDate.Unix()).To(Equal(playDate.Unix())) @@ -555,7 +556,7 @@ var _ = Describe("MediaRepository", func() { // Insert test data first for i := range testMediaFiles { - Expect(mr.Put(&testMediaFiles[i])).To(Succeed()) + Expect(mr.Put(ctx, &testMediaFiles[i])).To(Succeed()) } // Then manually update timestamps using direct SQL to bypass the repository logic @@ -597,7 +598,7 @@ var _ = Describe("MediaRepository", func() { AfterEach(func() { // Clean up test data for _, mf := range testMediaFiles { - _ = mr.Delete(mf.ID) + _ = mr.Delete(ctx, mf.ID) } }) @@ -607,14 +608,12 @@ var _ = Describe("MediaRepository", func() { BeforeEach(func() { conf.Server.RecentlyAddedByModTime = false // Create repository AFTER setting config - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid"}) - testRepo = NewMediaFileRepository(ctx, GetDBXBuilder()) + testRepo = NewMediaFileRepository(GetDBXBuilder()) }) It("sorts by created_at", func() { // Get results sorted by recently_added (should use created_at) - results, err := testRepo.GetAll(model.QueryOptions{ + results, err := testRepo.GetAll(ctx, model.QueryOptions{ Sort: "recently_added", Order: "desc", Filters: squirrel.Eq{"media_file.id": []string{testMediaFiles[0].ID, testMediaFiles[1].ID, testMediaFiles[2].ID}}, @@ -630,7 +629,7 @@ var _ = Describe("MediaRepository", func() { It("sorts in ascending order when specified", func() { // Get results sorted by recently_added in ascending order - results, err := testRepo.GetAll(model.QueryOptions{ + results, err := testRepo.GetAll(ctx, model.QueryOptions{ Sort: "recently_added", Order: "asc", Filters: squirrel.Eq{"media_file.id": []string{testMediaFiles[0].ID, testMediaFiles[1].ID, testMediaFiles[2].ID}}, @@ -651,14 +650,12 @@ var _ = Describe("MediaRepository", func() { BeforeEach(func() { conf.Server.RecentlyAddedByModTime = true // Create repository AFTER setting config - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid"}) - testRepo = NewMediaFileRepository(ctx, GetDBXBuilder()) + testRepo = NewMediaFileRepository(GetDBXBuilder()) }) It("sorts by updated_at", func() { // Get results sorted by recently_added (should use updated_at) - results, err := testRepo.GetAll(model.QueryOptions{ + results, err := testRepo.GetAll(ctx, model.QueryOptions{ Sort: "recently_added", Order: "desc", Filters: squirrel.Eq{"media_file.id": []string{testMediaFiles[0].ID, testMediaFiles[1].ID, testMediaFiles[2].ID}}, @@ -677,7 +674,7 @@ var _ = Describe("MediaRepository", func() { conf.Server.RecentlyAddedByModTime = false ctx := log.NewContext(GinkgoT().Context()) ctx = request.WithUser(ctx, model.User{ID: "userid"}) - repo := NewMediaFileRepository(ctx, GetDBXBuilder()) + repo := NewMediaFileRepository(GetDBXBuilder()) ids := []string{testMediaFiles[0].ID, testMediaFiles[1].ID, testMediaFiles[2].ID} sameTime := time.Date(2024, 3, 1, 0, 0, 0, 0, time.UTC) @@ -687,7 +684,7 @@ var _ = Describe("MediaRepository", func() { Expect(err).ToNot(HaveOccurred()) order := func() []string { - res, err := repo.GetAll(model.QueryOptions{ + res, err := repo.GetAll(ctx, model.QueryOptions{ Sort: "recently_added", Order: "desc", Filters: squirrel.Eq{"media_file.id": ids}}) Expect(err).ToNot(HaveOccurred()) @@ -709,20 +706,20 @@ var _ = Describe("MediaRepository", func() { BeforeEach(func() { mfWithoutAnnotation = model.MediaFile{ID: "no-annotation-file", LibraryID: 1, Path: "test/no-annotation.mp3", Title: "No Annotation"} - Expect(mr.Put(&mfWithoutAnnotation)).To(Succeed()) + Expect(mr.Put(ctx, &mfWithoutAnnotation)).To(Succeed()) }) AfterEach(func() { - _ = mr.Delete(mfWithoutAnnotation.ID) + _ = mr.Delete(ctx, mfWithoutAnnotation.ID) }) Describe("starred", func() { It("false includes items without annotations", func() { - res, err := mr.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := mr.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "false"}, }) Expect(err).ToNot(HaveOccurred()) - files := res.(model.MediaFiles) + files := res var found bool for _, f := range files { @@ -735,11 +732,11 @@ var _ = Describe("MediaRepository", func() { }) It("true excludes items without annotations", func() { - res, err := mr.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := mr.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "true"}, }) Expect(err).ToNot(HaveOccurred()) - files := res.(model.MediaFiles) + files := res for _, f := range files { Expect(f.ID).ToNot(Equal(mfWithoutAnnotation.ID)) @@ -749,11 +746,11 @@ var _ = Describe("MediaRepository", func() { Describe("path", func() { It("matches files whose path starts with the given prefix", func() { - res, err := mr.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := mr.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"path": "test/"}, }) Expect(err).ToNot(HaveOccurred()) - files := res.(model.MediaFiles) + files := res var found bool for _, f := range files { @@ -766,11 +763,11 @@ var _ = Describe("MediaRepository", func() { }) It("excludes files whose path does not start with the given prefix", func() { - res, err := mr.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := mr.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"path": "no-such-prefix/"}, }) Expect(err).ToNot(HaveOccurred()) - files := res.(model.MediaFiles) + files := res Expect(files).To(BeEmpty()) }) }) @@ -779,7 +776,7 @@ var _ = Describe("MediaRepository", func() { Describe("Search", func() { Context("text search", func() { It("finds media files by title", func() { - results, err := mr.Search("Antenna", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "Antenna", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(3)) // songAntenna, songAntennaWithLyrics, songAntenna2 for _, result := range results { @@ -788,7 +785,7 @@ var _ = Describe("MediaRepository", func() { }) It("finds media files case insensitively", func() { - results, err := mr.Search("antenna", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "antenna", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(3)) for _, result := range results { @@ -797,7 +794,7 @@ var _ = Describe("MediaRepository", func() { }) It("returns empty result when no matches found", func() { - results, err := mr.Search("nonexistent", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "nonexistent", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -820,17 +817,17 @@ var _ = Describe("MediaRepository", func() { } // Insert the test media file into the database - err := mr.Put(&mediaFileWithMBID) + err := mr.Put(ctx, &mediaFileWithMBID) Expect(err).ToNot(HaveOccurred()) }) AfterEach(func() { // Clean up test data using direct SQL - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": mediaFileWithMBID.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": mediaFileWithMBID.ID})) }) It("finds media file by mbz_recording_id", func() { - results, err := mr.Search("550e8400-e29b-41d4-a716-446655440020", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "550e8400-e29b-41d4-a716-446655440020", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("test-mbid-mediafile")) @@ -838,7 +835,7 @@ var _ = Describe("MediaRepository", func() { }) It("finds media file by mbz_release_track_id", func() { - results, err := mr.Search("550e8400-e29b-41d4-a716-446655440021", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "550e8400-e29b-41d4-a716-446655440021", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("test-mbid-mediafile")) @@ -846,7 +843,7 @@ var _ = Describe("MediaRepository", func() { }) It("returns empty result when MBID is not found", func() { - results, err := mr.Search("550e8400-e29b-41d4-a716-446655440099", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "550e8400-e29b-41d4-a716-446655440099", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -862,22 +859,22 @@ var _ = Describe("MediaRepository", func() { Missing: true, } - err := mr.Put(&missingMediaFile) + err := mr.Put(ctx, &missingMediaFile) Expect(err).ToNot(HaveOccurred()) // Search never returns missing media files (hardcoded behavior) - results, err := mr.Search("550e8400-e29b-41d4-a716-446655440022", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "550e8400-e29b-41d4-a716-446655440022", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) // Clean up - _, _ = raw.executeSQL(squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": missingMediaFile.ID})) + _, _ = raw.executeSQL(ctx, squirrel.Delete(raw.tableName).Where(squirrel.Eq{"id": missingMediaFile.ID})) }) }) Context("empty query (natural order pagination)", func() { It("returns all non-missing files in natural order", func() { - results, err := mr.Search("", model.QueryOptions{Max: 1000}) + results, err := mr.Search(ctx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) Expect(results).ToNot(BeEmpty()) for _, result := range results { @@ -886,22 +883,22 @@ var _ = Describe("MediaRepository", func() { }) It(`treats quoted empty query ("") the same as empty`, func() { - all, err := mr.Search("", model.QueryOptions{Max: 1000}) + all, err := mr.Search(ctx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) - quoted, err := mr.Search(`""`, model.QueryOptions{Max: 1000}) + quoted, err := mr.Search(ctx, `""`, model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) Expect(quoted).To(HaveLen(len(all))) }) It("paginates without overlaps or gaps", func() { - all, err := mr.Search("", model.QueryOptions{Max: 1000}) + all, err := mr.Search(ctx, "", model.QueryOptions{Max: 1000}) Expect(err).ToNot(HaveOccurred()) Expect(len(all)).To(BeNumerically(">", 3)) var paged model.MediaFiles pageSize := 3 for offset := 0; offset < len(all); offset += pageSize { - page, err := mr.Search("", model.QueryOptions{Max: pageSize, Offset: offset}) + page, err := mr.Search(ctx, "", model.QueryOptions{Max: pageSize, Offset: offset}) Expect(err).ToNot(HaveOccurred()) paged = append(paged, page...) } @@ -912,7 +909,7 @@ var _ = Describe("MediaRepository", func() { }) It("returns empty page when offset is beyond the total", func() { - results, err := mr.Search("", model.QueryOptions{Max: 10, Offset: 100000}) + results, err := mr.Search(ctx, "", model.QueryOptions{Max: 10, Offset: 100000}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -926,41 +923,41 @@ var _ = Describe("MediaRepository", func() { BeforeEach(func() { ctx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) - pr = NewPlaylistRepository(ctx, GetDBXBuilder()) + pr = NewPlaylistRepository(GetDBXBuilder()) prev = model.MediaFile{ID: "reassign-prev", LibraryID: 1, Path: "reassign/prev.mp3", Title: "Prev"} next = model.MediaFile{ID: "reassign-next", LibraryID: 1, Path: "reassign/next.mp3", Title: "Next"} - Expect(mr.Put(&prev)).To(Succeed()) - Expect(mr.Put(&next)).To(Succeed()) + Expect(mr.Put(ctx, &prev)).To(Succeed()) + Expect(mr.Put(ctx, &next)).To(Succeed()) pls = model.Playlist{Name: "Reassign", OwnerID: "userid"} pls.AddMediaFilesByID([]string{prev.ID}) - Expect(pr.Put(&pls)).To(Succeed()) + Expect(pr.Put(ctx, &pls)).To(Succeed()) }) AfterEach(func() { - _ = pr.Delete(pls.ID) - _ = mr.Delete(prev.ID) - _ = mr.Delete(next.ID) - _, _ = mr.(*mediaFileRepository).executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": []string{prev.ID, next.ID}})) - _, _ = mr.(*mediaFileRepository).executeSQL(squirrel.Delete("bookmark").Where(squirrel.Eq{"item_id": []string{prev.ID, next.ID}})) - _, _ = mr.(*mediaFileRepository).executeSQL(squirrel.Delete("scrobbles").Where(squirrel.Eq{"media_file_id": []string{prev.ID, next.ID}})) - _, _ = mr.(*mediaFileRepository).executeSQL(squirrel.Delete("scrobble_buffer").Where(squirrel.Eq{"media_file_id": []string{prev.ID, next.ID}})) + _ = pr.Delete(ctx, pls.ID) + _ = mr.Delete(ctx, prev.ID) + _ = mr.Delete(ctx, next.ID) + _, _ = mr.(*mediaFileRepository).executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": []string{prev.ID, next.ID}})) + _, _ = mr.(*mediaFileRepository).executeSQL(ctx, squirrel.Delete("bookmark").Where(squirrel.Eq{"item_id": []string{prev.ID, next.ID}})) + _, _ = mr.(*mediaFileRepository).executeSQL(ctx, squirrel.Delete("scrobbles").Where(squirrel.Eq{"media_file_id": []string{prev.ID, next.ID}})) + _, _ = mr.(*mediaFileRepository).executeSQL(ctx, squirrel.Delete("scrobble_buffer").Where(squirrel.Eq{"media_file_id": []string{prev.ID, next.ID}})) }) It("moves annotations, bookmarks and playlist entries onto the new id", func() { - Expect(mr.SetRating(5, prev.ID)).To(Succeed()) - Expect(mr.AddBookmark(prev.ID, "here", 42)).To(Succeed()) + Expect(mr.SetRating(ctx, 5, prev.ID)).To(Succeed()) + Expect(mr.AddBookmark(ctx, prev.ID, "here", 42)).To(Succeed()) - Expect(mr.ReassignReferences(prev.ID, next.ID)).To(Succeed()) + Expect(mr.ReassignReferences(ctx, prev.ID, next.ID)).To(Succeed()) - got, err := mr.Get(next.ID) + got, err := mr.Get(ctx, next.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.Rating).To(Equal(5)) - bookmarks, err := mr.GetBookmarks() + bookmarks, err := mr.GetBookmarks(ctx) Expect(err).ToNot(HaveOccurred()) Expect(bookmarks).To(ContainElement(HaveField("Item.ID", next.ID))) - withTracks, err := pr.GetWithTracks(pls.ID, false, false) + withTracks, err := pr.GetWithTracks(ctx, pls.ID, false, false) Expect(err).ToNot(HaveOccurred()) Expect(withTracks.Tracks).To(HaveLen(1)) Expect(withTracks.Tracks[0].MediaFileID).To(Equal(next.ID)) @@ -968,51 +965,52 @@ var _ = Describe("MediaRepository", func() { It("moves scrobbles and buffered scrobbles onto the new id", func() { ctx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) - scrobbles := NewScrobbleRepository(ctx, GetDBXBuilder()) - buffer := NewScrobbleBufferRepository(ctx, GetDBXBuilder()) - Expect(scrobbles.RecordScrobble(prev.ID, time.Now())).To(Succeed()) - Expect(buffer.Enqueue("lastfm", "userid", prev.ID, time.Now())).To(Succeed()) + scrobbles := NewScrobbleRepository(GetDBXBuilder()) + buffer := NewScrobbleBufferRepository(GetDBXBuilder()) + Expect(scrobbles.RecordScrobble(ctx, prev.ID, time.Now())).To(Succeed()) + Expect(buffer.Enqueue(ctx, "lastfm", "userid", prev.ID, time.Now())).To(Succeed()) - Expect(mr.ReassignReferences(prev.ID, next.ID)).To(Succeed()) - Expect(mr.Delete(prev.ID)).To(Succeed()) + Expect(mr.ReassignReferences(ctx, prev.ID, next.ID)).To(Succeed()) + Expect(mr.Delete(ctx, prev.ID)).To(Succeed()) - all, err := scrobbles.GetAll() + all, err := scrobbles.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) mine := slice.Map(slice.Filter(all, func(sc model.Scrobble) bool { return sc.MediaFileID == prev.ID || sc.MediaFileID == next.ID }), func(sc model.Scrobble) string { return sc.MediaFileID }) Expect(mine).To(ConsistOf(next.ID)) - entry, err := buffer.Next("lastfm", "userid") + entry, err := buffer.Next(ctx, "lastfm", "userid") Expect(err).ToNot(HaveOccurred()) Expect(entry).ToNot(BeNil()) Expect(entry.MediaFile.ID).To(Equal(next.ID)) }) It("recomputes the average rating after merging another user's annotation", func() { - other := NewMediaFileRepository(request.WithUser(log.NewContext(context.TODO()), model.User{ID: "2222"}), GetDBXBuilder()) - Expect(mr.SetRating(5, next.ID)).To(Succeed()) - Expect(other.SetRating(3, prev.ID)).To(Succeed()) + otherCtx := request.WithUser(ctx, model.User{ID: "2222"}) + other := NewMediaFileRepository(GetDBXBuilder()) + Expect(mr.SetRating(ctx, 5, next.ID)).To(Succeed()) + Expect(other.SetRating(otherCtx, 3, prev.ID)).To(Succeed()) - Expect(mr.ReassignReferences(prev.ID, next.ID)).To(Succeed()) + Expect(mr.ReassignReferences(ctx, prev.ID, next.ID)).To(Succeed()) - got, err := mr.Get(next.ID) + got, err := mr.Get(ctx, next.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.AverageRating).To(Equal(4.0)) }) It("keeps the new id's own annotation and bookmark when both exist", func() { - Expect(mr.SetRating(5, prev.ID)).To(Succeed()) - Expect(mr.SetRating(1, next.ID)).To(Succeed()) - Expect(mr.AddBookmark(prev.ID, "prev", 42)).To(Succeed()) - Expect(mr.AddBookmark(next.ID, "next", 7)).To(Succeed()) + Expect(mr.SetRating(ctx, 5, prev.ID)).To(Succeed()) + Expect(mr.SetRating(ctx, 1, next.ID)).To(Succeed()) + Expect(mr.AddBookmark(ctx, prev.ID, "prev", 42)).To(Succeed()) + Expect(mr.AddBookmark(ctx, next.ID, "next", 7)).To(Succeed()) - Expect(mr.ReassignReferences(prev.ID, next.ID)).To(Succeed()) + Expect(mr.ReassignReferences(ctx, prev.ID, next.ID)).To(Succeed()) - got, err := mr.Get(next.ID) + got, err := mr.Get(ctx, next.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.Rating).To(Equal(1)) - bookmarks, err := mr.GetBookmarks() + bookmarks, err := mr.GetBookmarks(ctx) Expect(err).ToNot(HaveOccurred()) Expect(bookmarks).To(ContainElement(SatisfyAll(HaveField("Item.ID", next.ID), HaveField("Comment", "next")))) }) @@ -1034,39 +1032,39 @@ var _ = Describe("MediaRepository", func() { {ID: "findpath-6", LibraryID: 1, Path: "1999: A Different Life/01.mp3", Title: "Numeric colon"}, } for _, mf := range testFiles { - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) } }) AfterEach(func() { for _, mf := range testFiles { - _ = mr.Delete(mf.ID) + _ = mr.Delete(ctx, mf.ID) } }) It("treats a path whose prefix is not a library id as unqualified", func() { - results, err := mr.FindByPaths([]string{"Bach: Goldberg Variations/01.mp3"}) + results, err := mr.FindByPaths(ctx, []string{"Bach: Goldberg Variations/01.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("findpath-5")) }) It("finds a plain path whose colon prefix looks like a library id", func() { - results, err := mr.FindByPaths([]string{"1999: A Different Life/01.mp3"}) + results, err := mr.FindByPaths(ctx, []string{"1999: A Different Life/01.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("findpath-6")) }) It("splits only the first colon of a library-qualified path", func() { - results, err := mr.FindByPaths([]string{"1:Bach: Goldberg Variations/01.mp3"}) + results, err := mr.FindByPaths(ctx, []string{"1:Bach: Goldberg Variations/01.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("findpath-5")) }) It("finds files by exact path", func() { - results, err := mr.FindByPaths([]string{"1:artist/Album/track.mp3"}) + results, err := mr.FindByPaths(ctx, []string{"1:artist/Album/track.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("findpath-1")) @@ -1074,7 +1072,7 @@ var _ = Describe("MediaRepository", func() { It("finds files case-insensitively for ASCII characters (NOCASE)", func() { // SQLite's COLLATE NOCASE handles ASCII case-insensitivity - results, err := mr.FindByPaths([]string{"1:ARTIST/ALBUM/TRACK.MP3"}) + results, err := mr.FindByPaths(ctx, []string{"1:ARTIST/ALBUM/TRACK.MP3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("findpath-1")) @@ -1083,20 +1081,20 @@ var _ = Describe("MediaRepository", func() { It("finds fullwidth characters only with exact case match (SQLite NOCASE limitation)", func() { // SQLite's NOCASE does NOT handle fullwidth uppercase/lowercase equivalence // The DB has fullwidth uppercase ACROSS, searching with exact match should work - results, err := mr.FindByPaths([]string{"1:plex/02 - ACROSS.flac"}) + results, err := mr.FindByPaths(ctx, []string{"1:plex/02 - ACROSS.flac"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("findpath-3")) // Searching with fullwidth lowercase across should NOT match // (this is the SQLite limitation that requires exact matching for non-ASCII) - results, err = mr.FindByPaths([]string{"1:plex/02 - across.flac"}) + results, err = mr.FindByPaths(ctx, []string{"1:plex/02 - across.flac"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) It("returns multiple files when querying multiple paths", func() { - results, err := mr.FindByPaths([]string{ + results, err := mr.FindByPaths(ctx, []string{ "1:artist/Album/track.mp3", "1:artist/Album/UPPER.mp3", }) @@ -1105,25 +1103,25 @@ var _ = Describe("MediaRepository", func() { }) It("returns empty slice for non-existent paths", func() { - results, err := mr.FindByPaths([]string{"1:nonexistent/path.mp3"}) + results, err := mr.FindByPaths(ctx, []string{"1:nonexistent/path.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) It("returns empty slice for empty input", func() { - results, err := mr.FindByPaths([]string{}) + results, err := mr.FindByPaths(ctx, []string{}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) It("handles library-qualified paths correctly", func() { // Library 1 should find the file - results, err := mr.FindByPaths([]string{"1:artist/Album/track.mp3"}) + results, err := mr.FindByPaths(ctx, []string{"1:artist/Album/track.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) // Library 2 should NOT find it (file is in library 1) - results, err = mr.FindByPaths([]string{"2:artist/Album/track.mp3"}) + results, err = mr.FindByPaths(ctx, []string{"2:artist/Album/track.mp3"}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty()) }) @@ -1134,54 +1132,57 @@ var _ = Describe("MediaRepository", func() { BeforeEach(func() { adminCtx := request.WithUser(GinkgoT().Context(), adminUser) - lr := NewLibraryRepository(adminCtx, GetDBXBuilder()) + lr := NewLibraryRepository(GetDBXBuilder()) // A second library the restricted user has no access to otherLib = model.Library{ID: 0, Name: "Other Library", Path: "/other/lib"} - Expect(lr.Put(&otherLib)).To(Succeed()) + Expect(lr.Put(adminCtx, &otherLib)).To(Succeed()) // A track that lives only in the other library (created as admin) - adminMr := NewMediaFileRepository(adminCtx, GetDBXBuilder()) - Expect(adminMr.Put(&model.MediaFile{ + adminMr := NewMediaFileRepository(GetDBXBuilder()) + Expect(adminMr.Put(adminCtx, &model.MediaFile{ ID: "otherlib-track", LibraryID: otherLib.ID, Path: "hidden/test.mp3", Title: "Hidden", })).To(Succeed()) // Non-admin user with access to library 1 ONLY restrictedUser = createUserWithLibraries("restricted-finder", []int{1}) - ur := NewUserRepository(adminCtx, GetDBXBuilder()) - Expect(ur.Put(&restrictedUser)).To(Succeed()) - Expect(ur.SetUserLibraries(restrictedUser.ID, []int{1})).To(Succeed()) + ur := NewUserRepository(GetDBXBuilder()) + Expect(ur.Put(adminCtx, &restrictedUser)).To(Succeed()) + Expect(ur.SetUserLibraries(adminCtx, restrictedUser.ID, []int{1})).To(Succeed()) }) AfterEach(func() { adminCtx := request.WithUser(GinkgoT().Context(), adminUser) - _ = NewMediaFileRepository(adminCtx, GetDBXBuilder()).Delete("otherlib-track") - lr := NewLibraryRepository(adminCtx, GetDBXBuilder()).(*libraryRepository) - _ = lr.delete(squirrel.Eq{"id": otherLib.ID}) - _ = NewUserRepository(adminCtx, GetDBXBuilder()).Delete(restrictedUser.ID) + _ = NewMediaFileRepository(GetDBXBuilder()).Delete(adminCtx, "otherlib-track") + lr := NewLibraryRepository(GetDBXBuilder()).(*libraryRepository) + _ = lr.delete(adminCtx, squirrel.Eq{"id": otherLib.ID}) + _ = NewUserRepository(GetDBXBuilder()).Delete(adminCtx, restrictedUser.ID) }) It("does not resolve paths in libraries the user cannot access", func() { - userMr := NewMediaFileRepository(request.WithUser(GinkgoT().Context(), restrictedUser), GetDBXBuilder()) + userCtx := request.WithUser(ctx, restrictedUser) + userMr := NewMediaFileRepository(GetDBXBuilder()) qualified := fmt.Sprintf("%d:hidden/test.mp3", otherLib.ID) - results, err := userMr.FindByPaths([]string{qualified}) + results, err := userMr.FindByPaths(userCtx, []string{qualified}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty(), "a track outside the user's libraries must not be resolvable") }) It("still resolves the path for an admin", func() { - adminMr := NewMediaFileRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) + adminCtx := request.WithUser(ctx, adminUser) + adminMr := NewMediaFileRepository(GetDBXBuilder()) qualified := fmt.Sprintf("%d:hidden/test.mp3", otherLib.ID) - results, err := adminMr.FindByPaths([]string{qualified}) + results, err := adminMr.FindByPaths(adminCtx, []string{qualified}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal("otherlib-track")) }) It("resolves paths from multiple libraries in a single call", func() { - adminMr := NewMediaFileRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) - results, err := adminMr.FindByPaths([]string{ + adminCtx := request.WithUser(ctx, adminUser) + adminMr := NewMediaFileRepository(GetDBXBuilder()) + results, err := adminMr.FindByPaths(adminCtx, []string{ "1:artist/Album/track.mp3", fmt.Sprintf("%d:hidden/test.mp3", otherLib.ID), }) @@ -1191,9 +1192,10 @@ var _ = Describe("MediaRepository", func() { }) It("keeps each path scoped to its own library when several are queried", func() { - adminMr := NewMediaFileRepository(request.WithUser(GinkgoT().Context(), adminUser), GetDBXBuilder()) + adminCtx := request.WithUser(ctx, adminUser) + adminMr := NewMediaFileRepository(GetDBXBuilder()) // Each path exists, but under the other library's ID, so neither must match. - results, err := adminMr.FindByPaths([]string{ + results, err := adminMr.FindByPaths(adminCtx, []string{ fmt.Sprintf("%d:artist/Album/track.mp3", otherLib.ID), "1:hidden/test.mp3", }) @@ -1254,9 +1256,9 @@ var _ = Describe("MediaRepository", func() { It("stores nil BPM and BitDepth as NULL and retrieves them as nil", func() { newID := id.NewRandom() mf := model.MediaFile{LibraryID: 1, ID: newID, Path: "test/bpm-nil.mp3"} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) - retrieved, err := mr.Get(newID) + retrieved, err := mr.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(retrieved.BPM).To(BeNil()) Expect(retrieved.BitDepth).To(BeNil()) @@ -1274,7 +1276,7 @@ var _ = Describe("MediaRepository", func() { Expect(row.BPM).To(BeNil(), "bpm should be stored as NULL in the database") Expect(row.BitDepth).To(BeNil(), "bit_depth should be stored as NULL in the database") - _ = mr.Delete(newID) + _ = mr.Delete(ctx, newID) }) It("stores non-nil BPM and BitDepth and retrieves correct values", func() { @@ -1282,16 +1284,16 @@ var _ = Describe("MediaRepository", func() { bpm := 120 bitDepth := 24 mf := model.MediaFile{LibraryID: 1, ID: newID, Path: "test/bpm-set.mp3", BPM: &bpm, BitDepth: &bitDepth} - Expect(mr.Put(&mf)).To(Succeed()) + Expect(mr.Put(ctx, &mf)).To(Succeed()) - retrieved, err := mr.Get(newID) + retrieved, err := mr.Get(ctx, newID) Expect(err).ToNot(HaveOccurred()) Expect(retrieved.BPM).ToNot(BeNil()) Expect(*retrieved.BPM).To(Equal(120)) Expect(retrieved.BitDepth).ToNot(BeNil()) Expect(*retrieved.BitDepth).To(Equal(24)) - _ = mr.Delete(newID) + _ = mr.Delete(ctx, newID) }) }) @@ -1326,34 +1328,34 @@ var _ = Describe("MediaRepository", func() { restricted := model.User{ID: "restricted_mf_user", UserName: "rm", Name: "RM", Email: "rm@t.com"} rctx := request.WithUser(GinkgoT().Context(), restricted) - Expect(mr.Exists(songAntenna.ID)).To(BeTrue(), "admin sees it") - Expect(NewMediaFileRepository(rctx, GetDBXBuilder()).Exists(songAntenna.ID)).To(BeFalse()) + Expect(mr.Exists(ctx, songAntenna.ID)).To(BeTrue(), "admin sees it") + Expect(NewMediaFileRepository(GetDBXBuilder()).Exists(rctx, songAntenna.ID)).To(BeFalse()) }) }) Describe("MatchesCriteria", func() { It("returns true when the track matches", func() { c := criteria.Criteria{Expression: criteria.All{criteria.Contains{"title": "Day"}}} - match, err := mr.MatchesCriteria(songDayInALife.ID, c) + match, err := mr.MatchesCriteria(ctx, songDayInALife.ID, c) Expect(err).ToNot(HaveOccurred()) Expect(match).To(BeTrue()) }) It("returns false when the track does not match", func() { c := criteria.Criteria{Expression: criteria.All{criteria.Contains{"title": "Nickelback"}}} - match, err := mr.MatchesCriteria(songDayInALife.ID, c) + match, err := mr.MatchesCriteria(ctx, songDayInALife.ID, c) Expect(err).ToNot(HaveOccurred()) Expect(match).To(BeFalse()) }) It("treats missing annotations as their COALESCE default", func() { // unrated track: rating coalesces to 0, so "rating < 4" matches c := criteria.Criteria{Expression: criteria.All{criteria.Lt{"rating": 4}}} - match, err := mr.MatchesCriteria(songDayInALife.ID, c) + match, err := mr.MatchesCriteria(ctx, songDayInALife.ID, c) Expect(err).ToNot(HaveOccurred()) Expect(match).To(BeTrue()) }) It("returns an error for an invalid field", func() { c := criteria.Criteria{Expression: criteria.All{criteria.Is{"bogusfield": 1}}} - _, err := mr.MatchesCriteria(songDayInALife.ID, c) + _, err := mr.MatchesCriteria(ctx, songDayInALife.ID, c) Expect(err).To(HaveOccurred()) }) }) diff --git a/persistence/persistence.go b/persistence/persistence.go index 589812266..44e944bff 100644 --- a/persistence/persistence.go +++ b/persistence/persistence.go @@ -4,7 +4,7 @@ import ( "context" "database/sql" "fmt" - "reflect" + "sync" "time" "github.com/navidrome/navidrome/db" @@ -15,128 +15,144 @@ import ( ) type SQLStore struct { - db dbx.Builder + db dbx.Builder + library func() model.LibraryRepository + folder func() model.FolderRepository + album func() model.AlbumRepository + artist func() model.ArtistRepository + mediaFile func() model.MediaFileRepository + genre func() model.GenreRepository + tag func() model.TagRepository + playlist func() model.PlaylistRepository + playQueue func() model.PlayQueueRepository + transcoding func() model.TranscodingRepository + player func() model.PlayerRepository + radio func() model.RadioRepository + share func() model.ShareRepository + property func() model.PropertyRepository + user func() model.UserRepository + userProps func() model.UserPropsRepository + scrobbleBuf func() model.ScrobbleBufferRepository + scrobble func() model.ScrobbleRepository + plugin func() model.PluginRepository + artwork func() model.ArtworkRepository + artworkQueue func() model.ArtworkQueueRepository +} + +// Repositories are built on first use, so a transaction store only pays for the ones its block touches. +func newSQLStore(db dbx.Builder) *SQLStore { + return &SQLStore{ + db: db, + library: sync.OnceValue(func() model.LibraryRepository { return NewLibraryRepository(db) }), + folder: sync.OnceValue(func() model.FolderRepository { return newFolderRepository(db) }), + album: sync.OnceValue(func() model.AlbumRepository { return NewAlbumRepository(db) }), + artist: sync.OnceValue(func() model.ArtistRepository { return NewArtistRepository(db) }), + mediaFile: sync.OnceValue(func() model.MediaFileRepository { return NewMediaFileRepository(db) }), + genre: sync.OnceValue(func() model.GenreRepository { return NewGenreRepository(db) }), + tag: sync.OnceValue(func() model.TagRepository { return NewTagRepository(db) }), + playlist: sync.OnceValue(func() model.PlaylistRepository { return NewPlaylistRepository(db) }), + playQueue: sync.OnceValue(func() model.PlayQueueRepository { return NewPlayQueueRepository(db) }), + transcoding: sync.OnceValue(func() model.TranscodingRepository { return NewTranscodingRepository(db) }), + player: sync.OnceValue(func() model.PlayerRepository { return NewPlayerRepository(db) }), + radio: sync.OnceValue(func() model.RadioRepository { return NewRadioRepository(db) }), + share: sync.OnceValue(func() model.ShareRepository { return NewShareRepository(db) }), + property: sync.OnceValue(func() model.PropertyRepository { return NewPropertyRepository(db) }), + user: sync.OnceValue(func() model.UserRepository { return NewUserRepository(db) }), + userProps: sync.OnceValue(func() model.UserPropsRepository { return NewUserPropsRepository(db) }), + scrobbleBuf: sync.OnceValue(func() model.ScrobbleBufferRepository { return NewScrobbleBufferRepository(db) }), + scrobble: sync.OnceValue(func() model.ScrobbleRepository { return NewScrobbleRepository(db) }), + plugin: sync.OnceValue(func() model.PluginRepository { return NewPluginRepository(db) }), + artwork: sync.OnceValue(func() model.ArtworkRepository { return NewArtworkRepository(db) }), + artworkQueue: sync.OnceValue(func() model.ArtworkQueueRepository { return NewArtworkQueueRepository(db) }), + } } func New(conn *sql.DB) model.DataStore { - return &SQLStore{db: dbx.NewFromDB(conn, db.Driver)} + return newSQLStore(dbx.NewFromDB(conn, db.Driver)) } -func (s *SQLStore) Album(ctx context.Context) model.AlbumRepository { - return NewAlbumRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Album() model.AlbumRepository { + return s.album() } -func (s *SQLStore) Artist(ctx context.Context) model.ArtistRepository { - return NewArtistRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Artist() model.ArtistRepository { + return s.artist() } -func (s *SQLStore) MediaFile(ctx context.Context) model.MediaFileRepository { - return NewMediaFileRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) MediaFile() model.MediaFileRepository { + return s.mediaFile() } -func (s *SQLStore) Library(ctx context.Context) model.LibraryRepository { - return NewLibraryRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Library() model.LibraryRepository { + return s.library() } -func (s *SQLStore) Folder(ctx context.Context) model.FolderRepository { - return newFolderRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Folder() model.FolderRepository { + return s.folder() } -func (s *SQLStore) Genre(ctx context.Context) model.GenreRepository { - return NewGenreRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Genre() model.GenreRepository { + return s.genre() } -func (s *SQLStore) Tag(ctx context.Context) model.TagRepository { - return NewTagRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Tag() model.TagRepository { + return s.tag() } -func (s *SQLStore) PlayQueue(ctx context.Context) model.PlayQueueRepository { - return NewPlayQueueRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) PlayQueue() model.PlayQueueRepository { + return s.playQueue() } -func (s *SQLStore) Playlist(ctx context.Context) model.PlaylistRepository { - return NewPlaylistRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Playlist() model.PlaylistRepository { + return s.playlist() } -func (s *SQLStore) Property(ctx context.Context) model.PropertyRepository { - return NewPropertyRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Property() model.PropertyRepository { + return s.property() } -func (s *SQLStore) Radio(ctx context.Context) model.RadioRepository { - return NewRadioRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Radio() model.RadioRepository { + return s.radio() } -func (s *SQLStore) UserProps(ctx context.Context) model.UserPropsRepository { - return NewUserPropsRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) UserProps() model.UserPropsRepository { + return s.userProps() } -func (s *SQLStore) Share(ctx context.Context) model.ShareRepository { - return NewShareRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Share() model.ShareRepository { + return s.share() } -func (s *SQLStore) User(ctx context.Context) model.UserRepository { - return NewUserRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) User() model.UserRepository { + return s.user() } -func (s *SQLStore) Transcoding(ctx context.Context) model.TranscodingRepository { - return NewTranscodingRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Transcoding() model.TranscodingRepository { + return s.transcoding() } -func (s *SQLStore) Player(ctx context.Context) model.PlayerRepository { - return NewPlayerRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Player() model.PlayerRepository { + return s.player() } -func (s *SQLStore) ScrobbleBuffer(ctx context.Context) model.ScrobbleBufferRepository { - return NewScrobbleBufferRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) ScrobbleBuffer() model.ScrobbleBufferRepository { + return s.scrobbleBuf() } -func (s *SQLStore) Scrobble(ctx context.Context) model.ScrobbleRepository { - return NewScrobbleRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Scrobble() model.ScrobbleRepository { + return s.scrobble() } -func (s *SQLStore) Plugin(ctx context.Context) model.PluginRepository { - return NewPluginRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Plugin() model.PluginRepository { + return s.plugin() } -func (s *SQLStore) Artwork(ctx context.Context) model.ArtworkRepository { - return NewArtworkRepository(ctx, s.getDBXBuilder()) +func (s *SQLStore) Artwork() model.ArtworkRepository { + return s.artwork() } -func (s *SQLStore) ArtworkQueue(ctx context.Context) model.ArtworkQueueRepository { - return NewArtworkQueueRepository(ctx, s.getDBXBuilder()) -} - -func (s *SQLStore) Resource(ctx context.Context, m any) model.ResourceRepository { - switch m.(type) { - case model.User: - return s.User(ctx).(model.ResourceRepository) - case model.Transcoding: - return s.Transcoding(ctx).(model.ResourceRepository) - case model.Player: - return s.Player(ctx).(model.ResourceRepository) - case model.Artist: - return s.Artist(ctx).(model.ResourceRepository) - case model.Album: - return s.Album(ctx).(model.ResourceRepository) - case model.MediaFile: - return s.MediaFile(ctx).(model.ResourceRepository) - case model.Genre: - return s.Genre(ctx).(model.ResourceRepository) - case model.Playlist: - return s.Playlist(ctx).(model.ResourceRepository) - case model.Radio: - return s.Radio(ctx).(model.ResourceRepository) - case model.Share: - return s.Share(ctx).(model.ResourceRepository) - case model.Tag: - return s.Tag(ctx).(model.ResourceRepository) - case model.Plugin: - return s.Plugin(ctx).(model.ResourceRepository) - case model.Scrobble: - return s.Scrobble(ctx).(model.ResourceRepository) - } - log.Error("Resource not implemented", "model", reflect.TypeOf(m).Name()) - return nil +func (s *SQLStore) ArtworkQueue() model.ArtworkQueueRepository { + return s.artworkQueue() } func scopeLabel(scope []string) string { @@ -157,7 +173,7 @@ func (s *SQLStore) WithTx(block func(tx model.DataStore) error, scope ...string) log.Trace("Transaction started", "scope", msg) } return conn.Transactional(func(tx *dbx.Tx) error { - newDb := &SQLStore{db: tx} + newDb := newSQLStore(tx) err := block(newDb) if !inTx { log.Trace("Nested Transaction finished", "scope", msg, "elapsed", time.Since(start), err) @@ -173,9 +189,9 @@ func (s *SQLStore) WithTxImmediate(block func(tx model.DataStore) error, scope . return s.WithTx(func(tx model.DataStore) error { // Workaround to force the transaction to be upgraded to immediate mode to avoid deadlocks // See https://berthub.eu/articles/posts/a-brief-post-on-sqlite3-database-locked-despite-timeout/ - _ = tx.Property(ctx).Put("tmp_lock_flag", "") + _ = tx.Property().Put(ctx, "tmp_lock_flag", "") defer func() { - _ = tx.Property(ctx).Delete("tmp_lock_flag") + _ = tx.Property().Delete(ctx, "tmp_lock_flag") }() return block(tx) @@ -244,27 +260,20 @@ func (s *SQLStore) GC(ctx context.Context, libraryIDs ...int) error { } err := run.Sequentially( - trace(ctx, "purge empty albums", func() error { return s.Album(ctx).(*albumRepository).purgeEmpty(libraryIDs...) }), - trace(ctx, "purge empty artists", func() error { return s.Artist(ctx).(*artistRepository).purgeEmpty() }), - trace(ctx, "mark missing artists", func() error { return s.Artist(ctx).(*artistRepository).markMissing() }), - trace(ctx, "purge empty folders", func() error { return s.Folder(ctx).(*folderRepository).purgeEmpty(libraryIDs...) }), - trace(ctx, "clean album annotations", func() error { return s.Album(ctx).(*albumRepository).cleanAnnotations() }), - trace(ctx, "clean artist annotations", func() error { return s.Artist(ctx).(*artistRepository).cleanAnnotations() }), - trace(ctx, "clean media file annotations", func() error { return s.MediaFile(ctx).(*mediaFileRepository).cleanAnnotations() }), - trace(ctx, "clean playlist annotations", func() error { return s.Playlist(ctx).(*playlistRepository).cleanAnnotations() }), - trace(ctx, "clean media file bookmarks", func() error { return s.MediaFile(ctx).(*mediaFileRepository).cleanBookmarks() }), - trace(ctx, "purge non used tags", func() error { return s.Tag(ctx).(*tagRepository).purgeUnused() }), - trace(ctx, "remove orphan playlist tracks", func() error { return s.Playlist(ctx).(*playlistRepository).removeOrphans() }), + trace(ctx, "purge empty albums", func() error { return s.album().(*albumRepository).purgeEmpty(ctx, libraryIDs...) }), + trace(ctx, "purge empty artists", func() error { return s.artist().(*artistRepository).purgeEmpty(ctx) }), + trace(ctx, "mark missing artists", func() error { return s.artist().(*artistRepository).markMissing(ctx) }), + trace(ctx, "purge empty folders", func() error { return s.folder().(*folderRepository).purgeEmpty(ctx, libraryIDs...) }), + trace(ctx, "clean album annotations", func() error { return s.album().(*albumRepository).cleanAnnotations(ctx) }), + trace(ctx, "clean artist annotations", func() error { return s.artist().(*artistRepository).cleanAnnotations(ctx) }), + trace(ctx, "clean media file annotations", func() error { return s.mediaFile().(*mediaFileRepository).cleanAnnotations(ctx) }), + trace(ctx, "clean playlist annotations", func() error { return s.playlist().(*playlistRepository).cleanAnnotations(ctx) }), + trace(ctx, "clean media file bookmarks", func() error { return s.mediaFile().(*mediaFileRepository).cleanBookmarks(ctx) }), + trace(ctx, "purge non used tags", func() error { return s.tag().(*tagRepository).purgeUnused(ctx) }), + trace(ctx, "remove orphan playlist tracks", func() error { return s.playlist().(*playlistRepository).removeOrphans(ctx) }), ) if err != nil { return fmt.Errorf("tidying up database: %w", err) } return nil } - -func (s *SQLStore) getDBXBuilder() dbx.Builder { - if s.db == nil { - return dbx.NewFromDB(db.Db(), db.Driver) - } - return s.db -} diff --git a/persistence/persistence_suite_test.go b/persistence/persistence_suite_test.go index 644284c0b..ee2794454 100644 --- a/persistence/persistence_suite_test.go +++ b/persistence/persistence_suite_test.go @@ -176,17 +176,17 @@ func restrictedFixture(name string) (context.Context, model.Library, model.User) db := GetDBXBuilder() lib := model.Library{Name: name + " Library", Path: "/" + name} - lr := NewLibraryRepository(adminCtx, db) - Expect(lr.Put(&lib)).To(Succeed()) + lr := NewLibraryRepository(db) + Expect(lr.Put(adminCtx, &lib)).To(Succeed()) user := createUserWithLibraries(name+"-restricted", []int{1}) - ur := NewUserRepository(adminCtx, db) - Expect(ur.Put(&user)).To(Succeed()) - Expect(ur.SetUserLibraries(user.ID, []int{1})).To(Succeed()) + ur := NewUserRepository(db) + Expect(ur.Put(adminCtx, &user)).To(Succeed()) + Expect(ur.SetUserLibraries(adminCtx, user.ID, []int{1})).To(Succeed()) DeferCleanup(func() { - _ = NewUserRepository(adminCtx, db).Delete(user.ID) - _ = NewLibraryRepository(adminCtx, db).(*libraryRepository).delete(squirrel.Eq{"id": lib.ID}) + _ = NewUserRepository(db).Delete(adminCtx, user.ID) + _ = NewLibraryRepository(db).(*libraryRepository).delete(adminCtx, squirrel.Eq{"id": lib.ID}) }) return adminCtx, lib, user } @@ -196,9 +196,9 @@ var _ = BeforeSuite(func() { ctx := log.NewContext(context.TODO()) ctx = request.WithUser(ctx, adminUser) - ur := NewUserRepository(ctx, conn) + ur := NewUserRepository(conn) for i := range testUsers { - err := ur.Put(&testUsers[i]) + err := ur.Put(ctx, &testUsers[i]) if err != nil { panic(err) } @@ -206,32 +206,32 @@ var _ = BeforeSuite(func() { // Associate users with library 1 (default test library) for i := range testUsers { - err := ur.SetUserLibraries(testUsers[i].ID, []int{1}) + err := ur.SetUserLibraries(ctx, testUsers[i].ID, []int{1}) if err != nil { panic(err) } } - alr := NewAlbumRepository(ctx, conn).(*albumRepository) + alr := NewAlbumRepository(conn).(*albumRepository) for i := range testAlbums { - err := alr.Put(new(testAlbums[i])) + err := alr.Put(ctx, new(testAlbums[i])) if err != nil { panic(err) } } - arr := NewArtistRepository(ctx, conn) + arr := NewArtistRepository(conn) for i := range testArtists { - err := arr.Put(new(testArtists[i])) + err := arr.Put(ctx, new(testArtists[i])) if err != nil { panic(err) } } // Associate artists with library 1 (default test library) - lr := NewLibraryRepository(ctx, conn) + lr := NewLibraryRepository(conn) for i := range testArtists { - err := lr.AddArtist(1, testArtists[i].ID) + err := lr.AddArtist(ctx, 1, testArtists[i].ID) if err != nil { panic(err) } @@ -247,7 +247,7 @@ var _ = BeforeSuite(func() { if a.AlbumArtistID == "" || !artistIDs[a.AlbumArtistID] { continue } - _, err := alr.executeSQL(squirrel.Insert("album_artists").SetMap(map[string]any{ + _, err := alr.executeSQL(ctx, squirrel.Insert("album_artists").SetMap(map[string]any{ "album_id": a.ID, "artist_id": a.AlbumArtistID, "role": "artist", @@ -258,17 +258,17 @@ var _ = BeforeSuite(func() { } } - mr := NewMediaFileRepository(ctx, conn) + mr := NewMediaFileRepository(conn) for i := range testSongs { - err := mr.Put(&testSongs[i]) + err := mr.Put(ctx, &testSongs[i]) if err != nil { panic(err) } } - rar := NewRadioRepository(ctx, conn) + rar := NewRadioRepository(conn) for i := range testRadios { - err := rar.Put(new(testRadios[i])) + err := rar.Put(ctx, new(testRadios[i])) if err != nil { panic(err) } @@ -287,19 +287,19 @@ var _ = BeforeSuite(func() { plsCool.AddMediaFilesByID([]string{"1004"}) testPlaylists = []*model.Playlist{&plsBest, &plsCool} - pr := NewPlaylistRepository(ctx, conn) + pr := NewPlaylistRepository(conn) for i := range testPlaylists { - err := pr.Put(testPlaylists[i]) + err := pr.Put(ctx, testPlaylists[i]) if err != nil { panic(err) } } // Prepare annotations - if err := arr.SetStar(true, artistBeatles.ID); err != nil { + if err := arr.SetStar(ctx, true, artistBeatles.ID); err != nil { panic(err) } - ar, err := arr.Get(artistBeatles.ID) + ar, err := arr.Get(ctx, artistBeatles.ID) if err != nil { panic(err) } @@ -310,10 +310,10 @@ var _ = BeforeSuite(func() { artistBeatles.StarredAt = ar.StarredAt testArtists[1] = artistBeatles - if err := alr.SetStar(true, albumRadioactivity.ID); err != nil { + if err := alr.SetStar(ctx, true, albumRadioactivity.ID); err != nil { panic(err) } - al, err := alr.Get(albumRadioactivity.ID) + al, err := alr.Get(ctx, albumRadioactivity.ID) if err != nil { panic(err) } @@ -324,10 +324,10 @@ var _ = BeforeSuite(func() { albumRadioactivity.StarredAt = al.StarredAt testAlbums[2] = albumRadioactivity - if err := mr.SetStar(true, songComeTogether.ID); err != nil { + if err := mr.SetStar(ctx, true, songComeTogether.ID); err != nil { panic(err) } - mf, err := mr.Get(songComeTogether.ID) + mf, err := mr.Get(ctx, songComeTogether.ID) if err != nil { panic(err) } @@ -335,9 +335,9 @@ var _ = BeforeSuite(func() { songComeTogether.StarredAt = mf.StarredAt testSongs[1] = songComeTogether - scrobbleRepo := NewScrobbleRepository(ctx, conn).(*scrobbleRepository) + scrobbleRepo := NewScrobbleRepository(conn).(*scrobbleRepository) for _, s := range scrobbles { - _, err := scrobbleRepo.executeSQL(squirrel.Insert("scrobbles").SetMap(map[string]any{ + _, err := scrobbleRepo.executeSQL(ctx, squirrel.Insert("scrobbles").SetMap(map[string]any{ "media_file_id": s.MediaFileID, "user_id": s.UserID, "submission_time": s.SubmissionTime, diff --git a/persistence/persistence_test.go b/persistence/persistence_test.go index e13f6a231..43d2e81ed 100644 --- a/persistence/persistence_test.go +++ b/persistence/persistence_test.go @@ -23,37 +23,37 @@ var _ = Describe("SQLStore", func() { Context("When block returns nil", func() { It("commits changes to the DB", func() { err := ds.WithTx(func(tx model.DataStore) error { - pl := tx.Player(ctx) - err := pl.Put(&model.Player{ID: "666", UserId: "userid"}) + pl := tx.Player() + err := pl.Put(ctx, &model.Player{ID: "666", UserId: "userid"}) Expect(err).ToNot(HaveOccurred()) - pr := tx.Property(ctx) - err = pr.Put("777", "value") + pr := tx.Property() + err = pr.Put(ctx, "777", "value") Expect(err).ToNot(HaveOccurred()) return nil }) Expect(err).ToNot(HaveOccurred()) - Expect(ds.Player(ctx).Get("666")).To(Equal(&model.Player{ID: "666", UserId: "userid", Username: "userid"})) - Expect(ds.Property(ctx).Get("777")).To(Equal("value")) + Expect(ds.Player().Get(ctx, "666")).To(Equal(&model.Player{ID: "666", UserId: "userid", Username: "userid"})) + Expect(ds.Property().Get(ctx, "777")).To(Equal("value")) }) }) Context("When block returns an error", func() { It("rollbacks changes to the DB", func() { err := ds.WithTx(func(tx model.DataStore) error { - pr := tx.Property(ctx) - err := pr.Put("999", "value") + pr := tx.Property() + err := pr.Put(ctx, "999", "value") Expect(err).ToNot(HaveOccurred()) // Will fail as it is missing the UserName - pl := tx.Player(ctx) - err = pl.Put(&model.Player{ID: "888"}) + pl := tx.Player() + err = pl.Put(ctx, &model.Player{ID: "888"}) Expect(err).To(HaveOccurred()) return err }) Expect(err).To(HaveOccurred()) - _, err = ds.Property(ctx).Get("999") + _, err = ds.Property().Get(ctx, "999") Expect(err).To(MatchError(model.ErrNotFound)) - _, err = ds.Player(ctx).Get("888") + _, err = ds.Player().Get(ctx, "888") Expect(err).To(MatchError(model.ErrNotFound)) }) }) @@ -70,7 +70,7 @@ var _ = Describe("SQLStore", func() { var attempts []bool err := ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { attempts = append(attempts, hasBusyRetry(ctx)) - Expect(tx.Property(ctx).Put("retry-key", "attempt")).To(Succeed()) + Expect(tx.Property().Put(ctx, "retry-key", "attempt")).To(Succeed()) if len(attempts) < 3 { return busy } @@ -78,7 +78,7 @@ var _ = Describe("SQLStore", func() { }) Expect(err).ToNot(HaveOccurred()) Expect(attempts).To(Equal([]bool{true, true, true})) - Expect(ds.Property(ctx).Get("retry-key")).To(Equal("attempt")) + Expect(ds.Property().Get(ctx, "retry-key")).To(Equal("attempt")) }) It("gives up after the last retry, which is not marked as retried", func() { @@ -116,15 +116,15 @@ var _ = Describe("SQLStore", func() { It("joins the enclosing transaction instead of opening another", func() { rollback := errors.New("rollback") err := ds.WithTx(func(tx model.DataStore) error { - Expect(tx.Property(ctx).Put("outer-key", "v")).To(Succeed()) + Expect(tx.Property().Put(ctx, "outer-key", "v")).To(Succeed()) Expect(tx.WithTxRetry(ctx, func(ctx context.Context, inner model.DataStore) error { - Expect(inner.Property(ctx).Get("outer-key")).To(Equal("v")) - return inner.Property(ctx).Put("inner-key", "v") + Expect(inner.Property().Get(ctx, "outer-key")).To(Equal("v")) + return inner.Property().Put(ctx, "inner-key", "v") })).To(Succeed()) return rollback }) Expect(err).To(MatchError(rollback)) - _, err = ds.Property(ctx).Get("inner-key") + _, err = ds.Property().Get(ctx, "inner-key") Expect(err).To(MatchError(model.ErrNotFound)) }) }) diff --git a/persistence/player_repository.go b/persistence/player_repository.go index a2114b060..e46a8d82d 100644 --- a/persistence/player_repository.go +++ b/persistence/player_repository.go @@ -13,9 +13,8 @@ type playerRepository struct { sqlRepository } -func NewPlayerRepository(ctx context.Context, db dbx.Builder) model.PlayerRepository { +func NewPlayerRepository(db dbx.Builder) model.PlayerRepository { r := &playerRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Player{}, map[string]filterFunc{ "name": containsFilter("player.name"), @@ -26,43 +25,43 @@ func NewPlayerRepository(ctx context.Context, db dbx.Builder) model.PlayerReposi return r } -func (r *playerRepository) Put(p *model.Player) error { - _, err := r.put(p.ID, p) +func (r *playerRepository) Put(ctx context.Context, p *model.Player) error { + _, err := r.put(ctx, p.ID, p) return err } -func (r *playerRepository) selectPlayer(options ...model.QueryOptions) SelectBuilder { - return r.newSelect(options...). +func (r *playerRepository) selectPlayer(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + return r.newSelect(ctx, options...). Columns("player.*"). Join("user ON player.user_id = user.id"). Columns("user.user_name username") } -func (r *playerRepository) Get(id string) (*model.Player, error) { - sel := r.selectPlayer().Where(Eq{"player.id": id}) +func (r *playerRepository) Get(ctx context.Context, id string) (*model.Player, error) { + sel := r.selectPlayer(ctx).Where(Eq{"player.id": id}) var res model.Player - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *playerRepository) FindMatch(userId, client, userAgent string) (*model.Player, error) { - sel := r.selectPlayer().Where(And{ +func (r *playerRepository) FindMatch(ctx context.Context, userId, client, userAgent string) (*model.Player, error) { + sel := r.selectPlayer(ctx).Where(And{ Eq{"client": client}, Eq{"user_agent": userAgent}, Eq{"user_id": userId}, }) var res model.Player - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *playerRepository) newRestSelect(options ...model.QueryOptions) SelectBuilder { - s := r.selectPlayer(options...) - return s.Where(r.addRestriction()) +func (r *playerRepository) newRestSelect(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + s := r.selectPlayer(ctx, options...) + return s.Where(r.addRestriction(ctx)) } -func (r *playerRepository) CountByClient(options ...model.QueryOptions) (map[string]int64, error) { - sel := r.newSelect(options...). +func (r *playerRepository) CountByClient(ctx context.Context, options ...model.QueryOptions) (map[string]int64, error) { + sel := r.newSelect(ctx, options...). Columns( "case when client = 'NavidromeUI' then name else client end as player", "count(*) as count", @@ -71,7 +70,7 @@ func (r *playerRepository) CountByClient(options ...model.QueryOptions) (map[str Player string Count int64 } - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } @@ -82,63 +81,54 @@ func (r *playerRepository) CountByClient(options ...model.QueryOptions) (map[str return counts, nil } -func (r *playerRepository) CountAll(options ...model.QueryOptions) (int64, error) { - return r.count(r.newRestSelect(), options...) +func (r *playerRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + return r.count(ctx, r.newRestSelect(ctx), options...) } -func (r *playerRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *playerRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *playerRepository) Read(id string) (any, error) { - sel := r.newRestSelect().Where(Eq{"player.id": id}) +func (r *playerRepository) Read(ctx context.Context, id string) (*model.Player, error) { + sel := r.newRestSelect(ctx).Where(Eq{"player.id": id}) var res model.Player - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *playerRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - sel := r.newRestSelect(r.parseRestOptions(r.ctx, options...)) +func (r *playerRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Player, error) { + sel := r.newRestSelect(ctx, r.parseRestOptions(ctx, options...)) res := model.Players{} - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) return res, err } -func (r *playerRepository) EntityName() string { - return "player" -} - -func (r *playerRepository) NewInstance() any { - return &model.Player{} -} - // isPermitted authorizes creating a new record, based on the owner declared in the request body. // This is only safe for inserts: there is no stored row yet, and a non-admin may only create a // player they own. Updates must not use this (the body owner is attacker-controlled); they go // through updateOwned, which authorizes against the persisted user_id in the WHERE clause. -func (r *playerRepository) isPermitted(p *model.Player) bool { - u := loggedUser(r.ctx) +func (r *playerRepository) isPermitted(ctx context.Context, p *model.Player) bool { + u := loggedUser(ctx) return u.IsAdmin || p.UserId == u.ID } -func (r *playerRepository) Save(entity any) (string, error) { - t := entity.(*model.Player) - if !r.isPermitted(t) { +func (r *playerRepository) Save(ctx context.Context, t *model.Player) (string, error) { + if !r.isPermitted(ctx, t) { return "", rest.ErrPermissionDenied } - return r.put("", t) // Save only creates; edits go through the owner-scoped Update + return r.put(ctx, "", t) // Save only creates; edits go through the owner-scoped Update } -func (r *playerRepository) Update(id string, entity any, cols ...string) error { - t := entity.(*model.Player) +func (r *playerRepository) Update(ctx context.Context, id string, entity model.Player, cols ...string) error { + t := &entity t.ID = id - return r.updateOwned(id, t, cols...) + return r.updateOwned(ctx, id, t, cols...) } -func (r *playerRepository) Delete(id string) error { - return r.deleteOwned(id) +func (r *playerRepository) Delete(ctx context.Context, ids ...string) error { + return r.deleteOwnedAll(ctx, ids...) } var _ model.PlayerRepository = (*playerRepository)(nil) -var _ rest.Repository = (*playerRepository)(nil) -var _ rest.Persistable = (*playerRepository)(nil) +var _ rest.Repository[model.Player] = (*playerRepository)(nil) +var _ rest.Persistable[model.Player] = (*playerRepository)(nil) diff --git a/persistence/player_repository_test.go b/persistence/player_repository_test.go index 4a7701ac2..f12f3e74e 100644 --- a/persistence/player_repository_test.go +++ b/persistence/player_repository_test.go @@ -15,6 +15,7 @@ import ( var _ = Describe("PlayerRepository", func() { var adminRepo *playerRepository var database *dbx.DB + var ctx context.Context var ( adminPlayer1 = model.Player{ID: "1", Name: "NavidromeUI [Firefox/Linux]", UserAgent: "Firefox/Linux", UserId: adminUser.ID, Username: adminUser.UserName, Client: "NavidromeUI", IP: "127.0.0.1", ReportRealPath: true, ScrobbleEnabled: true} @@ -25,77 +26,69 @@ var _ = Describe("PlayerRepository", func() { ) BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, adminUser) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) database = GetDBXBuilder() - adminRepo = NewPlayerRepository(ctx, database).(*playerRepository) + adminRepo = NewPlayerRepository(database).(*playerRepository) for idx := range players { - err := adminRepo.Put(&players[idx]) + err := adminRepo.Put(ctx, &players[idx]) Expect(err).To(BeNil()) } }) AfterEach(func() { - items, err := adminRepo.ReadAll() + players, err := adminRepo.ReadAll(ctx) Expect(err).To(BeNil()) - players, ok := items.(model.Players) - Expect(ok).To(BeTrue()) for i := range players { - err = adminRepo.Delete(players[i].ID) + err = adminRepo.Delete(ctx, players[i].ID) Expect(err).To(BeNil()) } }) - Describe("EntityName", func() { - It("returns the right name", func() { - Expect(adminRepo.EntityName()).To(Equal("player")) - }) - }) - Describe("FindMatch", func() { It("finds existing match", func() { - player, err := adminRepo.FindMatch(adminUser.ID, "NavidromeUI", "Firefox/Linux") + player, err := adminRepo.FindMatch(ctx, adminUser.ID, "NavidromeUI", "Firefox/Linux") Expect(err).To(BeNil()) Expect(*player).To(Equal(adminPlayer1)) }) It("doesn't find bad match", func() { - _, err := adminRepo.FindMatch(regularUser.ID, "NavidromeUI", "Firefox/Linux") + _, err := adminRepo.FindMatch(ctx, regularUser.ID, "NavidromeUI", "Firefox/Linux") Expect(err).To(Equal(model.ErrNotFound)) }) }) Describe("Get", func() { It("Gets an existing item from user", func() { - player, err := adminRepo.Get(adminPlayer1.ID) + player, err := adminRepo.Get(ctx, adminPlayer1.ID) Expect(err).To(BeNil()) Expect(*player).To(Equal(adminPlayer1)) }) It("Gets an existing item from another user", func() { - player, err := adminRepo.Get(regularPlayer.ID) + player, err := adminRepo.Get(ctx, regularPlayer.ID) Expect(err).To(BeNil()) Expect(*player).To(Equal(regularPlayer)) }) It("does not get nonexistent item", func() { - _, err := adminRepo.Get("i don't exist") + _, err := adminRepo.Get(ctx, "i don't exist") Expect(err).To(Equal(model.ErrNotFound)) }) }) DescribeTableSubtree("per context", func(admin bool, players model.Players, userPlayer model.Player, otherPlayer model.Player) { var repo *playerRepository + var repoCtx context.Context BeforeEach(func() { + repoCtx = ctx if admin { repo = adminRepo } else { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, regularUser) - repo = NewPlayerRepository(ctx, database).(*playerRepository) + repoCtx = request.WithUser(ctx, regularUser) + repo = NewPlayerRepository(database).(*playerRepository) } }) @@ -103,7 +96,7 @@ var _ = Describe("PlayerRepository", func() { Describe("Count", func() { It("should return all", func() { - count, err := repo.Count() + count, err := repo.Count(repoCtx) Expect(err).To(BeNil()) Expect(count).To(Equal(baseCount)) }) @@ -111,53 +104,53 @@ var _ = Describe("PlayerRepository", func() { Describe("Delete", func() { It("deletes a player owned by the current user", func() { - err := repo.Delete(userPlayer.ID) + err := repo.Delete(repoCtx, userPlayer.ID) Expect(err).To(BeNil()) - count, err := repo.Count() + count, err := repo.Count(repoCtx) Expect(err).To(BeNil()) Expect(count).To(Equal(baseCount - 1)) - _, err = repo.Get(userPlayer.ID) + _, err = repo.Get(repoCtx, userPlayer.ID) Expect(err).To(Equal(model.ErrNotFound)) }) It("does not delete another user's player when not admin", func() { - err := repo.Delete(otherPlayer.ID) + err := repo.Delete(repoCtx, otherPlayer.ID) if admin { // Admins may delete any player. Expect(err).To(BeNil()) - Expect(repo.Count()).To(Equal(baseCount - 1)) - _, err = repo.Get(otherPlayer.ID) + Expect(repo.Count(repoCtx)).To(Equal(baseCount - 1)) + _, err = repo.Get(repoCtx, otherPlayer.ID) Expect(err).To(Equal(model.ErrNotFound)) } else { // The ownership-restricted delete matches no owned row, so it reports // permission-denied and leaves the other user's player untouched. Expect(err).To(Equal(rest.ErrPermissionDenied)) - Expect(repo.Count()).To(Equal(baseCount)) - item, err := repo.Get(otherPlayer.ID) + Expect(repo.Count(repoCtx)).To(Equal(baseCount)) + item, err := repo.Get(repoCtx, otherPlayer.ID) Expect(err).To(BeNil()) Expect(*item).To(Equal(otherPlayer)) } }) It("returns not-found for a nonexistent player", func() { - err := repo.Delete("i don't exist") + err := repo.Delete(repoCtx, "i don't exist") Expect(err).To(Equal(rest.ErrNotFound)) - Expect(repo.Count()).To(Equal(baseCount)) + Expect(repo.Count(repoCtx)).To(Equal(baseCount)) }) }) Describe("Read", func() { It("can read from current user", func() { - player, err := repo.Read(userPlayer.ID) + player, err := repo.Read(repoCtx, userPlayer.ID) Expect(err).To(BeNil()) Expect(player).To(Equal(&userPlayer)) }) It("can read from other user or fail if not admin", func() { - player, err := repo.Read(otherPlayer.ID) + player, err := repo.Read(repoCtx, otherPlayer.ID) if admin { Expect(err).To(BeNil()) Expect(player).To(Equal(&otherPlayer)) @@ -167,16 +160,16 @@ var _ = Describe("PlayerRepository", func() { }) It("does not get nonexistent item", func() { - _, err := repo.Read("i don't exist") + _, err := repo.Read(repoCtx, "i don't exist") Expect(err).To(Equal(model.ErrNotFound)) }) }) Describe("ReadAll", func() { It("should get all items", func() { - data, err := repo.ReadAll() + data, err := repo.ReadAll(repoCtx) Expect(err).To(BeNil()) - Expect(data).To(Equal(players)) + Expect(model.Players(data)).To(Equal(players)) }) }) @@ -185,7 +178,7 @@ var _ = Describe("PlayerRepository", func() { clone := player clone.ID = "" clone.IP = "192.168.1.1" - id, err := repo.Save(&clone) + id, err := repo.Save(repoCtx, &clone) if clone.UserId == "" { Expect(err).To(HaveOccurred()) @@ -197,11 +190,11 @@ var _ = Describe("PlayerRepository", func() { Expect(id).ToNot(BeEmpty()) } - count, err := repo.Count() + count, err := repo.Count(repoCtx) Expect(err).To(BeNil()) clone.ID = id - newItem, err := repo.Get(id) + newItem, err := repo.Get(repoCtx, id) if clone.UserId == "" { Expect(count).To(Equal(baseCount)) @@ -223,7 +216,7 @@ var _ = Describe("PlayerRepository", func() { clone := player clone.IP = "192.168.1.1" clone.MaxBitRate = 10000 - err := repo.Update(clone.ID, &clone, "ip") + err := repo.Update(repoCtx, clone.ID, clone, "ip") if player.UserId == "" { Expect(err).To(HaveOccurred()) @@ -238,7 +231,7 @@ var _ = Describe("PlayerRepository", func() { } clone.MaxBitRate = player.MaxBitRate - newItem, err := repo.Get(clone.ID) + newItem, err := repo.Get(repoCtx, clone.ID) if player.UserId == "" { Expect(err).To(Equal(model.ErrNotFound)) @@ -260,11 +253,11 @@ var _ = Describe("PlayerRepository", func() { Describe("Ownership enforcement (cross-tenant write protection)", func() { var regularRepo *playerRepository + var regularCtx context.Context BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, regularUser) - regularRepo = NewPlayerRepository(ctx, database).(*playerRepository) + regularCtx = request.WithUser(ctx, regularUser) + regularRepo = NewPlayerRepository(database).(*playerRepository) }) It("does not let a regular user hijack another user's player by spoofing userId in the body", func() { @@ -279,11 +272,11 @@ var _ = Describe("PlayerRepository", func() { // The ownership-restricted update matches no row owned by the attacker, so the write // targets nothing and reports permission-denied rather than overwriting the victim's row. - err := regularRepo.Update(adminPlayer1.ID, &spoofed, "name", "user_id", "max_bit_rate") + err := regularRepo.Update(regularCtx, adminPlayer1.ID, spoofed, "name", "user_id", "max_bit_rate") Expect(err).To(Equal(rest.ErrPermissionDenied)) // The victim's player must remain untouched. - stored, err := adminRepo.Get(adminPlayer1.ID) + stored, err := adminRepo.Get(ctx, adminPlayer1.ID) Expect(err).To(BeNil()) Expect(*stored).To(Equal(adminPlayer1)) }) @@ -296,15 +289,15 @@ var _ = Describe("PlayerRepository", func() { ReportRealPath: true, } - id, err := regularRepo.Save(&spoofed) + id, err := regularRepo.Save(regularCtx, &spoofed) Expect(err).To(BeNil()) Expect(id).ToNot(Equal(adminPlayer1.ID)) - stored, err := adminRepo.Get(adminPlayer1.ID) + stored, err := adminRepo.Get(ctx, adminPlayer1.ID) Expect(err).To(BeNil()) Expect(*stored).To(Equal(adminPlayer1)) - created, err := adminRepo.Get(id) + created, err := adminRepo.Get(ctx, id) Expect(err).To(BeNil()) Expect(created.UserId).To(Equal(regularUser.ID)) }) @@ -316,11 +309,11 @@ var _ = Describe("PlayerRepository", func() { reassign.UserId = adminUser.ID reassign.Name = "given-away" - err := regularRepo.Update(regularPlayer.ID, &reassign, "name", "user_id") + err := regularRepo.Update(regularCtx, regularPlayer.ID, reassign, "name", "user_id") Expect(err).To(BeNil()) // Ownership must not have changed. - stored, err := adminRepo.Get(regularPlayer.ID) + stored, err := adminRepo.Get(ctx, regularPlayer.ID) Expect(err).To(BeNil()) Expect(stored.UserId).To(Equal(regularUser.ID)) }) @@ -331,11 +324,11 @@ var _ = Describe("PlayerRepository", func() { reassign.UserId = adminUser.ID reassign.Name = "admin-renamed" - err := adminRepo.Update(regularPlayer.ID, &reassign, "name", "user_id") + err := adminRepo.Update(regularCtx, regularPlayer.ID, reassign, "name", "user_id") Expect(err).To(BeNil()) // The name change applies, but ownership must not have moved. - stored, err := adminRepo.Get(regularPlayer.ID) + stored, err := adminRepo.Get(ctx, regularPlayer.ID) Expect(err).To(BeNil()) Expect(stored.Name).To(Equal("admin-renamed")) Expect(stored.UserId).To(Equal(regularUser.ID)) @@ -345,10 +338,10 @@ var _ = Describe("PlayerRepository", func() { update := regularPlayer update.Name = "renamed-by-owner" - err := regularRepo.Update(regularPlayer.ID, &update, "name") + err := regularRepo.Update(regularCtx, regularPlayer.ID, update, "name") Expect(err).To(BeNil()) - stored, err := adminRepo.Get(regularPlayer.ID) + stored, err := adminRepo.Get(ctx, regularPlayer.ID) Expect(err).To(BeNil()) Expect(stored.Name).To(Equal("renamed-by-owner")) Expect(stored.UserId).To(Equal(regularUser.ID)) @@ -356,7 +349,7 @@ var _ = Describe("PlayerRepository", func() { It("returns not found when updating a nonexistent player", func() { ghost := model.Player{ID: "does-not-exist", Name: "ghost", UserId: regularUser.ID} - err := regularRepo.Update("does-not-exist", &ghost, "name") + err := regularRepo.Update(regularCtx, "does-not-exist", ghost, "name") Expect(err).To(Equal(rest.ErrNotFound)) }) }) diff --git a/persistence/playlist_repository.go b/persistence/playlist_repository.go index 2afef9f45..41b75266c 100644 --- a/persistence/playlist_repository.go +++ b/persistence/playlist_repository.go @@ -50,9 +50,8 @@ func (p dbPlaylist) PostMapArgs(args map[string]any) error { return nil } -func NewPlaylistRepository(ctx context.Context, db dbx.Builder) model.PlaylistRepository { +func NewPlaylistRepository(db dbx.Builder) model.PlaylistRepository { r := &playlistRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Playlist{}, map[string]filterFunc{ "id": idFilter("playlist"), @@ -81,8 +80,8 @@ func smartPlaylistFilter(string, any) Sqlizer { } } -func (r *playlistRepository) userFilter() Sqlizer { - user := loggedUser(r.ctx) +func (r *playlistRepository) userFilter(ctx context.Context) Sqlizer { + user := loggedUser(ctx) if user.IsAdmin { return And{} } @@ -92,29 +91,29 @@ func (r *playlistRepository) userFilter() Sqlizer { } } -func (r *playlistRepository) CountAll(options ...model.QueryOptions) (int64, error) { - query := Select().Where(r.userFilter()) +func (r *playlistRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + query := Select().Where(r.userFilter(ctx)) if filtersNeedAnnotation(r.applyFilters(query, options...)) { - query = r.withAnnotation(query, "playlist.id") + query = r.withAnnotation(ctx, query, "playlist.id") } - return r.count(query, options...) + return r.count(ctx, query, options...) } -func (r *playlistRepository) Exists(id string) (bool, error) { - return r.exists(And{Eq{"id": id}, r.userFilter()}) +func (r *playlistRepository) Exists(ctx context.Context, id string) (bool, error) { + return r.exists(ctx, And{Eq{"id": id}, r.userFilter(ctx)}) } -func (r *playlistRepository) Delete(id string) error { - return r.delete(And{Eq{"id": id}, r.userFilter()}) +func (r *playlistRepository) Delete(ctx context.Context, ids ...string) error { + return r.delete(ctx, And{Eq{"id": ids}, r.userFilter(ctx)}) } -func (r *playlistRepository) Put(p *model.Playlist, cols ...string) error { +func (r *playlistRepository) Put(ctx context.Context, p *model.Playlist, cols ...string) error { pls := dbPlaylist{Playlist: *p} if len(cols) > 0 { if pls.ID == "" { return errors.New("playlist id is required for partial update") } - _, err := r.put(pls.ID, pls, cols...) + _, err := r.put(ctx, pls.ID, pls, cols...) return err } isNew := pls.ID == "" @@ -123,7 +122,7 @@ func (r *playlistRepository) Put(p *model.Playlist, cols ...string) error { } pls.UpdatedAt = time.Now() - id, err := r.put(pls.ID, pls) + id, err := r.put(ctx, pls.ID, pls) if err != nil { return err } @@ -135,48 +134,48 @@ func (r *playlistRepository) Put(p *model.Playlist, cols ...string) error { } // Only update tracks if they were specified if len(pls.Tracks) > 0 { - return r.updateTracks(id, p.MediaFiles()) + return r.updateTracks(ctx, id, p.MediaFiles()) } pls.ID = id // r.put assigns the generated id to p, not to this copy if isNew { // Even a trackless new playlist has art to find (an imported m3u can carry an // ExternalImageURL); an update landing here changed only metadata, so leave its cover be. - r.enqueueCoverRebuild(id) + r.enqueueCoverRebuild(ctx, id) } - return r.refreshCounters(&pls.Playlist) + return r.refreshCounters(ctx, &pls.Playlist) } -func (r *playlistRepository) Get(id string) (*model.Playlist, error) { - return r.findBy(And{Eq{"playlist.id": id}, r.userFilter()}) +func (r *playlistRepository) Get(ctx context.Context, id string) (*model.Playlist, error) { + return r.findBy(ctx, And{Eq{"playlist.id": id}, r.userFilter(ctx)}) } -func (r *playlistRepository) GetWithTracks(id string, refreshSmartPlaylist, includeMissing bool) (*model.Playlist, error) { - pls, err := r.Get(id) +func (r *playlistRepository) GetWithTracks(ctx context.Context, id string, refreshSmartPlaylist, includeMissing bool) (*model.Playlist, error) { + pls, err := r.Get(ctx, id) if err != nil { return nil, err } if refreshSmartPlaylist { - r.refreshSmartPlaylist(pls) + r.refreshSmartPlaylist(ctx, pls) } - tracks, err := r.loadTracks(Select().From("playlist_tracks"). + tracks, err := r.loadTracks(ctx, Select().From("playlist_tracks"). Where(Eq{"missing": false}). OrderBy("playlist_tracks.id"), id) if err != nil { - log.Error(r.ctx, "Error loading playlist tracks ", "playlist", pls.Name, "id", pls.ID, err) + log.Error(ctx, "Error loading playlist tracks ", "playlist", pls.Name, "id", pls.ID, err) return nil, err } pls.SetTracks(tracks) return pls, nil } -func (r *playlistRepository) FindByPath(path string) (*model.Playlist, error) { - return r.findBy(Eq{"path": path}) +func (r *playlistRepository) FindByPath(ctx context.Context, path string) (*model.Playlist, error) { + return r.findBy(ctx, Eq{"path": path}) } -func (r *playlistRepository) findBy(sql Sqlizer) (*model.Playlist, error) { - sel := r.selectPlaylist().Where(sql) +func (r *playlistRepository) findBy(ctx context.Context, sql Sqlizer) (*model.Playlist, error) { + sel := r.selectPlaylist(ctx).Where(sql) var pls []dbPlaylist - err := r.queryAll(sel, &pls) + err := r.queryAll(ctx, sel, &pls) if err != nil { return nil, err } @@ -185,19 +184,19 @@ func (r *playlistRepository) findBy(sql Sqlizer) (*model.Playlist, error) { } list := model.Playlists{pls[0].Playlist} - r.hydrateArtwork(list) + r.hydrateArtwork(ctx, list) return &list[0], nil } -func (r *playlistRepository) hydrateArtwork(playlists model.Playlists) { - hydrateItems(r.ctx, r.db, model.KindPlaylistArtwork, playlists, +func (r *playlistRepository) hydrateArtwork(ctx context.Context, playlists model.Playlists) { + hydrateItems(ctx, r.db, model.KindPlaylistArtwork, playlists, func(p *model.Playlist) (string, *model.ItemImage) { return p.ID, &p.ItemImage }) } -func (r *playlistRepository) GetAll(options ...model.QueryOptions) (model.Playlists, error) { - sel := r.selectPlaylist(options...).Where(r.userFilter()) +func (r *playlistRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Playlists, error) { + sel := r.selectPlaylist(ctx, options...).Where(r.userFilter(ctx)) var res []dbPlaylist - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } @@ -205,42 +204,42 @@ func (r *playlistRepository) GetAll(options ...model.QueryOptions) (model.Playli for i, p := range res { playlists[i] = p.Playlist } - r.hydrateArtwork(playlists) + r.hydrateArtwork(ctx, playlists) return playlists, err } // getAllIDs returns the IDs of GetAll's row set, skipping its per-row processing. -func (r *playlistRepository) getAllIDs(options ...model.QueryOptions) ([]string, error) { +func (r *playlistRepository) getAllIDs(ctx context.Context, options ...model.QueryOptions) ([]string, error) { // Joins a projection of user, not the table: its name/created_at columns would make an ORDER BY // on the playlist's own ambiguous. - sq := r.newSelect(options...).Columns("playlist.id", "user.user_name as owner_name"). - Join("(select id, user_name from user) user on user.id = owner_id").Where(r.userFilter()) + sq := r.newSelect(ctx, options...).Columns("playlist.id", "user.user_name as owner_name"). + Join("(select id, user_name from user) user on user.id = owner_id").Where(r.userFilter(ctx)) if filtersNeedAnnotation(sq) { - sq = r.withAnnotation(sq, "playlist.id") + sq = r.withAnnotation(ctx, sq, "playlist.id") } ids := []string{} - err := r.queryAllSlice(sq, &ids) + err := r.queryAllSlice(ctx, sq, &ids) return ids, err } -func (r *playlistRepository) GetCursor(options ...model.QueryOptions) (model.PlaylistCursor, error) { +func (r *playlistRepository) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.PlaylistCursor, error) { // Both passes apply userFilter, so a visibility change between them cannot widen the cursor. - ids, err := r.getAllIDs(options...) + ids, err := r.getAllIDs(ctx, options...) if err != nil { return nil, err } opts := chunkOptions(options, "playlist.id") return model.PlaylistCursor(streamByIDs(ids, func(chunk []string) (model.Playlists, error) { - return r.GetAll(opts(chunk)) + return r.GetAll(ctx, opts(chunk)) })), nil } -func (r *playlistRepository) GetPlaylists(mediaFileId string) (model.Playlists, error) { - sel := r.selectPlaylist(model.QueryOptions{Sort: "name"}). +func (r *playlistRepository) GetPlaylists(ctx context.Context, mediaFileId string) (model.Playlists, error) { + sel := r.selectPlaylist(ctx, model.QueryOptions{Sort: "name"}). Join("playlist_tracks on playlist.id = playlist_tracks.playlist_id"). - Where(And{Eq{"playlist_tracks.media_file_id": mediaFileId}, r.userFilter()}) + Where(And{Eq{"playlist_tracks.media_file_id": mediaFileId}, r.userFilter(ctx)}) var res []dbPlaylist - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { if errors.Is(err, model.ErrNotFound) { return model.Playlists{}, nil @@ -251,40 +250,40 @@ func (r *playlistRepository) GetPlaylists(mediaFileId string) (model.Playlists, for i, p := range res { playlists[i] = p.Playlist } - r.hydrateArtwork(playlists) + r.hydrateArtwork(ctx, playlists) return playlists, nil } -func (r *playlistRepository) selectPlaylist(options ...model.QueryOptions) SelectBuilder { - sel := r.newSelect(options...).Join("user on user.id = owner_id"). +func (r *playlistRepository) selectPlaylist(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + sel := r.newSelect(ctx, options...).Join("user on user.id = owner_id"). Columns(r.tableName+".*", "user.user_name as owner_name") - return r.withAnnotation(sel, r.tableName+".id") + return r.withAnnotation(ctx, sel, r.tableName+".id") } -func (r *playlistRepository) updateTracks(id string, tracks model.MediaFiles) error { +func (r *playlistRepository) updateTracks(ctx context.Context, id string, tracks model.MediaFiles) error { ids := make([]string, len(tracks)) for i := range tracks { ids[i] = tracks[i].ID } - return r.updatePlaylist(id, ids) + return r.updatePlaylist(ctx, id, ids) } -func (r *playlistRepository) updatePlaylist(playlistId string, mediaFileIds []string) error { +func (r *playlistRepository) updatePlaylist(ctx context.Context, playlistId string, mediaFileIds []string) error { // Remove old tracks del := Delete("playlist_tracks").Where(Eq{"playlist_id": playlistId}) - _, err := r.executeSQL(del) + _, err := r.executeSQL(ctx, del) if err != nil { return err } - _, err = r.addTracks(playlistId, 1, mediaFileIds) + _, err = r.addTracks(ctx, playlistId, 1, mediaFileIds) return err } // addTracks is the only path that writes playlist_tracks rows (smart playlists aside), so it owns // the library check: every caller, including a full replace through Put, goes through it. -func (r *playlistRepository) addTracks(playlistId string, startingPos int, mediaFileIds []string) (int, error) { - mediaFileIds, err := r.keepAccessible(mediaFileIds) +func (r *playlistRepository) addTracks(ctx context.Context, playlistId string, startingPos int, mediaFileIds []string) (int, error) { + mediaFileIds, err := r.keepAccessible(ctx, mediaFileIds) if err != nil { return 0, err } @@ -297,26 +296,26 @@ func (r *playlistRepository) addTracks(playlistId string, startingPos int, media ins = ins.Values(playlistId, t, pos) pos++ } - if _, err := r.executeSQL(ins); err != nil { + if _, err := r.executeSQL(ctx, ins); err != nil { return 0, err } } - r.enqueueCoverRebuild(playlistId) - return len(mediaFileIds), r.refreshCounters(&model.Playlist{ID: playlistId}) + r.enqueueCoverRebuild(ctx, playlistId) + return len(mediaFileIds), r.refreshCounters(ctx, &model.Playlist{ID: playlistId}) } // keepAccessible drops ids the caller cannot read, preserving order and duplicates. Chunked // because callers pass unbounded id lists (M3U import), well past SQLITE_MAX_VARIABLE_NUMBER. -func (r *playlistRepository) keepAccessible(mediaFileIds []string) ([]string, error) { - if visible, err := r.visibleLibraryIDs(); err == nil && r.userSeesAllLibraries(visible) { +func (r *playlistRepository) keepAccessible(ctx context.Context, mediaFileIds []string) ([]string, error) { + if visible, err := r.visibleLibraryIDs(ctx); err == nil && r.userSeesAllLibraries(ctx, visible) { return mediaFileIds, nil } accessible := make(map[string]struct{}, len(mediaFileIds)) for chunk := range slices.Chunk(slice.Unique(mediaFileIds), 200) { - sq := r.applyLibraryFilter(Select("id").From("media_file").Where(Eq{"id": chunk}), "media_file") + sq := r.applyLibraryFilter(ctx, Select("id").From("media_file").Where(Eq{"id": chunk}), "media_file") var found []string - if err := r.queryAllSlice(sq, &found); err != nil { + if err := r.queryAllSlice(ctx, sq, &found); err != nil { return nil, err } for _, id := range found { @@ -330,7 +329,7 @@ func (r *playlistRepository) keepAccessible(mediaFileIds []string) ([]string, er } // refreshCounters updates total playlist duration, size and count -func (r *playlistRepository) refreshCounters(pls *model.Playlist) error { +func (r *playlistRepository) refreshCounters(ctx context.Context, pls *model.Playlist) error { statsSql := Select( "coalesce(sum(duration), 0) as duration", "coalesce(sum(size), 0) as size", @@ -340,7 +339,7 @@ func (r *playlistRepository) refreshCounters(pls *model.Playlist) error { Join("playlist_tracks f on f.media_file_id = media_file.id"). Where(Eq{"playlist_id": pls.ID}) var res struct{ Duration, Size, Count float32 } - err := r.queryOne(statsSql, &res) + err := r.queryOne(ctx, statsSql, &res) if err != nil { return err } @@ -353,7 +352,7 @@ func (r *playlistRepository) refreshCounters(pls *model.Playlist) error { Set("song_count", res.Count). Set("updated_at", now). Where(Eq{"id": pls.ID}) - _, err = r.executeSQL(upd) + _, err = r.executeSQL(ctx, upd) if err != nil { return err } @@ -366,18 +365,18 @@ func (r *playlistRepository) refreshCounters(pls *model.Playlist) error { // enqueueCoverRebuild re-resolves the generated 2x2 grid. Call it only when the track set changes: // the grid samples albums at random, so rebuilding after a mere rename would change the cover. -func (r *playlistRepository) enqueueCoverRebuild(id string) { +func (r *playlistRepository) enqueueCoverRebuild(ctx context.Context, id string) { item := model.ArtworkQueueItem{ItemKind: model.KindPlaylistArtwork.Prefix(), ItemID: id, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityScan} - if err := NewArtworkQueueRepository(r.ctx, r.db).Enqueue(item); err != nil { - log.Warn(r.ctx, "could not enqueue playlist artwork after content change", "id", id, err) + if err := NewArtworkQueueRepository(r.db).Enqueue(ctx, item); err != nil { + log.Warn(ctx, "could not enqueue playlist artwork after content change", "id", id, err) } } // tracksQuery is shared by loadTracks and GetCursor, so both hydrate rows identically. -func (r *playlistRepository) tracksQuery(query SelectBuilder, id string) SelectBuilder { - query = r.applyLibraryFilter(query, "f") - userID := loggedUser(r.ctx).ID +func (r *playlistRepository) tracksQuery(ctx context.Context, query SelectBuilder, id string) SelectBuilder { + query = r.applyLibraryFilter(ctx, query, "f") + userID := loggedUser(ctx).ID return query. Columns( "coalesce(starred, 0) as starred", @@ -400,56 +399,47 @@ func (r *playlistRepository) tracksQuery(query SelectBuilder, id string) SelectB Where(Eq{"playlist_id": id}) } -func (r *playlistRepository) loadTracks(query SelectBuilder, id string) (model.PlaylistTracks, error) { +func (r *playlistRepository) loadTracks(ctx context.Context, query SelectBuilder, id string) (model.PlaylistTracks, error) { tracks := dbPlaylistTracks{} - err := r.queryAll(r.tracksQuery(query, id), &tracks) + err := r.queryAll(ctx, r.tracksQuery(ctx, query, id), &tracks) if err != nil { return nil, err } res := tracks.toModels() - hydratePlaylistTrackArtwork(r.ctx, r.db, res) + hydratePlaylistTrackArtwork(ctx, r.db, res) return res, err } -func (r *playlistRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *playlistRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *playlistRepository) Read(id string) (any, error) { - return r.Get(id) +func (r *playlistRepository) Read(ctx context.Context, id string) (*model.Playlist, error) { + return r.Get(ctx, id) } -func (r *playlistRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) +func (r *playlistRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Playlist, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *playlistRepository) EntityName() string { - return "playlist" -} - -func (r *playlistRepository) NewInstance() any { - return &model.Playlist{} -} - -func (r *playlistRepository) Save(entity any) (string, error) { - pls := entity.(*model.Playlist) +func (r *playlistRepository) Save(ctx context.Context, pls *model.Playlist) (string, error) { pls.ID = "" // Force new creation - err := r.Put(pls) + err := r.Put(ctx, pls) if err != nil { return "", err } return pls.ID, err } -func (r *playlistRepository) Update(id string, entity any, cols ...string) error { - pls := dbPlaylist{Playlist: *entity.(*model.Playlist)} +func (r *playlistRepository) Update(ctx context.Context, id string, entity model.Playlist, cols ...string) error { + pls := dbPlaylist{Playlist: entity} pls.ID = id pls.UpdatedAt = time.Now() - _, err := r.put(id, pls, append(cols, "updatedAt")...) + _, err := r.put(ctx, id, pls, append(cols, "updatedAt")...) return err } -func (r *playlistRepository) removeOrphans() error { +func (r *playlistRepository) removeOrphans(ctx context.Context) error { sel := Select("playlist_tracks.playlist_id as id", "p.name").From("playlist_tracks"). Join("playlist p on playlist_tracks.playlist_id = p.id"). LeftJoin("media_file mf on playlist_tracks.media_file_id = mf.id"). @@ -457,25 +447,25 @@ func (r *playlistRepository) removeOrphans() error { GroupBy("playlist_tracks.playlist_id") var pls []struct{ Id, Name string } - err := r.queryAll(sel, &pls) + err := r.queryAll(ctx, sel, &pls) if err != nil { return fmt.Errorf("fetching playlists with orphan tracks: %w", err) } for _, pl := range pls { - log.Debug(r.ctx, "Cleaning-up orphan tracks from playlist", "id", pl.Id, "name", pl.Name) + log.Debug(ctx, "Cleaning-up orphan tracks from playlist", "id", pl.Id, "name", pl.Name) del := Delete("playlist_tracks").Where(And{ ConcatExpr("media_file_id not in (select id from media_file)"), Eq{"playlist_id": pl.Id}, }) - n, err := r.executeSQL(del) + n, err := r.executeSQL(ctx, del) if n == 0 || err != nil { return fmt.Errorf("deleting orphan tracks from playlist %s: %w", pl.Name, err) } - log.Debug(r.ctx, "Deleted tracks, now reordering", "id", pl.Id, "name", pl.Name, "deleted", n) + log.Debug(ctx, "Deleted tracks, now reordering", "id", pl.Id, "name", pl.Name, "deleted", n) // Renumber the playlist if any track was removed - if err := r.renumber(pl.Id); err != nil { + if err := r.renumber(ctx, pl.Id); err != nil { return fmt.Errorf("renumbering playlist %s: %w", pl.Name, err) } } @@ -485,9 +475,9 @@ func (r *playlistRepository) removeOrphans() error { // renumber updates the position of all tracks in the playlist to be sequential starting from 1, ordered by their // current position. This is needed after removing orphan tracks, to ensure there are no gaps in the track numbering. // The two-step approach (negate then reassign via CTE) avoids UNIQUE constraint violations on (playlist_id, id). -func (r *playlistRepository) renumber(id string) error { +func (r *playlistRepository) renumber(ctx context.Context, id string) error { // Step 1: Negate all IDs to clear the positive ID space - _, err := r.executeSQL(Expr( + _, err := r.executeSQL(ctx, Expr( `UPDATE playlist_tracks SET id = -id WHERE playlist_id = ? AND id > 0`, id)) if err != nil { return err @@ -495,7 +485,7 @@ func (r *playlistRepository) renumber(id string) error { // Step 2: Assign new sequential positive IDs using UPDATE...FROM with a CTE. // The CTE is fully materialized before the UPDATE begins, avoiding self-referencing issues. // ORDER BY id DESC restores original order since IDs are now negative. - _, err = r.executeSQL(Expr( + _, err = r.executeSQL(ctx, Expr( `WITH new_ids AS ( SELECT rowid as rid, ROW_NUMBER() OVER (ORDER BY id DESC) as new_id FROM playlist_tracks WHERE playlist_id = ? @@ -506,10 +496,10 @@ func (r *playlistRepository) renumber(id string) error { if err != nil { return err } - r.enqueueCoverRebuild(id) - return r.refreshCounters(&model.Playlist{ID: id}) + r.enqueueCoverRebuild(ctx, id) + return r.refreshCounters(ctx, &model.Playlist{ID: id}) } var _ model.PlaylistRepository = (*playlistRepository)(nil) -var _ rest.Repository = (*playlistRepository)(nil) -var _ rest.Persistable = (*playlistRepository)(nil) +var _ rest.Repository[model.Playlist] = (*playlistRepository)(nil) +var _ rest.Persistable[model.Playlist] = (*playlistRepository)(nil) diff --git a/persistence/playlist_repository_test.go b/persistence/playlist_repository_test.go index 60263807a..93d37c928 100644 --- a/persistence/playlist_repository_test.go +++ b/persistence/playlist_repository_test.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "slices" "github.com/Masterminds/squirrel" @@ -19,11 +20,11 @@ import ( var _ = Describe("PlaylistRepository", func() { var repo model.PlaylistRepository + var ctx context.Context BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewPlaylistRepository(ctx, GetDBXBuilder()) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid", UserName: "userid", IsAdmin: true}) + repo = NewPlaylistRepository(GetDBXBuilder()) }) Describe("natural sorting", func() { @@ -34,23 +35,23 @@ var _ = Describe("PlaylistRepository", func() { conf.Server.EnableNaturalSorting = true ctx := log.NewContext(GinkgoT().Context()) ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewPlaylistRepository(ctx, GetDBXBuilder()) + repo = NewPlaylistRepository(GetDBXBuilder()) ids = nil for _, n := range []string{"mix 1", "mix 10", "mix 2"} { pls := model.Playlist{Name: n, OwnerID: "userid"} - Expect(repo.Put(&pls)).To(Succeed()) + Expect(repo.Put(ctx, &pls)).To(Succeed()) ids = append(ids, pls.ID) } DeferCleanup(func() { for _, id := range ids { - _ = repo.Delete(id) + _ = repo.Delete(ctx, id) } }) }) It("sorts playlist names by number value", func() { - all, err := repo.GetAll(model.QueryOptions{ + all, err := repo.GetAll(ctx, model.QueryOptions{ Sort: "name", Filters: squirrel.Eq{"playlist.id": ids}, }) Expect(err).ToNot(HaveOccurred()) @@ -61,25 +62,25 @@ var _ = Describe("PlaylistRepository", func() { Describe("Count", func() { It("returns the number of playlists in the DB", func() { - Expect(repo.CountAll()).To(Equal(int64(2))) + Expect(repo.CountAll(ctx)).To(Equal(int64(2))) }) }) Describe("GetCursor", func() { It("yields the same playlists as GetAll", func() { opts := model.QueryOptions{Sort: "name"} - want, err := repo.GetAll(opts) + want, err := repo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - Expect(collectCursor(repo.GetCursor(opts))).To(Equal([]model.Playlist(want))) + Expect(collectCursor(repo.GetCursor(ctx, opts))).To(Equal([]model.Playlist(want))) }) }) Describe("getAllIDs", func() { It("returns the same id set as GetAll", func() { - want, err := repo.GetAll() + want, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(want).ToNot(BeEmpty()) - ids, err := repo.(*playlistRepository).getAllIDs() + ids, err := repo.(*playlistRepository).getAllIDs(ctx) Expect(err).ToNot(HaveOccurred()) Expect(ids).To(ConsistOf(slice.Map(want, func(p model.Playlist) string { return p.ID }))) }) @@ -87,16 +88,16 @@ var _ = Describe("PlaylistRepository", func() { Describe("Exists", func() { It("returns true for an existing playlist", func() { - Expect(repo.Exists(plsCool.ID)).To(BeTrue()) + Expect(repo.Exists(ctx, plsCool.ID)).To(BeTrue()) }) It("returns false for a non-existing playlist", func() { - Expect(repo.Exists("666")).To(BeFalse()) + Expect(repo.Exists(ctx, "666")).To(BeFalse()) }) }) Describe("Get", func() { It("returns an existing playlist", func() { - p, err := repo.Get(plsBest.ID) + p, err := repo.Get(ctx, plsBest.ID) Expect(err).To(BeNil()) // Compare all but Tracks and timestamps p2 := *p @@ -110,11 +111,11 @@ var _ = Describe("PlaylistRepository", func() { } }) It("returns ErrNotFound for a non-existing playlist", func() { - _, err := repo.Get("666") + _, err := repo.Get(ctx, "666") Expect(err).To(MatchError(model.ErrNotFound)) }) It("returns all tracks", func() { - pls, err := repo.GetWithTracks(plsBest.ID, true, false) + pls, err := repo.GetWithTracks(ctx, plsBest.ID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(pls.Name).To(Equal(plsBest.Name)) Expect(pls.Tracks).To(HaveLen(2)) @@ -138,7 +139,7 @@ var _ = Describe("PlaylistRepository", func() { BeforeEach(func() { pls := model.Playlist{Name: "Annotated", OwnerID: "userid"} - Expect(repo.Put(&pls)).To(Succeed()) + Expect(repo.Put(ctx, &pls)).To(Succeed()) plsID = pls.ID }) @@ -151,18 +152,18 @@ var _ = Describe("PlaylistRepository", func() { } It("stores and reads back starred", func() { - Expect(repo.SetStar(true, plsID)).To(Succeed()) + Expect(repo.SetStar(ctx, true, plsID)).To(Succeed()) - p, err := repo.Get(plsID) + p, err := repo.Get(ctx, plsID) Expect(err).ToNot(HaveOccurred()) Expect(p.Starred).To(BeTrue()) Expect(p.StarredAt).ToNot(BeNil()) }) It("stores and reads back rating and average_rating", func() { - Expect(repo.SetRating(4, plsID)).To(Succeed()) + Expect(repo.SetRating(ctx, 4, plsID)).To(Succeed()) - p, err := repo.Get(plsID) + p, err := repo.Get(ctx, plsID) Expect(err).ToNot(HaveOccurred()) Expect(p.Rating).To(Equal(4)) Expect(p.RatedAt).ToNot(BeNil()) @@ -170,21 +171,21 @@ var _ = Describe("PlaylistRepository", func() { }) It("keeps annotations isolated per user", func() { - Expect(repo.SetStar(true, plsID)).To(Succeed()) + Expect(repo.SetStar(ctx, true, plsID)).To(Succeed()) otherCtx := request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "otheruser", UserName: "otheruser", IsAdmin: true}) - otherRepo := NewPlaylistRepository(otherCtx, GetDBXBuilder()) + otherRepo := NewPlaylistRepository(GetDBXBuilder()) - p, err := otherRepo.Get(plsID) + p, err := otherRepo.Get(otherCtx, plsID) Expect(err).ToNot(HaveOccurred()) Expect(p.Starred).To(BeFalse()) }) It("reads starred back through GetAll", func() { - Expect(repo.SetStar(true, plsID)).To(Succeed()) + Expect(repo.SetStar(ctx, true, plsID)).To(Succeed()) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) idx := slices.IndexFunc(all, func(p model.Playlist) bool { return p.ID == plsID }) Expect(idx).To(BeNumerically(">=", 0)) @@ -192,44 +193,43 @@ var _ = Describe("PlaylistRepository", func() { }) It("counts playlists using annotation filters", func() { - Expect(repo.SetStar(true, plsID)).To(Succeed()) + Expect(repo.SetStar(ctx, true, plsID)).To(Succeed()) options := model.QueryOptions{Filters: squirrel.Eq{"starred": true}} - starred, err := repo.GetAll(options) + starred, err := repo.GetAll(ctx, options) Expect(err).ToNot(HaveOccurred()) Expect(starred).To(ContainElement(HaveField("ID", plsID))) - count, err := repo.CountAll(options) + count, err := repo.CountAll(ctx, options) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(len(starred)))) }) It("filters starred playlists through the registered REST filter", func() { - Expect(repo.SetStar(true, plsID)).To(Succeed()) + Expect(repo.SetStar(ctx, true, plsID)).To(Succeed()) - res, err := repo.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := repo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "true"}, }) Expect(err).ToNot(HaveOccurred()) - starred := res.(model.Playlists) - Expect(starred).To(ContainElement(HaveField("ID", plsID))) - for _, p := range starred { + Expect(res).To(ContainElement(HaveField("ID", plsID))) + for _, p := range res { Expect(p.Starred).To(BeTrue()) } - res, err = repo.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err = repo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": "false"}, }) Expect(err).ToNot(HaveOccurred()) - Expect(res.(model.Playlists)).ToNot(ContainElement(HaveField("ID", plsID))) + Expect(res).ToNot(ContainElement(HaveField("ID", plsID))) }) It("reads a playlist by id through the REST id filter without ambiguity", func() { - res, err := repo.(model.ResourceRepository).ReadAll(rest.QueryOptions{ + res, err := repo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"id": plsID}, }) Expect(err).ToNot(HaveOccurred()) - Expect(res.(model.Playlists)).To(ContainElement(HaveField("ID", plsID))) + Expect(res).To(ContainElement(HaveField("ID", plsID))) }) It("does not leak an annotation row of another item_type sharing the playlist id", func() { @@ -240,11 +240,11 @@ var _ = Describe("PlaylistRepository", func() { Bind(dbx.Params{"uid": "userid", "id": plsID}).Execute() Expect(err).ToNot(HaveOccurred()) - p, err := repo.Get(plsID) + p, err := repo.Get(ctx, plsID) Expect(err).ToNot(HaveOccurred()) Expect(p.Starred).To(BeFalse()) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) matches := 0 for _, pl := range all { @@ -256,12 +256,12 @@ var _ = Describe("PlaylistRepository", func() { }) It("relies on the annotation sweep, not Delete, to clean up annotations", func() { - Expect(repo.SetStar(true, plsID)).To(Succeed()) + Expect(repo.SetStar(ctx, true, plsID)).To(Succeed()) - Expect(repo.Delete(plsID)).To(Succeed()) + Expect(repo.Delete(ctx, plsID)).To(Succeed()) Expect(countAnnotations()).To(Equal(1)) - Expect(repo.(*playlistRepository).cleanAnnotations()).To(Succeed()) + Expect(repo.(*playlistRepository).cleanAnnotations(ctx)).To(Succeed()) Expect(countAnnotations()).To(Equal(0)) }) }) @@ -271,8 +271,8 @@ var _ = Describe("PlaylistRepository", func() { pls := model.Playlist{Name: "Smart Counters", OwnerID: "userid", Rules: &criteria.Criteria{ Expression: criteria.All{criteria.Contains{"title": "love"}}, }} - Expect(repo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { Expect(repo.Delete(pls.ID)).To(Succeed()) }) + Expect(repo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { Expect(repo.Delete(ctx, pls.ID)).To(Succeed()) }) // Simulate a previous evaluation having stored the counters _, err := GetDBXBuilder().NewQuery("update playlist set song_count = 42, duration = 123, size = 456 where id = {:id}"). @@ -282,9 +282,9 @@ var _ = Describe("PlaylistRepository", func() { pls.SongCount = 0 pls.Duration = 0 pls.Size = 0 - Expect(repo.Put(&pls)).To(Succeed()) + Expect(repo.Put(ctx, &pls)).To(Succeed()) - saved, err := repo.Get(pls.ID) + saved, err := repo.Get(ctx, pls.ID) Expect(err).ToNot(HaveOccurred()) Expect(saved.SongCount).To(Equal(42)) Expect(saved.Duration).To(Equal(float32(123))) @@ -298,35 +298,35 @@ var _ = Describe("PlaylistRepository", func() { newPls.AddMediaFilesByID([]string{"1004", "1003"}) By("saves the playlist to the DB") - Expect(repo.Put(&newPls)).To(BeNil()) + Expect(repo.Put(ctx, &newPls)).To(BeNil()) By("adds repeated songs to a playlist and keeps the order") newPls.AddMediaFilesByID([]string{"1004"}) - Expect(repo.Put(&newPls)).To(BeNil()) - saved, _ := repo.GetWithTracks(newPls.ID, true, false) + Expect(repo.Put(ctx, &newPls)).To(BeNil()) + saved, _ := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(saved.Tracks).To(HaveLen(3)) Expect(saved.Tracks[0].MediaFileID).To(Equal("1004")) Expect(saved.Tracks[1].MediaFileID).To(Equal("1003")) Expect(saved.Tracks[2].MediaFileID).To(Equal("1004")) By("returns the newly created playlist") - Expect(repo.Exists(newPls.ID)).To(BeTrue()) + Expect(repo.Exists(ctx, newPls.ID)).To(BeTrue()) By("returns deletes the playlist") - Expect(repo.Delete(newPls.ID)).To(BeNil()) + Expect(repo.Delete(ctx, newPls.ID)).To(BeNil()) By("returns error if tries to retrieve the deleted playlist") - Expect(repo.Exists(newPls.ID)).To(BeFalse()) + Expect(repo.Exists(ctx, newPls.ID)).To(BeFalse()) }) It("enqueues a new empty playlist's artwork under its generated id, not an empty id", func() { ctx := request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid", UserName: "userid", IsAdmin: true}) newPls := model.Playlist{Name: "Empty PL", OwnerID: "userid"} // no tracks → refreshCounters path - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) Expect(newPls.ID).ToNot(BeEmpty()) - DeferCleanup(func() { _ = repo.Delete(newPls.ID) }) + DeferCleanup(func() { _ = repo.Delete(ctx, newPls.ID) }) - queued, err := NewArtworkQueueRepository(ctx, GetDBXBuilder()).DequeueBatch(1000) + queued, err := NewArtworkQueueRepository(GetDBXBuilder()).DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(queued).To(ContainElement(SatisfyAll(HaveField("ItemKind", "pl"), HaveField("ItemID", newPls.ID)))) Expect(queued).ToNot(ContainElement(HaveField("ItemID", "")), "must not enqueue an empty playlist id") @@ -336,23 +336,23 @@ var _ = Describe("PlaylistRepository", func() { It("does not enqueue artwork when only metadata changes", func() { ctx := request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid", UserName: "userid", IsAdmin: true}) newPls := model.Playlist{Name: "Rename Me", OwnerID: "userid"} - Expect(repo.Put(&newPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(newPls.ID) }) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, newPls.ID) }) // Clear the row creation just enqueued, so anything present afterwards came from the update. - queueRepo := NewArtworkQueueRepository(ctx, GetDBXBuilder()) - queued, err := queueRepo.DequeueBatch(1000) + queueRepo := NewArtworkQueueRepository(GetDBXBuilder()) + queued, err := queueRepo.DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) for _, q := range queued { if q.ItemID == newPls.ID { - Expect(queueRepo.DeleteIfUnchanged(q.ItemKind, q.ItemID, q.ImageType, q.RetryAt)).To(Succeed()) + Expect(queueRepo.DeleteIfUnchanged(ctx, q.ItemKind, q.ItemID, q.ImageType, q.RetryAt)).To(Succeed()) } } newPls.Name = "Renamed" newPls.Comment = "edited" - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) - queued, err = queueRepo.DequeueBatch(1000) + queued, err = queueRepo.DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(queued).ToNot(ContainElement(HaveField("ItemID", newPls.ID))) }) @@ -361,10 +361,10 @@ var _ = Describe("PlaylistRepository", func() { ctx := request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid", UserName: "userid", IsAdmin: true}) newPls := model.Playlist{Name: "Grid PL", OwnerID: "userid"} newPls.AddMediaFilesByID([]string{"1001", "1002"}) - Expect(repo.Put(&newPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(newPls.ID) }) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, newPls.ID) }) - queued, err := NewArtworkQueueRepository(ctx, GetDBXBuilder()).DequeueBatch(1000) + queued, err := NewArtworkQueueRepository(GetDBXBuilder()).DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(queued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "pl"), @@ -374,7 +374,7 @@ var _ = Describe("PlaylistRepository", func() { Describe("GetAll", func() { It("returns all playlists from DB", func() { - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) Expect(all[0].ID).To(Equal(plsBest.ID)) Expect(all[1].ID).To(Equal(plsCool.ID)) @@ -383,14 +383,14 @@ var _ = Describe("PlaylistRepository", func() { Describe("GetPlaylists", func() { It("returns playlists for a track", func() { - pls, err := repo.GetPlaylists(songRadioactivity.ID) + pls, err := repo.GetPlaylists(ctx, songRadioactivity.ID) Expect(err).ToNot(HaveOccurred()) Expect(pls).To(HaveLen(1)) Expect(pls[0].ID).To(Equal(plsBest.ID)) }) It("returns empty when none", func() { - pls, err := repo.GetPlaylists("9999") + pls, err := repo.GetPlaylists(ctx, "9999") Expect(err).ToNot(HaveOccurred()) Expect(pls).To(HaveLen(0)) }) @@ -401,14 +401,14 @@ var _ = Describe("PlaylistRepository", func() { AfterEach(func() { if testPlaylistID != "" { - Expect(repo.Delete(testPlaylistID)).To(BeNil()) + Expect(repo.Delete(ctx, testPlaylistID)).To(BeNil()) testPlaylistID = "" } }) // helper to get track positions and media file IDs getTrackInfo := func(playlistID string) (ids []string, mediaFileIDs []string) { - pls, err := repo.GetWithTracks(playlistID, false, false) + pls, err := repo.GetWithTracks(ctx, playlistID, false, false) Expect(err).ToNot(HaveOccurred()) for _, t := range pls.Tracks { ids = append(ids, t.ID) @@ -421,12 +421,12 @@ var _ = Describe("PlaylistRepository", func() { By("creating a playlist with 4 tracks") newPls := model.Playlist{Name: "Renumber Test Middle", OwnerID: "userid"} newPls.AddMediaFilesByID([]string{"1001", "1002", "1003", "1004"}) - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("deleting the second track (position 2)") - tracksRepo := repo.Tracks(newPls.ID, false) - Expect(tracksRepo.Delete("2")).To(Succeed()) + tracksRepo := repo.Tracks(ctx, newPls.ID, false) + Expect(tracksRepo.Delete(ctx, "2")).To(Succeed()) By("verifying remaining tracks are renumbered sequentially") ids, mediaFileIDs := getTrackInfo(newPls.ID) @@ -438,12 +438,12 @@ var _ = Describe("PlaylistRepository", func() { By("creating a playlist with 3 tracks") newPls := model.Playlist{Name: "Renumber Test First", OwnerID: "userid"} newPls.AddMediaFilesByID([]string{"1001", "1002", "1003"}) - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("deleting the first track (position 1)") - tracksRepo := repo.Tracks(newPls.ID, false) - Expect(tracksRepo.Delete("1")).To(Succeed()) + tracksRepo := repo.Tracks(ctx, newPls.ID, false) + Expect(tracksRepo.Delete(ctx, "1")).To(Succeed()) By("verifying remaining tracks are renumbered sequentially") ids, mediaFileIDs := getTrackInfo(newPls.ID) @@ -455,12 +455,12 @@ var _ = Describe("PlaylistRepository", func() { By("creating a playlist with 3 tracks") newPls := model.Playlist{Name: "Renumber Test Last", OwnerID: "userid"} newPls.AddMediaFilesByID([]string{"1001", "1002", "1003"}) - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("deleting the last track (position 3)") - tracksRepo := repo.Tracks(newPls.ID, false) - Expect(tracksRepo.Delete("3")).To(Succeed()) + tracksRepo := repo.Tracks(ctx, newPls.ID, false) + Expect(tracksRepo.Delete(ctx, "3")).To(Succeed()) By("verifying remaining tracks are renumbered sequentially") ids, mediaFileIDs := getTrackInfo(newPls.ID) @@ -476,18 +476,18 @@ var _ = Describe("PlaylistRepository", func() { // "userid" is the fixture user; playlist.owner_id has a FK to user(id). owner := model.User{ID: "userid", UserName: "userid"} octx := request.WithUser(GinkgoT().Context(), owner) - ownerRepo := NewPlaylistRepository(octx, GetDBXBuilder()) + ownerRepo := NewPlaylistRepository(GetDBXBuilder()) pls := model.Playlist{Name: "Private One", OwnerID: owner.ID, Public: false} - Expect(ownerRepo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { _ = ownerRepo.Delete(pls.ID) }) + Expect(ownerRepo.Put(octx, &pls)).To(Succeed()) + DeferCleanup(func() { _ = ownerRepo.Delete(octx, pls.ID) }) - Expect(ownerRepo.Exists(pls.ID)).To(BeTrue(), "the owner sees it") + Expect(ownerRepo.Exists(octx, pls.ID)).To(BeTrue(), "the owner sees it") - anon := NewPlaylistRepository(GinkgoT().Context(), GetDBXBuilder()) - Expect(anon.Exists(pls.ID)).To(BeFalse(), "no user: userFilter hides it") + anon := NewPlaylistRepository(GetDBXBuilder()) + Expect(anon.Exists(GinkgoT().Context(), pls.ID)).To(BeFalse(), "no user: userFilter hides it") admin := request.WithUser(GinkgoT().Context(), model.User{ID: "userid", IsAdmin: true}) - Expect(NewPlaylistRepository(admin, GetDBXBuilder()).Exists(pls.ID)).To(BeTrue(), + Expect(NewPlaylistRepository(GetDBXBuilder()).Exists(admin, pls.ID)).To(BeTrue(), "elevating is what the public image route relies on") }) }) diff --git a/persistence/playlist_track_repository.go b/persistence/playlist_track_repository.go index 848b67be0..341aab8af 100644 --- a/persistence/playlist_track_repository.go +++ b/persistence/playlist_track_repository.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "database/sql" "slices" @@ -40,11 +41,10 @@ func (t dbPlaylistTracks) toModels() model.PlaylistTracks { }) } -func (r *playlistRepository) Tracks(playlistId string, refreshSmartPlaylist bool) model.PlaylistTrackRepository { +func (r *playlistRepository) Tracks(ctx context.Context, playlistId string, refreshSmartPlaylist bool) model.PlaylistTrackRepository { p := &playlistTrackRepository{} p.playlistRepo = r p.playlistId = playlistId - p.ctx = r.ctx p.db = r.db p.tableName = "playlist_tracks" p.registerModel(&model.PlaylistTrack{}, map[string]filterFunc{ @@ -67,37 +67,37 @@ func (r *playlistRepository) Tracks(playlistId string, refreshSmartPlaylist bool }, "f") // TODO I don't like this solution, but I won't change it now as it's not the focus of BFR. - pls, err := r.Get(playlistId) + pls, err := r.Get(ctx, playlistId) if err != nil { - log.Warn(r.ctx, "Error getting playlist's tracks", "playlistId", playlistId, err) + log.Warn(ctx, "Error getting playlist's tracks", "playlistId", playlistId, err) return nil } if refreshSmartPlaylist { - r.refreshSmartPlaylist(pls) + r.refreshSmartPlaylist(ctx, pls) } p.playlist = pls return p } -func (r *playlistTrackRepository) CountAll(options ...model.QueryOptions) (int64, error) { +func (r *playlistTrackRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { query := Select(). Join("media_file f on f.id = media_file_id"). Where(Eq{"playlist_id": r.playlistId}) - query = r.applyLibraryFilter(query, "f") - return r.count(query, options...) + query = r.applyLibraryFilter(ctx, query, "f") + return r.count(ctx, query, options...) } -func (r *playlistTrackRepository) Count(options ...rest.QueryOptions) (int64, error) { +func (r *playlistTrackRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { query := Select(). LeftJoin("media_file f on f.id = media_file_id"). Where(Eq{"playlist_id": r.playlistId}) - query = r.applyLibraryFilter(query, "f") - return r.count(query, r.parseRestOptions(r.ctx, options...)) + query = r.applyLibraryFilter(ctx, query, "f") + return r.count(ctx, query, r.parseRestOptions(ctx, options...)) } -func (r *playlistTrackRepository) Read(id string) (any, error) { - userID := loggedUser(r.ctx).ID - sel := r.newSelect(). +func (r *playlistTrackRepository) Read(ctx context.Context, id string) (*model.PlaylistTrack, error) { + userID := loggedUser(ctx).ID + sel := r.newSelect(ctx). LeftJoin("annotation on ("+ "annotation.item_id = media_file_id"+ " AND annotation.item_type = 'media_file'"+ @@ -114,23 +114,23 @@ func (r *playlistTrackRepository) Read(id string) (any, error) { ). Join("media_file f on f.id = media_file_id"). Where(And{Eq{"playlist_id": r.playlistId}, Eq{"playlist_tracks.id": id}}) - sel = r.applyLibraryFilter(sel, "f") + sel = r.applyLibraryFilter(ctx, sel, "f") var trk dbPlaylistTrack - err := r.queryOne(sel, &trk) + err := r.queryOne(ctx, sel, &trk) return trk.PlaylistTrack, err } -func (r *playlistTrackRepository) GetAll(options ...model.QueryOptions) (model.PlaylistTracks, error) { - tracks, err := r.playlistRepo.loadTracks(r.newSelect(options...), r.playlistId) +func (r *playlistTrackRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.PlaylistTracks, error) { + tracks, err := r.playlistRepo.loadTracks(ctx, r.newSelect(ctx, options...), r.playlistId) if err != nil { return nil, err } return tracks, err } -func (r *playlistTrackRepository) GetCursor(options ...model.QueryOptions) (model.PlaylistTrackCursor, error) { - sel := r.playlistRepo.tracksQuery(r.newSelect(options...), r.playlistId) - cursor, err := queryWithStableResults[dbPlaylistTrack](r.sqlRepository, sel) +func (r *playlistTrackRepository) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.PlaylistTrackCursor, error) { + sel := r.playlistRepo.tracksQuery(ctx, r.newSelect(ctx, options...), r.playlistId) + cursor, err := queryWithStableResults[dbPlaylistTrack](ctx, r.sqlRepository, sel) if err != nil { return nil, err } @@ -138,114 +138,106 @@ func (r *playlistTrackRepository) GetCursor(options ...model.QueryOptions) (mode return t.PlaylistTrack }) return model.PlaylistTrackCursor(hydrateCursor(tracks, func(batch []model.PlaylistTrack) { - hydratePlaylistTrackArtwork(r.ctx, r.db, batch) + hydratePlaylistTrackArtwork(ctx, r.db, batch) })), nil } // GetMediaFileIDs returns the tracks' song ids, for callers that need every id but no track data. -func (r *playlistTrackRepository) GetMediaFileIDs(options ...model.QueryOptions) ([]string, error) { - query := r.newSelect(options...).Columns("media_file_id"). +func (r *playlistTrackRepository) GetMediaFileIDs(ctx context.Context, options ...model.QueryOptions) ([]string, error) { + query := r.newSelect(ctx, options...).Columns("media_file_id"). Join("media_file f on f.id = media_file_id"). Where(Eq{"playlist_id": r.playlistId}) - query = r.applyLibraryFilter(query, "f") + query = r.applyLibraryFilter(ctx, query, "f") var ids []string - if err := r.queryAllSlice(query, &ids); err != nil { + if err := r.queryAllSlice(ctx, query, &ids); err != nil { return nil, err } return ids, nil } -func (r *playlistTrackRepository) GetAlbumIDs(options ...model.QueryOptions) ([]string, error) { - query := r.newSelect(options...).Columns("distinct mf.album_id"). +func (r *playlistTrackRepository) GetAlbumIDs(ctx context.Context, options ...model.QueryOptions) ([]string, error) { + query := r.newSelect(ctx, options...).Columns("distinct mf.album_id"). Join("media_file mf on mf.id = media_file_id"). Where(Eq{"playlist_id": r.playlistId}) - query = r.applyLibraryFilter(query, "mf") + query = r.applyLibraryFilter(ctx, query, "mf") var ids []string - err := r.queryAllSlice(query, &ids) + err := r.queryAllSlice(ctx, query, &ids) if err != nil { return nil, err } return ids, nil } -func (r *playlistTrackRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) +func (r *playlistTrackRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.PlaylistTrack, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *playlistTrackRepository) EntityName() string { - return "playlist_tracks" -} - -func (r *playlistTrackRepository) NewInstance() any { - return &model.PlaylistTrack{} -} - -func (r *playlistTrackRepository) Add(mediaFileIds []string) (int, error) { +func (r *playlistTrackRepository) Add(ctx context.Context, mediaFileIds []string) (int, error) { if len(mediaFileIds) > 0 { - log.Debug(r.ctx, "Adding songs to playlist", "playlistId", r.playlistId, "mediaFileIds", mediaFileIds) + log.Debug(ctx, "Adding songs to playlist", "playlistId", r.playlistId, "mediaFileIds", mediaFileIds) } else { return 0, nil } // Get next pos (ID) in playlist - sq := r.newSelect().Columns("max(id) as max").Where(Eq{"playlist_id": r.playlistId}) + sq := r.newSelect(ctx).Columns("max(id) as max").Where(Eq{"playlist_id": r.playlistId}) var res struct{ Max sql.NullInt32 } - if err := r.queryOne(sq, &res); err != nil { + if err := r.queryOne(ctx, sq, &res); err != nil { return 0, err } - return r.playlistRepo.addTracks(r.playlistId, int(res.Max.Int32+1), mediaFileIds) + return r.playlistRepo.addTracks(ctx, r.playlistId, int(res.Max.Int32+1), mediaFileIds) } // Insert adds tracks before the 1-based position pos, shifting the following entries down; a // position past the end appends. Callers must run it in a transaction. -func (r *playlistTrackRepository) Insert(mediaFileIds []string, pos int) (int, error) { +func (r *playlistTrackRepository) Insert(ctx context.Context, mediaFileIds []string, pos int) (int, error) { if len(mediaFileIds) == 0 { return 0, nil } pos = max(pos, 1) n := len(mediaFileIds) // Negate while shifting, so no intermediate row hits the unique (playlist_id, id) index. - _, err := r.executeSQL(Expr(`UPDATE playlist_tracks SET id = -(id + ?) WHERE playlist_id = ? AND id >= ?`, n, r.playlistId, pos)) + _, err := r.executeSQL(ctx, Expr(`UPDATE playlist_tracks SET id = -(id + ?) WHERE playlist_id = ? AND id >= ?`, n, r.playlistId, pos)) if err != nil { return 0, err } - res, err := r.executeSQL(Expr(`UPDATE playlist_tracks SET id = -id WHERE playlist_id = ? AND id < 0`, r.playlistId)) + res, err := r.executeSQL(ctx, Expr(`UPDATE playlist_tracks SET id = -id WHERE playlist_id = ? AND id < 0`, r.playlistId)) if err != nil { return 0, err } if res == 0 { - return r.Add(mediaFileIds) + return r.Add(ctx, mediaFileIds) } - inserted, err := r.playlistRepo.addTracks(r.playlistId, pos, mediaFileIds) + inserted, err := r.playlistRepo.addTracks(ctx, r.playlistId, pos, mediaFileIds) if err != nil || inserted == n { return inserted, err } // The shift above reserved a slot per requested id, so ids dropped by the library filter // leave a hole. Close it. - return inserted, r.playlistRepo.renumber(r.playlistId) + return inserted, r.playlistRepo.renumber(ctx, r.playlistId) } -func (r *playlistTrackRepository) addMediaFileIds(cond Sqlizer) (int, error) { +func (r *playlistTrackRepository) addMediaFileIds(ctx context.Context, cond Sqlizer) (int, error) { sq := Select("id").From("media_file").Where(cond).OrderBy("album_artist, album, release_date, disc_number, track_number") var ids []string - err := r.queryAllSlice(sq, &ids) + err := r.queryAllSlice(ctx, sq, &ids) if err != nil { - log.Error(r.ctx, "Error getting tracks to add to playlist", err) + log.Error(ctx, "Error getting tracks to add to playlist", err) return 0, err } - return r.Add(ids) + return r.Add(ctx, ids) } -func (r *playlistTrackRepository) AddAlbums(albumIds []string) (int, error) { - return r.addMediaFileIds(Eq{"album_id": albumIds}) +func (r *playlistTrackRepository) AddAlbums(ctx context.Context, albumIds []string) (int, error) { + return r.addMediaFileIds(ctx, Eq{"album_id": albumIds}) } -func (r *playlistTrackRepository) AddArtists(artistIds []string) (int, error) { - return r.addMediaFileIds(Eq{"album_artist_id": artistIds}) +func (r *playlistTrackRepository) AddArtists(ctx context.Context, artistIds []string) (int, error) { + return r.addMediaFileIds(ctx, Eq{"album_artist_id": artistIds}) } -func (r *playlistTrackRepository) AddDiscs(discs []model.DiscID) (int, error) { +func (r *playlistTrackRepository) AddDiscs(ctx context.Context, discs []model.DiscID) (int, error) { if len(discs) == 0 { return 0, nil } @@ -253,36 +245,36 @@ func (r *playlistTrackRepository) AddDiscs(discs []model.DiscID) (int, error) { for _, d := range discs { clauses = append(clauses, And{Eq{"album_id": d.AlbumID}, Eq{"release_date": d.ReleaseDate}, Eq{"disc_number": d.DiscNumber}}) } - return r.addMediaFileIds(clauses) + return r.addMediaFileIds(ctx, clauses) } // deleteChunkSize keeps each DELETE under SQLITE_MAX_VARIABLE_NUMBER, matching addTracks. const deleteChunkSize = 200 -func (r *playlistTrackRepository) Delete(ids ...string) error { +func (r *playlistTrackRepository) Delete(ctx context.Context, ids ...string) error { for chunk := range slices.Chunk(ids, deleteChunkSize) { - if err := r.delete(And{Eq{"playlist_id": r.playlistId}, Eq{"id": chunk}}); err != nil { + if err := r.delete(ctx, And{Eq{"playlist_id": r.playlistId}, Eq{"id": chunk}}); err != nil { return err } } - return r.playlistRepo.renumber(r.playlistId) + return r.playlistRepo.renumber(ctx, r.playlistId) } -func (r *playlistTrackRepository) DeleteAll() error { - err := r.delete(Eq{"playlist_id": r.playlistId}) +func (r *playlistTrackRepository) DeleteAll(ctx context.Context) error { + err := r.delete(ctx, Eq{"playlist_id": r.playlistId}) if err != nil { return err } - return r.playlistRepo.renumber(r.playlistId) + return r.playlistRepo.renumber(ctx, r.playlistId) } // Reorder moves a track from pos to newPos, shifting other tracks accordingly. newPos is clamped // to the playlist; a pos outside it is ErrNotFound, since shifting around it would leave a gap. -func (r *playlistTrackRepository) Reorder(pos int, newPos int) error { +func (r *playlistTrackRepository) Reorder(ctx context.Context, pos int, newPos int) error { var res struct{ Max sql.NullInt32 } - if err := r.queryOne(r.newSelect().Columns("max(id) as max").Where(Eq{"playlist_id": r.playlistId}), &res); err != nil { + if err := r.queryOne(ctx, r.newSelect(ctx).Columns("max(id) as max").Where(Eq{"playlist_id": r.playlistId}), &res); err != nil { return err } last := int(res.Max.Int32) @@ -296,7 +288,7 @@ func (r *playlistTrackRepository) Reorder(pos int, newPos int) error { pid := r.playlistId // Step 1: Move the source track out of the way (temporary sentinel value) - _, err := r.executeSQL(Expr( + _, err := r.executeSQL(ctx, Expr( `UPDATE playlist_tracks SET id = -999999 WHERE playlist_id = ? AND id = ?`, pid, pos)) if err != nil { return err @@ -304,11 +296,11 @@ func (r *playlistTrackRepository) Reorder(pos int, newPos int) error { // Step 2: Shift the affected range using negative values to avoid unique constraint violations if pos < newPos { - _, err = r.executeSQL(Expr( + _, err = r.executeSQL(ctx, Expr( `UPDATE playlist_tracks SET id = -(id - 1) WHERE playlist_id = ? AND id > ? AND id <= ?`, pid, pos, newPos)) } else { - _, err = r.executeSQL(Expr( + _, err = r.executeSQL(ctx, Expr( `UPDATE playlist_tracks SET id = -(id + 1) WHERE playlist_id = ? AND id >= ? AND id < ?`, pid, newPos, pos)) } @@ -317,14 +309,14 @@ func (r *playlistTrackRepository) Reorder(pos int, newPos int) error { } // Step 3: Flip the shifted range back to positive - _, err = r.executeSQL(Expr( + _, err = r.executeSQL(ctx, Expr( `UPDATE playlist_tracks SET id = -id WHERE playlist_id = ? AND id < 0 AND id != -999999`, pid)) if err != nil { return err } // Step 4: Place the source track at its new position - _, err = r.executeSQL(Expr( + _, err = r.executeSQL(ctx, Expr( `UPDATE playlist_tracks SET id = ? WHERE playlist_id = ? AND id = -999999`, newPos, pid)) return err } diff --git a/persistence/playlist_track_repository_test.go b/persistence/playlist_track_repository_test.go index 1a6bc9dc6..3c532c405 100644 --- a/persistence/playlist_track_repository_test.go +++ b/persistence/playlist_track_repository_test.go @@ -17,30 +17,31 @@ const sqliteMaxVariables = 32766 var _ = Describe("PlaylistTrackRepository", func() { var repo model.PlaylistTrackRepository + var ctx context.Context BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) + ctx = log.NewContext(GinkgoT().Context()) ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewPlaylistRepository(ctx, GetDBXBuilder()).Tracks(plsBest.ID, true) + repo = NewPlaylistRepository(GetDBXBuilder()).Tracks(ctx, plsBest.ID, true) }) Describe("GetCursor", func() { It("yields the same tracks as GetAll", func() { opts := model.QueryOptions{Sort: "id"} - want, err := repo.GetAll(opts) + want, err := repo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).To(HaveLen(2)) - Expect(collectCursor(repo.GetCursor(opts))).To(Equal([]model.PlaylistTrack(want))) + Expect(collectCursor(repo.GetCursor(ctx, opts))).To(Equal([]model.PlaylistTrack(want))) }) It("honors Max and Offset", func() { opts := model.QueryOptions{Sort: "id", Max: 1, Offset: 1} - want, err := repo.GetAll(opts) + want, err := repo.GetAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) Expect(want).To(HaveLen(1)) - Expect(collectCursor(repo.GetCursor(opts))).To(Equal([]model.PlaylistTrack(want))) + Expect(collectCursor(repo.GetCursor(ctx, opts))).To(Equal([]model.PlaylistTrack(want))) }) }) @@ -48,11 +49,11 @@ var _ = Describe("PlaylistTrackRepository", func() { It("returns every row under a random sort, despite the integer id", func() { // playlist_tracks.id is an INTEGER, so SEEDEDRAND drops every row unless it is cast to // TEXT, and it fails silently: no error, just no rows. - all, err := repo.GetAll(model.QueryOptions{Sort: "random"}) + all, err := repo.GetAll(ctx, model.QueryOptions{Sort: "random"}) Expect(err).ToNot(HaveOccurred()) Expect(all).To(HaveLen(2), "a random sort must not silently drop rows") - got, err := repo.GetAll(model.QueryOptions{Sort: "random", Max: 1}) + got, err := repo.GetAll(ctx, model.QueryOptions{Sort: "random", Max: 1}) Expect(err).ToNot(HaveOccurred()) Expect(got).To(HaveLen(1)) }) @@ -60,22 +61,22 @@ var _ = Describe("PlaylistTrackRepository", func() { Describe("CountAll", func() { It("returns the number of tracks in the playlist", func() { - Expect(repo.CountAll()).To(Equal(int64(2))) + Expect(repo.CountAll(ctx)).To(Equal(int64(2))) }) It("ignores Max and Offset", func() { - Expect(repo.CountAll(model.QueryOptions{Max: 1, Offset: 1})).To(Equal(int64(2))) + Expect(repo.CountAll(ctx, model.QueryOptions{Max: 1, Offset: 1})).To(Equal(int64(2))) }) }) Describe("GetMediaFileIDs", func() { It("returns the song ids in playlist order", func() { - Expect(repo.GetMediaFileIDs(model.QueryOptions{Sort: "id"})). + Expect(repo.GetMediaFileIDs(ctx, model.QueryOptions{Sort: "id"})). To(Equal([]string{songDayInALife.ID, songRadioactivity.ID})) }) It("honors Max and Offset", func() { - Expect(repo.GetMediaFileIDs(model.QueryOptions{Sort: "id", Max: 1, Offset: 1})). + Expect(repo.GetMediaFileIDs(ctx, model.QueryOptions{Sort: "id", Max: 1, Offset: 1})). To(Equal([]string{songRadioactivity.ID})) }) }) @@ -84,28 +85,26 @@ var _ = Describe("PlaylistTrackRepository", func() { var tracks model.PlaylistTrackRepository BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - plsRepo := NewPlaylistRepository(ctx, GetDBXBuilder()) + plsRepo := NewPlaylistRepository(GetDBXBuilder()) pls := model.Playlist{Name: "Insert", OwnerID: "userid", OwnerName: "userid"} - Expect(plsRepo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { Expect(plsRepo.Delete(pls.ID)).To(Succeed()) }) + Expect(plsRepo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { Expect(plsRepo.Delete(ctx, pls.ID)).To(Succeed()) }) - tracks = plsRepo.Tracks(pls.ID, false) - Expect(tracks.Add([]string{songDayInALife.ID, songRadioactivity.ID})).To(Equal(2)) + tracks = plsRepo.Tracks(ctx, pls.ID, false) + Expect(tracks.Add(ctx, []string{songDayInALife.ID, songRadioactivity.ID})).To(Equal(2)) }) order := func() []string { - ids, err := tracks.GetMediaFileIDs(model.QueryOptions{Sort: "id"}) + ids, err := tracks.GetMediaFileIDs(ctx, model.QueryOptions{Sort: "id"}) Expect(err).ToNot(HaveOccurred()) return ids } DescribeTable("inserts before a 1-based position, keeping the new tracks' order", func(pos int, want func() []string) { - Expect(tracks.Insert([]string{songComeTogether.ID, songAntenna.ID}, pos)).To(Equal(2)) + Expect(tracks.Insert(ctx, []string{songComeTogether.ID, songAntenna.ID}, pos)).To(Equal(2)) Expect(order()).To(Equal(want())) - Expect(tracks.CountAll()).To(Equal(int64(4))) + Expect(tracks.CountAll(ctx)).To(Equal(int64(4))) }, Entry("in the middle", 2, func() []string { return []string{songDayInALife.ID, songComeTogether.ID, songAntenna.ID, songRadioactivity.ID} @@ -119,8 +118,8 @@ var _ = Describe("PlaylistTrackRepository", func() { ) It("renumbers positions contiguously", func() { - Expect(tracks.Insert([]string{songComeTogether.ID}, 1)).To(Equal(1)) - all, err := tracks.GetAll(model.QueryOptions{Sort: "id"}) + Expect(tracks.Insert(ctx, []string{songComeTogether.ID}, 1)).To(Equal(1)) + all, err := tracks.GetAll(ctx, model.QueryOptions{Sort: "id"}) Expect(err).ToNot(HaveOccurred()) Expect([]string{all[0].ID, all[1].ID, all[2].ID}).To(Equal([]string{"1", "2", "3"})) }) @@ -130,19 +129,17 @@ var _ = Describe("PlaylistTrackRepository", func() { var tracks model.PlaylistTrackRepository BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - plsRepo := NewPlaylistRepository(ctx, GetDBXBuilder()) + plsRepo := NewPlaylistRepository(GetDBXBuilder()) pls := model.Playlist{Name: "Reorder", OwnerID: "userid", OwnerName: "userid"} - Expect(plsRepo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { Expect(plsRepo.Delete(pls.ID)).To(Succeed()) }) + Expect(plsRepo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { Expect(plsRepo.Delete(ctx, pls.ID)).To(Succeed()) }) - tracks = plsRepo.Tracks(pls.ID, false) - Expect(tracks.Add([]string{songDayInALife.ID, songRadioactivity.ID, songComeTogether.ID})).To(Equal(3)) + tracks = plsRepo.Tracks(ctx, pls.ID, false) + Expect(tracks.Add(ctx, []string{songDayInALife.ID, songRadioactivity.ID, songComeTogether.ID})).To(Equal(3)) }) rows := func() ([]string, []string) { - all, err := tracks.GetAll(model.QueryOptions{Sort: "id"}) + all, err := tracks.GetAll(ctx, model.QueryOptions{Sort: "id"}) Expect(err).ToNot(HaveOccurred()) var ids, songs []string for _, t := range all { @@ -154,7 +151,7 @@ var _ = Describe("PlaylistTrackRepository", func() { DescribeTable("clamps the destination to the playlist", func(newPos int, want func() []string) { - Expect(tracks.Reorder(1, newPos)).To(Succeed()) + Expect(tracks.Reorder(ctx, 1, newPos)).To(Succeed()) ids, songs := rows() Expect(ids).To(Equal([]string{"1", "2", "3"})) Expect(songs).To(Equal(want())) @@ -169,7 +166,7 @@ var _ = Describe("PlaylistTrackRepository", func() { DescribeTable("rejects a source position outside the playlist, leaving rows untouched", func(pos int) { - Expect(tracks.Reorder(pos, 1)).To(MatchError(model.ErrNotFound)) + Expect(tracks.Reorder(ctx, pos, 1)).To(MatchError(model.ErrNotFound)) ids, songs := rows() Expect(ids).To(Equal([]string{"1", "2", "3"})) Expect(songs).To(Equal([]string{songDayInALife.ID, songRadioactivity.ID, songComeTogether.ID})) @@ -192,35 +189,33 @@ var _ = Describe("PlaylistTrackRepository", func() { } BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - plsRepo := NewPlaylistRepository(ctx, GetDBXBuilder()) + plsRepo := NewPlaylistRepository(GetDBXBuilder()) pls := model.Playlist{Name: "Chunked Delete", OwnerID: "userid", OwnerName: "userid"} - Expect(plsRepo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { Expect(plsRepo.Delete(pls.ID)).To(Succeed()) }) + Expect(plsRepo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { Expect(plsRepo.Delete(ctx, pls.ID)).To(Succeed()) }) - tracks = plsRepo.Tracks(pls.ID, false) + tracks = plsRepo.Tracks(ctx, pls.ID, false) songIds := make([]string, numTracks) for i := range songIds { songIds[i] = songDayInALife.ID } - Expect(tracks.Add(songIds)).To(Equal(numTracks)) + Expect(tracks.Add(ctx, songIds)).To(Equal(numTracks)) }) It("removes positions spanning several chunks, and renumbers what is left", func() { - Expect(tracks.Delete(positionsUpTo(numTracks - 1)...)).To(Succeed()) + Expect(tracks.Delete(ctx, positionsUpTo(numTracks-1)...)).To(Succeed()) - Expect(tracks.CountAll()).To(Equal(int64(1))) - remaining, err := tracks.GetAll(model.QueryOptions{Sort: "id"}) + Expect(tracks.CountAll(ctx)).To(Equal(int64(1))) + remaining, err := tracks.GetAll(ctx, model.QueryOptions{Sort: "id"}) Expect(err).ToNot(HaveOccurred()) Expect(remaining[0].ID).To(Equal("1"), "the surviving track must be renumbered to position 1") }) It("accepts more ids than SQLite allows as bind variables", func() { - Expect(tracks.Delete(positionsUpTo(sqliteMaxVariables + 100)...)).To(Succeed()) + Expect(tracks.Delete(ctx, positionsUpTo(sqliteMaxVariables+100)...)).To(Succeed()) - Expect(tracks.CountAll()).To(BeZero()) + Expect(tracks.CountAll(ctx)).To(BeZero()) }) }) @@ -236,69 +231,69 @@ var _ = Describe("PlaylistTrackRepository", func() { userCtx = request.WithUser(log.NewContext(GinkgoT().Context()), restrictedUser) db := GetDBXBuilder() - adminMr := NewMediaFileRepository(adminCtx, db) - Expect(adminMr.Put(&model.MediaFile{ + adminMr := NewMediaFileRepository(db) + Expect(adminMr.Put(adminCtx, &model.MediaFile{ ID: "pls-otherlib-track", LibraryID: otherLib.ID, AlbumID: "pls-hidden-album", Path: "hidden/in-playlist.mp3", Title: "Hidden In Playlist", })).To(Succeed()) - DeferCleanup(func() { _ = adminMr.Delete("pls-otherlib-track") }) + DeferCleanup(func() { _ = adminMr.Delete(adminCtx, "pls-otherlib-track") }) - adminPls := NewPlaylistRepository(adminCtx, db) + adminPls := NewPlaylistRepository(db) pls := model.Playlist{Name: "Public Mixed", OwnerID: adminUser.ID, OwnerName: adminUser.UserName, Public: true} - Expect(adminPls.Put(&pls)).To(Succeed()) + Expect(adminPls.Put(adminCtx, &pls)).To(Succeed()) plsID = pls.ID - DeferCleanup(func() { _ = adminPls.Delete(plsID) }) - Expect(adminPls.Tracks(plsID, false).Add([]string{songDayInALife.ID, "pls-otherlib-track"})).To(Equal(2)) + DeferCleanup(func() { _ = adminPls.Delete(adminCtx, plsID) }) + Expect(adminPls.Tracks(adminCtx, plsID, false).Add(adminCtx, []string{songDayInALife.ID, "pls-otherlib-track"})).To(Equal(2)) - userTracks = NewPlaylistRepository(userCtx, db).Tracks(plsID, false) + userTracks = NewPlaylistRepository(db).Tracks(userCtx, plsID, false) }) It("Read does not return a track outside the user's libraries", func() { - _, err := userTracks.Read("2") + _, err := userTracks.Read(userCtx, "2") Expect(err).To(MatchError(model.ErrNotFound), "position 2 holds a track the user cannot access") }) It("Read still returns a track inside the user's libraries", func() { - trk, err := userTracks.Read("1") + trk, err := userTracks.Read(userCtx, "1") Expect(err).ToNot(HaveOccurred()) - Expect(trk.(*model.PlaylistTrack).MediaFile.ID).To(Equal(songDayInALife.ID)) + Expect(trk.MediaFile.ID).To(Equal(songDayInALife.ID)) }) It("Count excludes tracks outside the user's libraries", func() { - Expect(userTracks.Count()).To(Equal(int64(1)), "Count must agree with the filtered listing") + Expect(userTracks.Count(userCtx)).To(Equal(int64(1)), "Count must agree with the filtered listing") }) It("GetAlbumIDs excludes albums outside the user's libraries", func() { - Expect(userTracks.GetAlbumIDs()).ToNot(ContainElement("pls-hidden-album")) + Expect(userTracks.GetAlbumIDs(userCtx)).ToNot(ContainElement("pls-hidden-album")) }) Describe("Add", func() { var ownTracks model.PlaylistTrackRepository BeforeEach(func() { - userPls := NewPlaylistRepository(userCtx, GetDBXBuilder()) + userPls := NewPlaylistRepository(GetDBXBuilder()) own := model.Playlist{Name: "Own Playlist", OwnerID: restrictedUser.ID, OwnerName: restrictedUser.UserName} - Expect(userPls.Put(&own)).To(Succeed()) - DeferCleanup(func() { _ = NewPlaylistRepository(adminCtx, GetDBXBuilder()).Delete(own.ID) }) - ownTracks = userPls.Tracks(own.ID, false) + Expect(userPls.Put(userCtx, &own)).To(Succeed()) + DeferCleanup(func() { _ = NewPlaylistRepository(GetDBXBuilder()).Delete(adminCtx, own.ID) }) + ownTracks = userPls.Tracks(userCtx, own.ID, false) }) It("drops ids outside the user's libraries", func() { - Expect(ownTracks.Add([]string{songDayInALife.ID, "pls-otherlib-track"})).To(Equal(1)) - Expect(ownTracks.GetMediaFileIDs()).To(ConsistOf(songDayInALife.ID)) + Expect(ownTracks.Add(userCtx, []string{songDayInALife.ID, "pls-otherlib-track"})).To(Equal(1)) + Expect(ownTracks.GetMediaFileIDs(userCtx)).To(ConsistOf(songDayInALife.ID)) }) It("drops them when reached through AddAlbums", func() { - Expect(ownTracks.AddAlbums([]string{"pls-hidden-album"})).To(BeZero()) + Expect(ownTracks.AddAlbums(userCtx, []string{"pls-hidden-album"})).To(BeZero()) }) It("drops them when reached through Insert", func() { - Expect(ownTracks.Add([]string{songDayInALife.ID})).To(Equal(1)) + Expect(ownTracks.Add(userCtx, []string{songDayInALife.ID})).To(Equal(1)) - Expect(ownTracks.Insert([]string{"pls-otherlib-track", songComeTogether.ID}, 1)).To(Equal(1)) + Expect(ownTracks.Insert(userCtx, []string{"pls-otherlib-track", songComeTogether.ID}, 1)).To(Equal(1)) - Expect(ownTracks.GetMediaFileIDs()).To(Equal([]string{songComeTogether.ID, songDayInALife.ID})) - trks, err := ownTracks.GetAll(model.QueryOptions{Sort: "id"}) + Expect(ownTracks.GetMediaFileIDs(userCtx)).To(Equal([]string{songComeTogether.ID, songDayInALife.ID})) + trks, err := ownTracks.GetAll(userCtx, model.QueryOptions{Sort: "id"}) Expect(err).ToNot(HaveOccurred()) Expect(slice.Map(trks, func(t model.PlaylistTrack) string { return t.ID })).To(Equal([]string{"1", "2"}), "positions must stay contiguous when an id is dropped") @@ -307,7 +302,7 @@ var _ = Describe("PlaylistTrackRepository", func() { Describe("Put", func() { storedIDs := func(id string) []string { - ids, err := NewPlaylistRepository(adminCtx, GetDBXBuilder()).Tracks(id, false).GetMediaFileIDs() + ids, err := NewPlaylistRepository(GetDBXBuilder()).Tracks(adminCtx, id, false).GetMediaFileIDs(adminCtx) Expect(err).ToNot(HaveOccurred()) return ids } @@ -315,8 +310,8 @@ var _ = Describe("PlaylistTrackRepository", func() { pls.OwnerID = owner.ID pls.Tracks = nil pls.AddMediaFilesByID(ids) - Expect(NewPlaylistRepository(ctx, GetDBXBuilder()).Put(pls)).To(Succeed()) - DeferCleanup(func() { _ = NewPlaylistRepository(adminCtx, GetDBXBuilder()).Delete(pls.ID) }) + Expect(NewPlaylistRepository(GetDBXBuilder()).Put(ctx, pls)).To(Succeed()) + DeferCleanup(func() { _ = NewPlaylistRepository(GetDBXBuilder()).Delete(adminCtx, pls.ID) }) return pls.ID } @@ -339,10 +334,10 @@ var _ = Describe("PlaylistTrackRepository", func() { hidden := put(userCtx, restrictedUser, &model.Playlist{Name: "Hidden"}, songDayInALife.ID, "pls-otherlib-track") unknown := put(userCtx, restrictedUser, &model.Playlist{Name: "Unknown"}, songDayInALife.ID, "no-such-track") - userPls := NewPlaylistRepository(userCtx, GetDBXBuilder()) - h, err := userPls.Get(hidden) + userPls := NewPlaylistRepository(GetDBXBuilder()) + h, err := userPls.Get(userCtx, hidden) Expect(err).ToNot(HaveOccurred()) - u, err := userPls.Get(unknown) + u, err := userPls.Get(userCtx, unknown) Expect(err).ToNot(HaveOccurred()) Expect(h.SongCount).To(Equal(u.SongCount)) Expect(h.Duration).To(Equal(u.Duration)) @@ -364,9 +359,9 @@ var _ = Describe("PlaylistTrackRepository", func() { }) It("still shows everything to an admin", func() { - adminTracks := NewPlaylistRepository(adminCtx, GetDBXBuilder()).Tracks(plsID, false) - Expect(adminTracks.Count()).To(Equal(int64(2))) - _, err := adminTracks.Read("2") + adminTracks := NewPlaylistRepository(GetDBXBuilder()).Tracks(adminCtx, plsID, false) + Expect(adminTracks.Count(adminCtx)).To(Equal(int64(2))) + _, err := adminTracks.Read(adminCtx, "2") Expect(err).ToNot(HaveOccurred()) }) }) diff --git a/persistence/playqueue_repository.go b/persistence/playqueue_repository.go index ba69ec746..9dc520c44 100644 --- a/persistence/playqueue_repository.go +++ b/persistence/playqueue_repository.go @@ -17,9 +17,8 @@ type playQueueRepository struct { sqlRepository } -func NewPlayQueueRepository(ctx context.Context, db dbx.Builder) model.PlayQueueRepository { +func NewPlayQueueRepository(db dbx.Builder) model.PlayQueueRepository { r := &playQueueRepository{} - r.ctx = ctx r.db = db r.tableName = "playqueue" return r @@ -36,13 +35,13 @@ type playQueue struct { UpdatedAt time.Time `structs:"updated_at"` } -func (r *playQueueRepository) Store(q *model.PlayQueue, colNames ...string) error { - u := loggedUser(r.ctx) +func (r *playQueueRepository) Store(ctx context.Context, q *model.PlayQueue, colNames ...string) error { + u := loggedUser(ctx) // Always find existing playqueue for this user - existingQueue, err := r.Retrieve(q.UserID) + existingQueue, err := r.Retrieve(ctx, q.UserID) if err != nil && !errors.Is(err, model.ErrNotFound) { - log.Error(r.ctx, "Error retrieving existing playqueue", "user", u.UserName, err) + log.Error(ctx, "Error retrieving existing playqueue", "user", u.UserName, err) return err } @@ -53,9 +52,9 @@ func (r *playQueueRepository) Store(q *model.PlayQueue, colNames ...string) erro // When no specific columns are provided, we replace the whole queue if len(colNames) == 0 { - err := r.clearPlayQueue(q.UserID) + err := r.clearPlayQueue(ctx, q.UserID) if err != nil { - log.Error(r.ctx, "Error deleting previous playqueue", "user", u.UserName, err) + log.Error(ctx, "Error deleting previous playqueue", "user", u.UserName, err) return err } if len(q.Items) == 0 { @@ -68,27 +67,27 @@ func (r *playQueueRepository) Store(q *model.PlayQueue, colNames ...string) erro pq.CreatedAt = time.Now() } pq.UpdatedAt = time.Now() - _, err = r.put(pq.ID, pq, colNames...) + _, err = r.put(ctx, pq.ID, pq, colNames...) if err != nil { - log.Error(r.ctx, "Error saving playqueue", "user", u.UserName, err) + log.Error(ctx, "Error saving playqueue", "user", u.UserName, err) return err } return nil } -func (r *playQueueRepository) RetrieveWithMediaFiles(userId string) (*model.PlayQueue, error) { - sel := r.newSelect().Columns("*").Where(Eq{"user_id": userId}) +func (r *playQueueRepository) RetrieveWithMediaFiles(ctx context.Context, userId string) (*model.PlayQueue, error) { + sel := r.newSelect(ctx).Columns("*").Where(Eq{"user_id": userId}) var res playQueue - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) q := r.toModel(&res) - q.Items = r.loadTracks(q.Items) + q.Items = r.loadTracks(ctx, q.Items) return &q, err } -func (r *playQueueRepository) Retrieve(userId string) (*model.PlayQueue, error) { - sel := r.newSelect().Columns("*").Where(Eq{"user_id": userId}) +func (r *playQueueRepository) Retrieve(ctx context.Context, userId string) (*model.PlayQueue, error) { + sel := r.newSelect(ctx).Columns("*").Where(Eq{"user_id": userId}) var res playQueue - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return new(r.toModel(&res)), err } @@ -131,12 +130,12 @@ func (r *playQueueRepository) toModel(pq *playQueue) model.PlayQueue { // loadTracks loads the tracks from the database. It receives a list of track IDs and returns a list of MediaFiles // in the same order as the input list. -func (r *playQueueRepository) loadTracks(tracks model.MediaFiles) model.MediaFiles { +func (r *playQueueRepository) loadTracks(ctx context.Context, tracks model.MediaFiles) model.MediaFiles { if len(tracks) == 0 { return nil } - mfRepo := NewMediaFileRepository(r.ctx, r.db) + mfRepo := NewMediaFileRepository(r.db) trackMap := map[string]model.MediaFile{} // Create an iterator to collect all track IDs @@ -145,10 +144,10 @@ func (r *playQueueRepository) loadTracks(tracks model.MediaFiles) model.MediaFil // Break the list in chunks, up to 500 items, to avoid hitting SQLITE_MAX_VARIABLE_NUMBER limit for chunk := range slice.CollectChunks(ids, 500) { idsFilter := Eq{"media_file.id": chunk} - tracks, err := mfRepo.GetAll(model.QueryOptions{Filters: idsFilter}) + tracks, err := mfRepo.GetAll(ctx, model.QueryOptions{Filters: idsFilter}) if err != nil { - u := loggedUser(r.ctx) - log.Error(r.ctx, "Could not load playqueue/bookmark's tracks", "user", u.UserName, err) + u := loggedUser(ctx) + log.Error(ctx, "Could not load playqueue/bookmark's tracks", "user", u.UserName, err) } for _, t := range tracks { trackMap[t.ID] = t @@ -166,12 +165,12 @@ func (r *playQueueRepository) loadTracks(tracks model.MediaFiles) model.MediaFil return newTracks } -func (r *playQueueRepository) clearPlayQueue(userId string) error { - return r.delete(Eq{"user_id": userId}) +func (r *playQueueRepository) clearPlayQueue(ctx context.Context, userId string) error { + return r.delete(ctx, Eq{"user_id": userId}) } -func (r *playQueueRepository) Clear(userId string) error { - return r.clearPlayQueue(userId) +func (r *playQueueRepository) Clear(ctx context.Context, userId string) error { + return r.clearPlayQueue(ctx, userId) } var _ model.PlayQueueRepository = (*playQueueRepository)(nil) diff --git a/persistence/playqueue_repository_test.go b/persistence/playqueue_repository_test.go index 2bcc88fd0..877faddfc 100644 --- a/persistence/playqueue_repository_test.go +++ b/persistence/playqueue_repository_test.go @@ -22,15 +22,15 @@ var _ = Describe("PlayQueueRepository", func() { DeferCleanup(configtest.SetupConfig()) ctx = log.NewContext(context.TODO()) ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewPlayQueueRepository(ctx, GetDBXBuilder()) + repo = NewPlayQueueRepository(GetDBXBuilder()) }) Describe("Store", func() { It("stores a complete playqueue", func() { expected := aPlayQueue("userid", 1, 123, songComeTogether, songDayInALife) - Expect(repo.Store(expected)).To(Succeed()) + Expect(repo.Store(ctx, expected)).To(Succeed()) - actual, err := repo.RetrieveWithMediaFiles("userid") + actual, err := repo.RetrieveWithMediaFiles(ctx, "userid") Expect(err).ToNot(HaveOccurred()) AssertPlayQueue(expected, actual) Expect(countPlayQueues(repo, "userid")).To(Equal(1)) @@ -39,13 +39,13 @@ var _ = Describe("PlayQueueRepository", func() { It("replaces existing playqueue when storing without column names", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) By("Storing replacement playqueue") replacement := aPlayQueue("userid", 1, 200, songDayInALife, songAntenna) - Expect(repo.Store(replacement)).To(Succeed()) + Expect(repo.Store(ctx, replacement)).To(Succeed()) - actual, err := repo.RetrieveWithMediaFiles("userid") + actual, err := repo.RetrieveWithMediaFiles(ctx, "userid") Expect(err).ToNot(HaveOccurred()) AssertPlayQueue(replacement, actual) Expect(countPlayQueues(repo, "userid")).To(Equal(1)) @@ -54,24 +54,24 @@ var _ = Describe("PlayQueueRepository", func() { It("clears playqueue when storing empty items", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) By("Storing empty playqueue") empty := aPlayQueue("userid", 0, 0) - Expect(repo.Store(empty)).To(Succeed()) + Expect(repo.Store(ctx, empty)).To(Succeed()) By("Verifying playqueue is cleared") - _, err := repo.Retrieve("userid") + _, err := repo.Retrieve(ctx, "userid") Expect(err).To(MatchError(model.ErrNotFound)) }) It("updates only current field when specified", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether, songDayInALife) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) By("Getting the existing playqueue to obtain its ID") - existing, err := repo.Retrieve("userid") + existing, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) By("Updating only current field") @@ -81,10 +81,10 @@ var _ = Describe("PlayQueueRepository", func() { Current: 1, ChangedBy: "test-update", } - Expect(repo.Store(update, "current")).To(Succeed()) + Expect(repo.Store(ctx, update, "current")).To(Succeed()) By("Verifying only current was updated") - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) Expect(actual.Current).To(Equal(1)) Expect(actual.Position).To(Equal(int64(100))) // Should remain unchanged @@ -94,10 +94,10 @@ var _ = Describe("PlayQueueRepository", func() { It("updates only position field when specified", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 1, 100, songComeTogether, songDayInALife) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) By("Getting the existing playqueue to obtain its ID") - existing, err := repo.Retrieve("userid") + existing, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) By("Updating only position field") @@ -107,10 +107,10 @@ var _ = Describe("PlayQueueRepository", func() { Position: 500, ChangedBy: "test-update", } - Expect(repo.Store(update, "position")).To(Succeed()) + Expect(repo.Store(ctx, update, "position")).To(Succeed()) By("Verifying only position was updated") - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) Expect(actual.Position).To(Equal(int64(500))) Expect(actual.Current).To(Equal(1)) // Should remain unchanged @@ -120,10 +120,10 @@ var _ = Describe("PlayQueueRepository", func() { It("updates multiple specified fields", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) By("Getting the existing playqueue to obtain its ID") - existing, err := repo.Retrieve("userid") + existing, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) By("Updating current and position fields") @@ -134,10 +134,10 @@ var _ = Describe("PlayQueueRepository", func() { Position: 300, ChangedBy: "test-update", } - Expect(repo.Store(update, "current", "position")).To(Succeed()) + Expect(repo.Store(ctx, update, "current", "position")).To(Succeed()) By("Verifying both fields were updated") - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) Expect(actual.Current).To(Equal(1)) Expect(actual.Position).To(Equal(int64(300))) @@ -147,10 +147,10 @@ var _ = Describe("PlayQueueRepository", func() { It("preserves existing data when updating with empty items list and column names", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether, songDayInALife) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) By("Getting the existing playqueue to obtain its ID") - existing, err := repo.Retrieve("userid") + existing, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) By("Updating only position with empty items") @@ -161,10 +161,10 @@ var _ = Describe("PlayQueueRepository", func() { ChangedBy: "test-update", Items: []model.MediaFile{}, // Empty items } - Expect(repo.Store(update, "position")).To(Succeed()) + Expect(repo.Store(ctx, update, "position")).To(Succeed()) By("Verifying items are preserved") - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) Expect(actual.Position).To(Equal(int64(200))) Expect(actual.Items).To(HaveLen(2)) // Should remain unchanged @@ -173,21 +173,21 @@ var _ = Describe("PlayQueueRepository", func() { It("ensures only one record per user by reusing existing record ID", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) initialCount := countPlayQueues(repo, "userid") Expect(initialCount).To(Equal(1)) By("Storing another playqueue with different ID but same user") different := aPlayQueue("userid", 1, 200, songDayInALife) different.ID = "different-id" // Force a different ID - Expect(repo.Store(different)).To(Succeed()) + Expect(repo.Store(ctx, different)).To(Succeed()) By("Verifying only one record exists for the user") finalCount := countPlayQueues(repo, "userid") Expect(finalCount).To(Equal(1)) By("Verifying the record was updated, not duplicated") - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) Expect(actual.Current).To(Equal(1)) // Should be updated value Expect(actual.Position).To(Equal(int64(200))) // Should be updated value @@ -198,7 +198,7 @@ var _ = Describe("PlayQueueRepository", func() { It("ensures only one record per user even with partial updates", func() { By("Storing initial playqueue") initial := aPlayQueue("userid", 0, 100, songComeTogether, songDayInALife) - Expect(repo.Store(initial)).To(Succeed()) + Expect(repo.Store(ctx, initial)).To(Succeed()) initialCount := countPlayQueues(repo, "userid") Expect(initialCount).To(Equal(1)) @@ -209,14 +209,14 @@ var _ = Describe("PlayQueueRepository", func() { Current: 1, ChangedBy: "test-partial", } - Expect(repo.Store(partialUpdate, "current")).To(Succeed()) + Expect(repo.Store(ctx, partialUpdate, "current")).To(Succeed()) By("Verifying only one record still exists for the user") finalCount := countPlayQueues(repo, "userid") Expect(finalCount).To(Equal(1)) By("Verifying the existing record was updated with new current value") - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) Expect(actual.Current).To(Equal(1)) // Should be updated value Expect(actual.Position).To(Equal(int64(100))) // Should remain unchanged @@ -226,7 +226,7 @@ var _ = Describe("PlayQueueRepository", func() { Describe("Retrieve", func() { It("returns notfound error if there's no playqueue for the user", func() { - _, err := repo.Retrieve("user999") + _, err := repo.Retrieve(ctx, "user999") Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -234,9 +234,9 @@ var _ = Describe("PlayQueueRepository", func() { By("Storing a playqueue for the user") expected := aPlayQueue("userid", 1, 123, songComeTogether, songDayInALife) - Expect(repo.Store(expected)).To(Succeed()) + Expect(repo.Store(ctx, expected)).To(Succeed()) - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) // Basic playqueue properties should match @@ -263,19 +263,19 @@ var _ = Describe("PlayQueueRepository", func() { newSong := songRadioactivity newSong.ID = "temp-track" newSong.Path = "/new-path" - mfRepo := NewMediaFileRepository(ctx, GetDBXBuilder()) + mfRepo := NewMediaFileRepository(GetDBXBuilder()) - Expect(mfRepo.Put(&newSong)).To(Succeed()) + Expect(mfRepo.Put(ctx, &newSong)).To(Succeed()) // Create a playqueue with the new song pq := aPlayQueue("userid", 0, 0, newSong, songAntenna) - Expect(repo.Store(pq)).To(Succeed()) + Expect(repo.Store(ctx, pq)).To(Succeed()) // Delete the new song from the database - Expect(mfRepo.Delete("temp-track")).To(Succeed()) + Expect(mfRepo.Delete(ctx, "temp-track")).To(Succeed()) // Retrieve the playqueue with Retrieve method - actual, err := repo.Retrieve("userid") + actual, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) // The playqueue should still contain both track IDs (including the deleted one) @@ -295,7 +295,7 @@ var _ = Describe("PlayQueueRepository", func() { Describe("RetrieveWithMediaFiles", func() { It("returns notfound error if there's no playqueue for the user", func() { - _, err := repo.RetrieveWithMediaFiles("user999") + _, err := repo.RetrieveWithMediaFiles(ctx, "user999") Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -303,9 +303,9 @@ var _ = Describe("PlayQueueRepository", func() { By("Storing a playqueue for the user") expected := aPlayQueue("userid", 1, 123, songComeTogether, songDayInALife) - Expect(repo.Store(expected)).To(Succeed()) + Expect(repo.Store(ctx, expected)).To(Succeed()) - actual, err := repo.RetrieveWithMediaFiles("userid") + actual, err := repo.RetrieveWithMediaFiles(ctx, "userid") Expect(err).ToNot(HaveOccurred()) AssertPlayQueue(expected, actual) @@ -316,26 +316,26 @@ var _ = Describe("PlayQueueRepository", func() { newSong := songRadioactivity newSong.ID = "temp-track" newSong.Path = "/new-path" - mfRepo := NewMediaFileRepository(ctx, GetDBXBuilder()) + mfRepo := NewMediaFileRepository(GetDBXBuilder()) - Expect(mfRepo.Put(&newSong)).To(Succeed()) + Expect(mfRepo.Put(ctx, &newSong)).To(Succeed()) // Create a playqueue with the new song pq := aPlayQueue("userid", 0, 0, newSong, songAntenna) - Expect(repo.Store(pq)).To(Succeed()) + Expect(repo.Store(ctx, pq)).To(Succeed()) // Retrieve the playqueue - actual, err := repo.RetrieveWithMediaFiles("userid") + actual, err := repo.RetrieveWithMediaFiles(ctx, "userid") Expect(err).ToNot(HaveOccurred()) // The playqueue should contain both tracks AssertPlayQueue(pq, actual) // Delete the new song - Expect(mfRepo.Delete("temp-track")).To(Succeed()) + Expect(mfRepo.Delete(ctx, "temp-track")).To(Succeed()) // Retrieve the playqueue - actual, err = repo.RetrieveWithMediaFiles("userid") + actual, err = repo.RetrieveWithMediaFiles(ctx, "userid") Expect(err).ToNot(HaveOccurred()) // The playqueue should not contain the deleted track @@ -348,48 +348,48 @@ var _ = Describe("PlayQueueRepository", func() { It("clears an existing playqueue", func() { By("Storing a playqueue") expected := aPlayQueue("userid", 1, 123, songComeTogether, songDayInALife) - Expect(repo.Store(expected)).To(Succeed()) + Expect(repo.Store(ctx, expected)).To(Succeed()) By("Verifying playqueue exists") - _, err := repo.Retrieve("userid") + _, err := repo.Retrieve(ctx, "userid") Expect(err).ToNot(HaveOccurred()) By("Clearing the playqueue") - Expect(repo.Clear("userid")).To(Succeed()) + Expect(repo.Clear(ctx, "userid")).To(Succeed()) By("Verifying playqueue is cleared") - _, err = repo.Retrieve("userid") + _, err = repo.Retrieve(ctx, "userid") Expect(err).To(MatchError(model.ErrNotFound)) }) It("does not error when clearing non-existent playqueue", func() { // Clear should not error even if no playqueue exists - Expect(repo.Clear("nonexistent-user")).To(Succeed()) + Expect(repo.Clear(ctx, "nonexistent-user")).To(Succeed()) }) It("only clears the specified user's playqueue", func() { By("Creating users in the database to avoid foreign key constraints") - userRepo := NewUserRepository(ctx, GetDBXBuilder()) + userRepo := NewUserRepository(GetDBXBuilder()) user1 := &model.User{ID: "user1", UserName: "user1", Name: "User 1", Email: "user1@test.com"} user2 := &model.User{ID: "user2", UserName: "user2", Name: "User 2", Email: "user2@test.com"} - Expect(userRepo.Put(user1)).To(Succeed()) - Expect(userRepo.Put(user2)).To(Succeed()) + Expect(userRepo.Put(ctx, user1)).To(Succeed()) + Expect(userRepo.Put(ctx, user2)).To(Succeed()) By("Storing playqueues for two users") user1Queue := aPlayQueue("user1", 0, 100, songComeTogether) user2Queue := aPlayQueue("user2", 1, 200, songDayInALife) - Expect(repo.Store(user1Queue)).To(Succeed()) - Expect(repo.Store(user2Queue)).To(Succeed()) + Expect(repo.Store(ctx, user1Queue)).To(Succeed()) + Expect(repo.Store(ctx, user2Queue)).To(Succeed()) By("Clearing only user1's playqueue") - Expect(repo.Clear("user1")).To(Succeed()) + Expect(repo.Clear(ctx, "user1")).To(Succeed()) By("Verifying user1's playqueue is cleared") - _, err := repo.Retrieve("user1") + _, err := repo.Retrieve(ctx, "user1") Expect(err).To(MatchError(model.ErrNotFound)) By("Verifying user2's playqueue still exists") - actual, err := repo.Retrieve("user2") + actual, err := repo.Retrieve(ctx, "user2") Expect(err).ToNot(HaveOccurred()) Expect(actual.UserID).To(Equal("user2")) Expect(actual.Current).To(Equal(1)) @@ -400,7 +400,7 @@ var _ = Describe("PlayQueueRepository", func() { func countPlayQueues(repo model.PlayQueueRepository, userId string) int { r := repo.(*playQueueRepository) - c, err := r.count(squirrel.Select().Where(squirrel.Eq{"user_id": userId})) + c, err := r.count(GinkgoT().Context(), squirrel.Select().Where(squirrel.Eq{"user_id": userId})) if err != nil { panic(err) } diff --git a/persistence/plugin_cleanup_test.go b/persistence/plugin_cleanup_test.go index bfe6d60ca..08959075d 100644 --- a/persistence/plugin_cleanup_test.go +++ b/persistence/plugin_cleanup_test.go @@ -1,6 +1,8 @@ package persistence import ( + "context" + "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" . "github.com/onsi/ginkgo/v2" @@ -11,30 +13,47 @@ var _ = Describe("Plugin Cleanup", func() { var pluginRepo model.PluginRepository var userRepo model.UserRepository var libraryRepo model.LibraryRepository + var ctx context.Context BeforeEach(func() { - ctx := GinkgoT().Context() - ctx = request.WithUser(ctx, model.User{ID: "admin", UserName: "admin", IsAdmin: true}) + ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "admin", UserName: "admin", IsAdmin: true}) db := GetDBXBuilder() - pluginRepo = NewPluginRepository(ctx, db) - userRepo = NewUserRepository(ctx, db) - libraryRepo = NewLibraryRepository(ctx, db) + pluginRepo = NewPluginRepository(db) + userRepo = NewUserRepository(db) + libraryRepo = NewLibraryRepository(db) // Clean up any existing plugins - all, _ := pluginRepo.GetAll() + all, _ := pluginRepo.GetAll(ctx) for _, p := range all { - _ = pluginRepo.Delete(p.ID) + _ = pluginRepo.Delete(ctx, p.ID) } }) AfterEach(func() { // Clean up after tests - all, _ := pluginRepo.GetAll() + all, _ := pluginRepo.GetAll(ctx) for _, p := range all { - _ = pluginRepo.Delete(p.ID) + _ = pluginRepo.Delete(ctx, p.ID) } }) + Describe("UserRepository.Delete", func() { + It("cleans up plugin references for users deleted before a later id fails", func() { + Expect(userRepo.Put(ctx, &model.User{ID: "bulk-1", UserName: "bulk-1", NewPassword: "x"})).To(Succeed()) + DeferCleanup(func() { _ = userRepo.Delete(ctx, "bulk-1") }) + Expect(pluginRepo.Put(ctx, &model.Plugin{ + ID: "bulk-plugin", Path: "/plugins/bulk.wasm", Manifest: `{"name":"bulk"}`, SHA256: "def456", + Users: `["bulk-1","other"]`, Enabled: true, + })).To(Succeed()) + + Expect(userRepo.Delete(ctx, "bulk-1", "does-not-exist")).To(MatchError(model.ErrNotFound)) + + updated, err := pluginRepo.Get(ctx, "bulk-plugin") + Expect(err).ToNot(HaveOccurred()) + Expect(updated.Users).To(Equal(`["other"]`)) + }) + }) + Describe("cleanupPluginUserReferences", func() { It("removes user ID from plugin users array", func() { // Create a plugin with multiple users @@ -46,14 +65,14 @@ var _ = Describe("Plugin Cleanup", func() { Users: `["user1","user2","user3"]`, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Clean up user2 reference db := GetDBXBuilder() Expect(cleanupPluginUserReferences(db, "user2")).To(Succeed()) // Verify user2 was removed - updated, err := pluginRepo.Get("test-plugin") + updated, err := pluginRepo.Get(ctx, "test-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Users).To(Equal(`["user1","user3"]`)) Expect(updated.Enabled).To(BeTrue()) // Still has users, should remain enabled @@ -70,14 +89,14 @@ var _ = Describe("Plugin Cleanup", func() { AllUsers: false, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Remove the only user db := GetDBXBuilder() Expect(cleanupPluginUserReferences(db, "only-user")).To(Succeed()) // Verify plugin was auto-disabled - updated, err := pluginRepo.Get("user-plugin") + updated, err := pluginRepo.Get(ctx, "user-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Users).To(Equal(`[]`)) Expect(updated.Enabled).To(BeFalse()) @@ -93,14 +112,14 @@ var _ = Describe("Plugin Cleanup", func() { AllUsers: true, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Remove the user (but allUsers is true) db := GetDBXBuilder() Expect(cleanupPluginUserReferences(db, "user1")).To(Succeed()) // Plugin should still be enabled because allUsers is true - updated, err := pluginRepo.Get("all-users-plugin") + updated, err := pluginRepo.Get(ctx, "all-users-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Enabled).To(BeTrue()) }) @@ -114,14 +133,14 @@ var _ = Describe("Plugin Cleanup", func() { Users: `["user1"]`, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Remove the user db := GetDBXBuilder() Expect(cleanupPluginUserReferences(db, "user1")).To(Succeed()) // Plugin should still be enabled (no users permission requirement) - updated, err := pluginRepo.Get("no-users-perm") + updated, err := pluginRepo.Get(ctx, "no-users-perm") Expect(err).ToNot(HaveOccurred()) Expect(updated.Users).To(Equal(`[]`)) Expect(updated.Enabled).To(BeTrue()) @@ -139,14 +158,14 @@ var _ = Describe("Plugin Cleanup", func() { Libraries: `[1,2,3]`, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Clean up library 2 reference db := GetDBXBuilder() Expect(cleanupPluginLibraryReferences(db, 2)).To(Succeed()) // Verify library 2 was removed - updated, err := pluginRepo.Get("lib-plugin") + updated, err := pluginRepo.Get(ctx, "lib-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Libraries).To(Equal(`[1,3]`)) }) @@ -162,14 +181,14 @@ var _ = Describe("Plugin Cleanup", func() { AllLibraries: false, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Remove the only library db := GetDBXBuilder() Expect(cleanupPluginLibraryReferences(db, 99)).To(Succeed()) // Verify plugin was auto-disabled - updated, err := pluginRepo.Get("lib-only-plugin") + updated, err := pluginRepo.Get(ctx, "lib-only-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Libraries).To(Equal(`[]`)) Expect(updated.Enabled).To(BeFalse()) @@ -185,14 +204,14 @@ var _ = Describe("Plugin Cleanup", func() { AllLibraries: true, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Remove the library (but allLibraries is true) db := GetDBXBuilder() Expect(cleanupPluginLibraryReferences(db, 1)).To(Succeed()) // Plugin should still be enabled - updated, err := pluginRepo.Get("all-libs-plugin") + updated, err := pluginRepo.Get(ctx, "all-libs-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Enabled).To(BeTrue()) }) @@ -207,7 +226,7 @@ var _ = Describe("Plugin Cleanup", func() { IsAdmin: false, } user.NewPassword = "password123" - Expect(userRepo.Put(user)).To(Succeed()) + Expect(userRepo.Put(ctx, user)).To(Succeed()) // Create a plugin referencing this user plugin := &model.Plugin{ @@ -218,13 +237,13 @@ var _ = Describe("Plugin Cleanup", func() { Users: `["test-delete-user","other-user"]`, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Delete the user - Expect(userRepo.Delete("test-delete-user")).To(Succeed()) + Expect(userRepo.Delete(ctx, "test-delete-user")).To(Succeed()) // Verify user was removed from plugin - updated, err := pluginRepo.Get("user-ref-plugin") + updated, err := pluginRepo.Get(ctx, "user-ref-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Users).To(Equal(`["other-user"]`)) }) @@ -238,7 +257,7 @@ var _ = Describe("Plugin Cleanup", func() { Name: "Test Library", Path: "/tmp/test-lib", } - Expect(libraryRepo.Put(library)).To(Succeed()) + Expect(libraryRepo.Put(ctx, library)).To(Succeed()) // Create a plugin referencing this library plugin := &model.Plugin{ @@ -249,13 +268,13 @@ var _ = Describe("Plugin Cleanup", func() { Libraries: `[99,1]`, Enabled: true, } - Expect(pluginRepo.Put(plugin)).To(Succeed()) + Expect(pluginRepo.Put(ctx, plugin)).To(Succeed()) // Delete the library - Expect(libraryRepo.Delete(99)).To(Succeed()) + Expect(libraryRepo.Delete(ctx, 99)).To(Succeed()) // Verify library was removed from plugin - updated, err := pluginRepo.Get("lib-ref-plugin") + updated, err := pluginRepo.Get(ctx, "lib-ref-plugin") Expect(err).ToNot(HaveOccurred()) Expect(updated.Libraries).To(Equal(`[1]`)) }) diff --git a/persistence/plugin_repository.go b/persistence/plugin_repository.go index 7d5781f49..1545ce2ba 100644 --- a/persistence/plugin_repository.go +++ b/persistence/plugin_repository.go @@ -15,9 +15,8 @@ type pluginRepository struct { sqlRepository } -func NewPluginRepository(ctx context.Context, db dbx.Builder) model.PluginRepository { +func NewPluginRepository(db dbx.Builder) model.PluginRepository { r := &pluginRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Plugin{}, map[string]filterFunc{ "id": idFilter("plugin"), @@ -26,13 +25,13 @@ func NewPluginRepository(ctx context.Context, db dbx.Builder) model.PluginReposi return r } -func (r *pluginRepository) isPermitted() bool { - user := loggedUser(r.ctx) +func (r *pluginRepository) isPermitted(ctx context.Context) bool { + user := loggedUser(ctx) return user.IsAdmin } -func (r *pluginRepository) ClearErrors() error { - if !r.isPermitted() { +func (r *pluginRepository) ClearErrors(ctx context.Context) error { + if !r.isPermitted(ctx) { return rest.ErrPermissionDenied } // An UPDATE takes the write lock even when nothing matches, so only run it when there is an error to clear @@ -47,43 +46,43 @@ func (r *pluginRepository) ClearErrors() error { return err } -func (r *pluginRepository) CountAll(options ...model.QueryOptions) (int64, error) { - if !r.isPermitted() { +func (r *pluginRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + if !r.isPermitted(ctx) { return 0, rest.ErrPermissionDenied } - sql := r.newSelect() - return r.count(sql, options...) + sql := r.newSelect(ctx) + return r.count(ctx, sql, options...) } -func (r *pluginRepository) Delete(id string) error { - if !r.isPermitted() { +func (r *pluginRepository) Delete(ctx context.Context, id string) error { + if !r.isPermitted(ctx) { return rest.ErrPermissionDenied } - return r.delete(Eq{"id": id}) + return r.delete(ctx, Eq{"id": id}) } -func (r *pluginRepository) Get(id string) (*model.Plugin, error) { - if !r.isPermitted() { +func (r *pluginRepository) Get(ctx context.Context, id string) (*model.Plugin, error) { + if !r.isPermitted(ctx) { return nil, rest.ErrPermissionDenied } - sel := r.newSelect().Where(Eq{"id": id}).Columns("*") + sel := r.newSelect(ctx).Where(Eq{"id": id}).Columns("*") res := model.Plugin{} - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *pluginRepository) GetAll(options ...model.QueryOptions) (model.Plugins, error) { - if !r.isPermitted() { +func (r *pluginRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Plugins, error) { + if !r.isPermitted(ctx) { return nil, rest.ErrPermissionDenied } - sel := r.newSelect(options...).Columns("*") + sel := r.newSelect(ctx, options...).Columns("*") res := model.Plugins{} - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) return res, err } -func (r *pluginRepository) Put(plugin *model.Plugin) error { - if !r.isPermitted() { +func (r *pluginRepository) Put(ctx context.Context, plugin *model.Plugin) error { + if !r.isPermitted(ctx) { return rest.ErrPermissionDenied } @@ -129,25 +128,17 @@ func (r *pluginRepository) Put(plugin *model.Plugin) error { return err } -func (r *pluginRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *pluginRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *pluginRepository) EntityName() string { - return "plugin" +func (r *pluginRepository) Read(ctx context.Context, id string) (*model.Plugin, error) { + return r.Get(ctx, id) } -func (r *pluginRepository) NewInstance() any { - return &model.Plugin{} -} - -func (r *pluginRepository) Read(id string) (any, error) { - return r.Get(id) -} - -func (r *pluginRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) +func (r *pluginRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Plugin, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } var _ model.PluginRepository = (*pluginRepository)(nil) -var _ rest.Repository = (*pluginRepository)(nil) +var _ rest.Repository[model.Plugin] = (*pluginRepository)(nil) diff --git a/persistence/plugin_repository_test.go b/persistence/plugin_repository_test.go index 9b135057e..44330250b 100644 --- a/persistence/plugin_repository_test.go +++ b/persistence/plugin_repository_test.go @@ -13,50 +13,54 @@ import ( var _ = Describe("PluginRepository", func() { var repo model.PluginRepository + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) Describe("Admin User", func() { BeforeEach(func() { - ctx := GinkgoT().Context() ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewPluginRepository(ctx, GetDBXBuilder()) + repo = NewPluginRepository(GetDBXBuilder()) // Clean up any existing plugins - all, _ := repo.GetAll() + all, _ := repo.GetAll(ctx) for _, p := range all { - _ = repo.Delete(p.ID) + _ = repo.Delete(ctx, p.ID) } }) AfterEach(func() { // Clean up after tests - all, _ := repo.GetAll() + all, _ := repo.GetAll(ctx) for _, p := range all { - _ = repo.Delete(p.ID) + _ = repo.Delete(ctx, p.ID) } }) Describe("CountAll", func() { It("returns 0 when no plugins exist", func() { - Expect(repo.CountAll()).To(Equal(int64(0))) + Expect(repo.CountAll(ctx)).To(Equal(int64(0))) }) It("returns the number of plugins in the DB", func() { - _ = repo.Put(&model.Plugin{ID: "test-plugin-1", Path: "/plugins/test1.wasm", Manifest: "{}", SHA256: "abc123"}) - _ = repo.Put(&model.Plugin{ID: "test-plugin-2", Path: "/plugins/test2.wasm", Manifest: "{}", SHA256: "def456"}) + _ = repo.Put(ctx, &model.Plugin{ID: "test-plugin-1", Path: "/plugins/test1.wasm", Manifest: "{}", SHA256: "abc123"}) + _ = repo.Put(ctx, &model.Plugin{ID: "test-plugin-2", Path: "/plugins/test2.wasm", Manifest: "{}", SHA256: "def456"}) - Expect(repo.CountAll()).To(Equal(int64(2))) + Expect(repo.CountAll(ctx)).To(Equal(int64(2))) }) }) Describe("Delete", func() { It("deletes existing item", func() { plugin := &model.Plugin{ID: "to-delete", Path: "/plugins/delete.wasm", Manifest: "{}", SHA256: "hash"} - _ = repo.Put(plugin) + _ = repo.Put(ctx, plugin) - err := repo.Delete(plugin.ID) + err := repo.Delete(ctx, plugin.ID) Expect(err).To(BeNil()) - _, err = repo.Get(plugin.ID) + _, err = repo.Get(ctx, plugin.ID) Expect(err).To(MatchError(model.ErrNotFound)) }) }) @@ -64,9 +68,9 @@ var _ = Describe("PluginRepository", func() { Describe("Get", func() { It("returns an existing item", func() { plugin := &model.Plugin{ID: "test-get", Path: "/plugins/test.wasm", Manifest: `{"name":"test"}`, SHA256: "hash123"} - _ = repo.Put(plugin) + _ = repo.Put(ctx, plugin) - res, err := repo.Get(plugin.ID) + res, err := repo.Get(ctx, plugin.ID) Expect(err).To(BeNil()) Expect(res.ID).To(Equal(plugin.ID)) Expect(res.Path).To(Equal(plugin.Path)) @@ -74,31 +78,31 @@ var _ = Describe("PluginRepository", func() { }) It("errors when missing", func() { - _, err := repo.Get("notanid") + _, err := repo.Get(ctx, "notanid") Expect(err).To(MatchError(model.ErrNotFound)) }) }) Describe("GetAll", func() { It("returns all items from the DB", func() { - _ = repo.Put(&model.Plugin{ID: "plugin-a", Path: "/plugins/a.wasm", Manifest: "{}", SHA256: "hash1"}) - _ = repo.Put(&model.Plugin{ID: "plugin-b", Path: "/plugins/b.wasm", Manifest: "{}", SHA256: "hash2"}) + _ = repo.Put(ctx, &model.Plugin{ID: "plugin-a", Path: "/plugins/a.wasm", Manifest: "{}", SHA256: "hash1"}) + _ = repo.Put(ctx, &model.Plugin{ID: "plugin-b", Path: "/plugins/b.wasm", Manifest: "{}", SHA256: "hash2"}) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) Expect(all).To(HaveLen(2)) }) It("supports pagination", func() { - _ = repo.Put(&model.Plugin{ID: "plugin-1", Path: "/plugins/1.wasm", Manifest: "{}", SHA256: "h1"}) - _ = repo.Put(&model.Plugin{ID: "plugin-2", Path: "/plugins/2.wasm", Manifest: "{}", SHA256: "h2"}) - _ = repo.Put(&model.Plugin{ID: "plugin-3", Path: "/plugins/3.wasm", Manifest: "{}", SHA256: "h3"}) + _ = repo.Put(ctx, &model.Plugin{ID: "plugin-1", Path: "/plugins/1.wasm", Manifest: "{}", SHA256: "h1"}) + _ = repo.Put(ctx, &model.Plugin{ID: "plugin-2", Path: "/plugins/2.wasm", Manifest: "{}", SHA256: "h2"}) + _ = repo.Put(ctx, &model.Plugin{ID: "plugin-3", Path: "/plugins/3.wasm", Manifest: "{}", SHA256: "h3"}) - page1, err := repo.GetAll(model.QueryOptions{Max: 2, Offset: 0, Sort: "id"}) + page1, err := repo.GetAll(ctx, model.QueryOptions{Max: 2, Offset: 0, Sort: "id"}) Expect(err).To(BeNil()) Expect(page1).To(HaveLen(2)) - page2, err := repo.GetAll(model.QueryOptions{Max: 2, Offset: 2, Sort: "id"}) + page2, err := repo.GetAll(ctx, model.QueryOptions{Max: 2, Offset: 2, Sort: "id"}) Expect(err).To(BeNil()) Expect(page2).To(HaveLen(1)) }) @@ -115,10 +119,10 @@ var _ = Describe("PluginRepository", func() { Enabled: false, } - err := repo.Put(plugin) + err := repo.Put(ctx, plugin) Expect(err).To(BeNil()) - saved, err := repo.Get(plugin.ID) + saved, err := repo.Get(ctx, plugin.ID) Expect(err).To(BeNil()) Expect(saved.Path).To(Equal(plugin.Path)) Expect(saved.Manifest).To(Equal(plugin.Manifest)) @@ -136,15 +140,15 @@ var _ = Describe("PluginRepository", func() { SHA256: "original", Enabled: false, } - _ = repo.Put(plugin) + _ = repo.Put(ctx, plugin) plugin.Enabled = true plugin.Config = `{"new":"config"}` plugin.SHA256 = "updated" - err := repo.Put(plugin) + err := repo.Put(ctx, plugin) Expect(err).To(BeNil()) - saved, err := repo.Get(plugin.ID) + saved, err := repo.Get(ctx, plugin.ID) Expect(err).To(BeNil()) Expect(saved.Enabled).To(BeTrue()) Expect(saved.Config).To(Equal(`{"new":"config"}`)) @@ -159,10 +163,10 @@ var _ = Describe("PluginRepository", func() { SHA256: "hash", LastError: "failed to load: missing export", } - err := repo.Put(plugin) + err := repo.Put(ctx, plugin) Expect(err).To(BeNil()) - saved, err := repo.Get(plugin.ID) + saved, err := repo.Get(ctx, plugin.ID) Expect(err).To(BeNil()) Expect(saved.LastError).To(Equal("failed to load: missing export")) }) @@ -173,7 +177,7 @@ var _ = Describe("PluginRepository", func() { Manifest: "{}", SHA256: "hash", } - err := repo.Put(plugin) + err := repo.Put(ctx, plugin) Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("ID cannot be empty")) }) @@ -181,14 +185,14 @@ var _ = Describe("PluginRepository", func() { Describe("ClearErrors", func() { It("clears last_error on all plugins with errors", func() { - _ = repo.Put(&model.Plugin{ID: "ok-plugin", Path: "/plugins/ok.wasm", Manifest: "{}", SHA256: "h1"}) - _ = repo.Put(&model.Plugin{ID: "err-plugin-1", Path: "/plugins/e1.wasm", Manifest: "{}", SHA256: "h2", LastError: "incompatible version"}) - _ = repo.Put(&model.Plugin{ID: "err-plugin-2", Path: "/plugins/e2.wasm", Manifest: "{}", SHA256: "h3", LastError: "missing export"}) + _ = repo.Put(ctx, &model.Plugin{ID: "ok-plugin", Path: "/plugins/ok.wasm", Manifest: "{}", SHA256: "h1"}) + _ = repo.Put(ctx, &model.Plugin{ID: "err-plugin-1", Path: "/plugins/e1.wasm", Manifest: "{}", SHA256: "h2", LastError: "incompatible version"}) + _ = repo.Put(ctx, &model.Plugin{ID: "err-plugin-2", Path: "/plugins/e2.wasm", Manifest: "{}", SHA256: "h3", LastError: "missing export"}) - err := repo.ClearErrors() + err := repo.ClearErrors(ctx) Expect(err).To(BeNil()) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) for _, p := range all { Expect(p.LastError).To(BeEmpty(), "plugin %s should have no error", p.ID) @@ -196,14 +200,14 @@ var _ = Describe("PluginRepository", func() { }) It("succeeds when no plugins have errors", func() { - _ = repo.Put(&model.Plugin{ID: "clean-plugin", Path: "/plugins/c.wasm", Manifest: "{}", SHA256: "h1"}) + _ = repo.Put(ctx, &model.Plugin{ID: "clean-plugin", Path: "/plugins/c.wasm", Manifest: "{}", SHA256: "h1"}) - err := repo.ClearErrors() + err := repo.ClearErrors(ctx) Expect(err).To(BeNil()) }) It("does not need the write lock when no plugins have errors", func() { - _ = repo.Put(&model.Plugin{ID: "clean-plugin", Path: "/plugins/c.wasm", Manifest: "{}", SHA256: "h1"}) + _ = repo.Put(ctx, &model.Plugin{ID: "clean-plugin", Path: "/plugins/c.wasm", Manifest: "{}", SHA256: "h1"}) conn, err := db.Db().Conn(GinkgoT().Context()) Expect(err).ToNot(HaveOccurred()) DeferCleanup(conn.Close) @@ -211,49 +215,48 @@ var _ = Describe("PluginRepository", func() { Expect(err).ToNot(HaveOccurred()) DeferCleanup(func() { _, _ = conn.ExecContext(context.Background(), "ROLLBACK") }) - Expect(repo.ClearErrors()).To(Succeed()) + Expect(repo.ClearErrors(ctx)).To(Succeed()) }) }) }) Describe("Regular User", func() { BeforeEach(func() { - ctx := GinkgoT().Context() ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: false}) - repo = NewPluginRepository(ctx, GetDBXBuilder()) + repo = NewPluginRepository(GetDBXBuilder()) }) Describe("CountAll", func() { It("fails to count items", func() { - _, err := repo.CountAll() + _, err := repo.CountAll(ctx) Expect(err).To(Equal(rest.ErrPermissionDenied)) }) }) Describe("Delete", func() { It("fails to delete items", func() { - err := repo.Delete("any-id") + err := repo.Delete(ctx, "any-id") Expect(err).To(Equal(rest.ErrPermissionDenied)) }) }) Describe("Get", func() { It("fails to get items", func() { - _, err := repo.Get("any-id") + _, err := repo.Get(ctx, "any-id") Expect(err).To(Equal(rest.ErrPermissionDenied)) }) }) Describe("GetAll", func() { It("fails to get all items", func() { - _, err := repo.GetAll() + _, err := repo.GetAll(ctx) Expect(err).To(Equal(rest.ErrPermissionDenied)) }) }) Describe("Put", func() { It("fails to create/update item", func() { - err := repo.Put(&model.Plugin{ + err := repo.Put(ctx, &model.Plugin{ ID: "user-create", Path: "/plugins/create.wasm", Manifest: "{}", diff --git a/persistence/property_repository.go b/persistence/property_repository.go index 14f9051f7..29bb2b564 100644 --- a/persistence/property_repository.go +++ b/persistence/property_repository.go @@ -13,17 +13,16 @@ type propertyRepository struct { sqlRepository } -func NewPropertyRepository(ctx context.Context, db dbx.Builder) model.PropertyRepository { +func NewPropertyRepository(db dbx.Builder) model.PropertyRepository { r := &propertyRepository{} - r.ctx = ctx r.db = db r.tableName = "property" return r } -func (r propertyRepository) Put(id string, value string) error { +func (r propertyRepository) Put(ctx context.Context, id string, value string) error { update := Update(r.tableName).Set("value", value).Where(Eq{"id": id}) - count, err := r.executeSQL(update) + count, err := r.executeSQL(ctx, update) if err != nil { return err } @@ -31,24 +30,24 @@ func (r propertyRepository) Put(id string, value string) error { return nil } insert := Insert(r.tableName).Columns("id", "value").Values(id, value) - _, err = r.executeSQL(insert) + _, err = r.executeSQL(ctx, insert) return err } -func (r propertyRepository) Get(id string) (string, error) { +func (r propertyRepository) Get(ctx context.Context, id string) (string, error) { sel := Select("value").From(r.tableName).Where(Eq{"id": id}) resp := struct { Value string }{} - err := r.queryOne(sel, &resp) + err := r.queryOne(ctx, sel, &resp) if err != nil { return "", err } return resp.Value, nil } -func (r propertyRepository) DefaultGet(id string, defaultValue string) (string, error) { - value, err := r.Get(id) +func (r propertyRepository) DefaultGet(ctx context.Context, id string, defaultValue string) (string, error) { + value, err := r.Get(ctx, id) if errors.Is(err, model.ErrNotFound) { return defaultValue, nil } @@ -58,6 +57,6 @@ func (r propertyRepository) DefaultGet(id string, defaultValue string) (string, return value, nil } -func (r propertyRepository) Delete(id string) error { - return r.delete(Eq{"id": id}) +func (r propertyRepository) Delete(ctx context.Context, id string) error { + return r.delete(ctx, Eq{"id": id}) } diff --git a/persistence/property_repository_test.go b/persistence/property_repository_test.go index 3a0495e9f..880b315ec 100644 --- a/persistence/property_repository_test.go +++ b/persistence/property_repository_test.go @@ -10,25 +10,27 @@ import ( ) var _ = Describe("Property Repository", func() { + var ctx context.Context var pr model.PropertyRepository BeforeEach(func() { - pr = NewPropertyRepository(log.NewContext(context.TODO()), GetDBXBuilder()) + ctx = log.NewContext(GinkgoT().Context()) + pr = NewPropertyRepository(GetDBXBuilder()) }) It("saves and restore a new property", func() { id := "1" value := "a_value" - Expect(pr.Put(id, value)).To(BeNil()) - Expect(pr.Get(id)).To(Equal("a_value")) + Expect(pr.Put(ctx, id, value)).To(BeNil()) + Expect(pr.Get(ctx, id)).To(Equal("a_value")) }) It("updates a property", func() { - Expect(pr.Put("1", "another_value")).To(BeNil()) - Expect(pr.Get("1")).To(Equal("another_value")) + Expect(pr.Put(ctx, "1", "another_value")).To(BeNil()) + Expect(pr.Get(ctx, "1")).To(Equal("another_value")) }) It("returns a default value if property does not exist", func() { - Expect(pr.DefaultGet("2", "default")).To(Equal("default")) + Expect(pr.DefaultGet(ctx, "2", "default")).To(Equal("default")) }) }) diff --git a/persistence/radio_repository.go b/persistence/radio_repository.go index 4584ceaed..b5d4f3a07 100644 --- a/persistence/radio_repository.go +++ b/persistence/radio_repository.go @@ -16,9 +16,8 @@ type radioRepository struct { sqlRepository } -func NewRadioRepository(ctx context.Context, db dbx.Builder) model.RadioRepository { +func NewRadioRepository(db dbx.Builder) model.RadioRepository { r := &radioRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Radio{}, map[string]filterFunc{ "name": containsFilter("name"), @@ -26,60 +25,65 @@ func NewRadioRepository(ctx context.Context, db dbx.Builder) model.RadioReposito return r } -func (r *radioRepository) isPermitted() bool { - user := loggedUser(r.ctx) +func (r *radioRepository) isPermitted(ctx context.Context) bool { + user := loggedUser(ctx) return user.IsAdmin } -func (r *radioRepository) CountAll(options ...model.QueryOptions) (int64, error) { - sql := r.newSelect() - return r.count(sql, options...) +func (r *radioRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + sql := r.newSelect(ctx) + return r.count(ctx, sql, options...) } // Exists needs no library or ownership filter: radios are visible to every user. -func (r *radioRepository) Exists(id string) (bool, error) { - return r.exists(Eq{"id": id}) +func (r *radioRepository) Exists(ctx context.Context, id string) (bool, error) { + return r.exists(ctx, Eq{"id": id}) } -func (r *radioRepository) Delete(id string) error { - if !r.isPermitted() { +func (r *radioRepository) Delete(ctx context.Context, ids ...string) error { + if !r.isPermitted(ctx) { return rest.ErrPermissionDenied } - return r.deleteByID(id) + for _, id := range ids { + if err := r.deleteByID(ctx, id); err != nil { + return err + } + } + return nil } -func (r *radioRepository) Get(id string) (*model.Radio, error) { - sel := r.newSelect().Where(Eq{"id": id}).Columns("*") +func (r *radioRepository) Get(ctx context.Context, id string) (*model.Radio, error) { + sel := r.newSelect(ctx).Where(Eq{"id": id}).Columns("*") res := model.Radio{} - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) if err != nil { return &res, err } list := model.Radios{res} - r.hydrateArtwork(list) + r.hydrateArtwork(ctx, list) return &list[0], nil } -func (r *radioRepository) GetAll(options ...model.QueryOptions) (model.Radios, error) { - sel := r.newSelect(options...).Columns("*") +func (r *radioRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Radios, error) { + sel := r.newSelect(ctx, options...).Columns("*") res := model.Radios{} - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return res, err } - r.hydrateArtwork(res) + r.hydrateArtwork(ctx, res) return res, nil } // hydrateArtwork fills each radio's ImageHash/ImageAbsent from one batched item_artwork lookup. -func (r *radioRepository) hydrateArtwork(radios model.Radios) { - hydrateItems(r.ctx, r.db, model.KindRadioArtwork, radios, +func (r *radioRepository) hydrateArtwork(ctx context.Context, radios model.Radios) { + hydrateItems(ctx, r.db, model.KindRadioArtwork, radios, func(rd *model.Radio) (string, *model.ItemImage) { return rd.ID, &rd.ItemImage }) } -func (r *radioRepository) Put(radio *model.Radio, colsToUpdate ...string) error { - if !r.isPermitted() { +func (r *radioRepository) Put(ctx context.Context, radio *model.Radio, colsToUpdate ...string) error { + if !r.isPermitted(ctx) { return rest.ErrPermissionDenied } @@ -91,7 +95,7 @@ func (r *radioRepository) Put(radio *model.Radio, colsToUpdate ...string) error if len(colsToUpdate) > 0 { colsToUpdate = append(colsToUpdate, "UpdatedAt") } - _, err := r.put(radio.ID, radio, colsToUpdate...) + _, err := r.put(ctx, radio.ID, radio, colsToUpdate...) if err != nil { return err } @@ -99,50 +103,41 @@ func (r *radioRepository) Put(radio *model.Radio, colsToUpdate ...string) error // radio's cover resolves proactively. Never fails the save. item := model.ArtworkQueueItem{ItemKind: model.KindRadioArtwork.Prefix(), ItemID: radio.ID, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBump} - if err := NewArtworkQueueRepository(r.ctx, r.db).Enqueue(item); err != nil { - log.Warn(r.ctx, "could not enqueue radio artwork", "id", radio.ID, err) + if err := NewArtworkQueueRepository(r.db).Enqueue(ctx, item); err != nil { + log.Warn(ctx, "could not enqueue radio artwork", "id", radio.ID, err) } return nil } -func (r *radioRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *radioRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *radioRepository) EntityName() string { - return "radio" +func (r *radioRepository) Read(ctx context.Context, id string) (*model.Radio, error) { + return r.Get(ctx, id) } -func (r *radioRepository) NewInstance() any { - return &model.Radio{} +func (r *radioRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Radio, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *radioRepository) Read(id string) (any, error) { - return r.Get(id) -} - -func (r *radioRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) -} - -func (r *radioRepository) Save(entity any) (string, error) { - t := entity.(*model.Radio) - if !r.isPermitted() { +func (r *radioRepository) Save(ctx context.Context, t *model.Radio) (string, error) { + if !r.isPermitted(ctx) { return "", rest.ErrPermissionDenied } - err := r.Put(t) + err := r.Put(ctx, t) return t.ID, err } -func (r *radioRepository) Update(id string, entity any, cols ...string) error { - t := entity.(*model.Radio) +func (r *radioRepository) Update(ctx context.Context, id string, entity model.Radio, cols ...string) error { + t := &entity t.ID = id - if !r.isPermitted() { + if !r.isPermitted(ctx) { return rest.ErrPermissionDenied } - return r.Put(t, cols...) + return r.Put(ctx, t, cols...) } var _ model.RadioRepository = (*radioRepository)(nil) -var _ rest.Repository = (*radioRepository)(nil) -var _ rest.Persistable = (*radioRepository)(nil) +var _ rest.Repository[model.Radio] = (*radioRepository)(nil) +var _ rest.Persistable[model.Radio] = (*radioRepository)(nil) diff --git a/persistence/radio_repository_test.go b/persistence/radio_repository_test.go index 8d4b685d9..aba776853 100644 --- a/persistence/radio_repository_test.go +++ b/persistence/radio_repository_test.go @@ -13,24 +13,28 @@ import ( var _ = Describe("RadioRepository", func() { var repo model.RadioRepository + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) Describe("Admin User", func() { BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewRadioRepository(ctx, GetDBXBuilder()) - _ = repo.Put(&radioWithHomePage) + ctx = request.WithUser(log.NewContext(ctx), model.User{ID: "userid", UserName: "userid", IsAdmin: true}) + repo = NewRadioRepository(GetDBXBuilder()) + _ = repo.Put(ctx, &radioWithHomePage) }) AfterEach(func() { - all, _ := repo.GetAll() + all, _ := repo.GetAll(ctx) for _, radio := range all { - _ = repo.Delete(radio.ID) + _ = repo.Delete(ctx, radio.ID) } for i := range testRadios { - err := repo.Put(new(testRadios[i])) + err := repo.Put(ctx, new(testRadios[i])) if err != nil { panic(err) } @@ -39,35 +43,35 @@ var _ = Describe("RadioRepository", func() { Describe("Count", func() { It("returns the number of radios in the DB", func() { - Expect(repo.CountAll()).To(Equal(int64(2))) + Expect(repo.CountAll(ctx)).To(Equal(int64(2))) }) }) Describe("Delete", func() { It("deletes existing item", func() { - err := repo.Delete(radioWithHomePage.ID) + err := repo.Delete(ctx, radioWithHomePage.ID) Expect(err).To(BeNil()) - _, err = repo.Get(radioWithHomePage.ID) + _, err = repo.Get(ctx, radioWithHomePage.ID) Expect(err).To(MatchError(model.ErrNotFound)) }) It("errors when missing", func() { - Expect(repo.Delete("notanid")).To(MatchError(model.ErrNotFound)) + Expect(repo.Delete(ctx, "notanid")).To(MatchError(model.ErrNotFound)) }) }) Describe("Get", func() { It("returns an existing item", func() { - res, err := repo.Get(radioWithHomePage.ID) + res, err := repo.Get(ctx, radioWithHomePage.ID) Expect(err).To(BeNil()) Expect(res.ID).To(Equal(radioWithHomePage.ID)) }) It("errors when missing", func() { - _, err := repo.Get("notanid") + _, err := repo.Get(ctx, "notanid") Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -75,7 +79,7 @@ var _ = Describe("RadioRepository", func() { Describe("GetAll", func() { It("returns all items from the DB", func() { - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) Expect(all[0].ID).To(Equal(radioWithoutHomePage.ID)) Expect(all[1].ID).To(Equal(radioWithHomePage.ID)) @@ -84,7 +88,7 @@ var _ = Describe("RadioRepository", func() { Describe("Put", func() { It("successfully updates item", func() { - err := repo.Put(&model.Radio{ + err := repo.Put(ctx, &model.Radio{ ID: radioWithHomePage.ID, Name: "New Name", StreamUrl: "https://example.com:4533/app", @@ -92,39 +96,39 @@ var _ = Describe("RadioRepository", func() { Expect(err).To(BeNil()) - item, err := repo.Get(radioWithHomePage.ID) + item, err := repo.Get(ctx, radioWithHomePage.ID) Expect(err).To(BeNil()) Expect(item.HomePageUrl).To(Equal("")) }) It("successfully creates item", func() { - err := repo.Put(&model.Radio{ + err := repo.Put(ctx, &model.Radio{ Name: "New radio", StreamUrl: "https://example.com:4533/app", }) Expect(err).To(BeNil()) - Expect(repo.CountAll()).To(Equal(int64(3))) + Expect(repo.CountAll(ctx)).To(Equal(int64(3))) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) Expect(all[2].StreamUrl).To(Equal("https://example.com:4533/app")) }) It("enqueues artwork resolution for the saved radio", func() { - err := repo.Put(&model.Radio{ + err := repo.Put(ctx, &model.Radio{ Name: "Artwork radio", StreamUrl: "https://example.com:4533/artwork", }) Expect(err).To(BeNil()) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) created := all[len(all)-1] - queueRepo := NewArtworkQueueRepository(context.Background(), GetDBXBuilder()) - queued, err := queueRepo.DequeueBatch(1000) + queueRepo := NewArtworkQueueRepository(GetDBXBuilder()) + queued, err := queueRepo.DequeueBatch(ctx, 1000) Expect(err).To(BeNil()) Expect(queued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "ra"), @@ -138,12 +142,11 @@ var _ = Describe("RadioRepository", func() { It("only writes the columns sent by the client", func() { radio := radioWithHomePage radio.UploadedImage = "cover.png" - Expect(repo.Put(&radio)).To(Succeed()) + Expect(repo.Put(ctx, &radio)).To(Succeed()) - persistable := repo.(rest.Persistable) - Expect(persistable.Update(radio.ID, &model.Radio{Name: "Renamed"}, "name")).To(Succeed()) + Expect(repo.Update(ctx, radio.ID, model.Radio{Name: "Renamed"}, "name")).To(Succeed()) - item, err := repo.Get(radio.ID) + item, err := repo.Get(ctx, radio.ID) Expect(err).To(BeNil()) Expect(item.Name).To(Equal("Renamed")) Expect(item.UploadedImage).To(Equal("cover.png")) @@ -155,20 +158,19 @@ var _ = Describe("RadioRepository", func() { Describe("Regular User", func() { BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: false}) - repo = NewRadioRepository(ctx, GetDBXBuilder()) + ctx = request.WithUser(log.NewContext(ctx), model.User{ID: "userid", UserName: "userid", IsAdmin: false}) + repo = NewRadioRepository(GetDBXBuilder()) }) Describe("Count", func() { It("returns the number of radios in the DB", func() { - Expect(repo.CountAll()).To(Equal(int64(2))) + Expect(repo.CountAll(ctx)).To(Equal(int64(2))) }) }) Describe("Delete", func() { It("fails to delete items", func() { - err := repo.Delete(radioWithHomePage.ID) + err := repo.Delete(ctx, radioWithHomePage.ID) Expect(err).To(Equal(rest.ErrPermissionDenied)) }) @@ -176,14 +178,14 @@ var _ = Describe("RadioRepository", func() { Describe("Get", func() { It("returns an existing item", func() { - res, err := repo.Get(radioWithHomePage.ID) + res, err := repo.Get(ctx, radioWithHomePage.ID) Expect(err).To(BeNil()) Expect(res.ID).To(Equal(radioWithHomePage.ID)) }) It("errors when missing", func() { - _, err := repo.Get("notanid") + _, err := repo.Get(ctx, "notanid") Expect(err).To(MatchError(model.ErrNotFound)) }) @@ -191,7 +193,7 @@ var _ = Describe("RadioRepository", func() { Describe("GetAll", func() { It("returns all items from the DB", func() { - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).To(BeNil()) Expect(all[0].ID).To(Equal(radioWithoutHomePage.ID)) Expect(all[1].ID).To(Equal(radioWithHomePage.ID)) @@ -200,7 +202,7 @@ var _ = Describe("RadioRepository", func() { Describe("Put", func() { It("fails to update item", func() { - err := repo.Put(&model.Radio{ + err := repo.Put(ctx, &model.Radio{ ID: radioWithHomePage.ID, Name: "New Name", StreamUrl: "https://example.com:4533/app", diff --git a/persistence/scrobble_buffer_repository.go b/persistence/scrobble_buffer_repository.go index cf54c664a..0daf2b0e7 100644 --- a/persistence/scrobble_buffer_repository.go +++ b/persistence/scrobble_buffer_repository.go @@ -29,15 +29,14 @@ func (t *dbScrobbleBuffer) PostScan() error { return nil } -func NewScrobbleBufferRepository(ctx context.Context, db dbx.Builder) model.ScrobbleBufferRepository { +func NewScrobbleBufferRepository(db dbx.Builder) model.ScrobbleBufferRepository { r := &scrobbleBufferRepository{} - r.ctx = ctx r.db = db r.tableName = "scrobble_buffer" return r } -func (r *scrobbleBufferRepository) UserIDs(service string) ([]string, error) { +func (r *scrobbleBufferRepository) UserIDs(ctx context.Context, service string) ([]string, error) { sql := Select().Columns("user_id"). From(r.tableName). Where(And{ @@ -46,11 +45,11 @@ func (r *scrobbleBufferRepository) UserIDs(service string) ([]string, error) { GroupBy("user_id"). OrderBy("count(*)") var userIds []string - err := r.queryAllSlice(sql, &userIds) + err := r.queryAllSlice(ctx, sql, &userIds) return userIds, err } -func (r *scrobbleBufferRepository) Enqueue(service, userId, mediaFileId string, playTime time.Time) error { +func (r *scrobbleBufferRepository) Enqueue(ctx context.Context, service, userId, mediaFileId string, playTime time.Time) error { ins := Insert(r.tableName).SetMap(map[string]any{ "id": id.NewRandom(), "user_id": userId, @@ -59,11 +58,11 @@ func (r *scrobbleBufferRepository) Enqueue(service, userId, mediaFileId string, "play_time": playTime, "enqueue_time": time.Now(), }) - _, err := r.executeSQL(ins) + _, err := r.executeSQL(ctx, ins) return err } -func (r *scrobbleBufferRepository) Next(service string, userId string) (*model.ScrobbleEntry, error) { +func (r *scrobbleBufferRepository) Next(ctx context.Context, service string, userId string) (*model.ScrobbleEntry, error) { // Put `s.*` last or else m.id overrides s.id sql := Select().Columns("m.*, s.*"). From(r.tableName+" s"). @@ -75,30 +74,30 @@ func (r *scrobbleBufferRepository) Next(service string, userId string) (*model.S OrderBy("play_time", "s.rowid").Limit(1) var res dbScrobbleBuffer - err := r.queryOne(sql, &res) + err := r.queryOne(ctx, sql, &res) if errors.Is(err, model.ErrNotFound) { return nil, nil } if err != nil { return nil, err } - res.ScrobbleEntry.Participants, err = r.getParticipants(&res.ScrobbleEntry.MediaFile) + res.ScrobbleEntry.Participants, err = r.getParticipants(ctx, &res.ScrobbleEntry.MediaFile) if err != nil { return nil, err } return res.ScrobbleEntry, nil } -func (r *scrobbleBufferRepository) Dequeue(entry *model.ScrobbleEntry) error { - return r.delete(Eq{"id": entry.ID}) +func (r *scrobbleBufferRepository) Dequeue(ctx context.Context, entry *model.ScrobbleEntry) error { + return r.delete(ctx, Eq{"id": entry.ID}) } -func (r *scrobbleBufferRepository) Discard(service string) error { - return r.delete(Eq{"service": service}) +func (r *scrobbleBufferRepository) Discard(ctx context.Context, service string) error { + return r.delete(ctx, Eq{"service": service}) } -func (r *scrobbleBufferRepository) Length() (int64, error) { - return r.count(Select()) +func (r *scrobbleBufferRepository) Length(ctx context.Context) (int64, error) { + return r.count(ctx, Select()) } var _ model.ScrobbleBufferRepository = (*scrobbleBufferRepository)(nil) diff --git a/persistence/scrobble_buffer_repository_test.go b/persistence/scrobble_buffer_repository_test.go index 3aa71070e..38b0e4e63 100644 --- a/persistence/scrobble_buffer_repository_test.go +++ b/persistence/scrobble_buffer_repository_test.go @@ -16,6 +16,7 @@ import ( var _ = Describe("ScrobbleBufferRepository", func() { var scrobble model.ScrobbleBufferRepository var rawRepo sqlRepository + var ctx context.Context enqueueTime := time.Date(2025, 01, 01, 00, 00, 00, 00, time.Local) var ids []string @@ -32,17 +33,16 @@ var _ = Describe("ScrobbleBufferRepository", func() { "play_time": playTime, "enqueue_time": enqueueTime, }) - _, err := rawRepo.executeSQL(ins) + _, err := rawRepo.executeSQL(ctx, ins) Expect(err).ToNot(HaveOccurred()) } BeforeEach(func() { - ctx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid", UserName: "johndoe", IsAdmin: true}) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid", UserName: "johndoe", IsAdmin: true}) db := GetDBXBuilder() - scrobble = NewScrobbleBufferRepository(ctx, db) + scrobble = NewScrobbleBufferRepository(db) rawRepo = sqlRepository{ - ctx: ctx, tableName: "scrobble_buffer", db: db, } @@ -51,14 +51,14 @@ var _ = Describe("ScrobbleBufferRepository", func() { AfterEach(func() { del := squirrel.Delete(rawRepo.tableName) - _, err := rawRepo.executeSQL(del) + _, err := rawRepo.executeSQL(ctx, del) Expect(err).ToNot(HaveOccurred()) }) Describe("Without data", func() { Describe("Count", func() { It("returns zero when empty", func() { - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeZero()) }) @@ -66,10 +66,10 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("Dequeue", func() { It("is a no-op when deleting a nonexistent item", func() { - err := scrobble.Dequeue(&model.ScrobbleEntry{ID: "fake"}) + err := scrobble.Dequeue(ctx, &model.ScrobbleEntry{ID: "fake"}) Expect(err).ToNot(HaveOccurred()) - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(0))) }) @@ -77,7 +77,7 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("Next", func() { It("should not fail with no item for the service", func() { - entry, err := scrobble.Next("fake", "userid") + entry, err := scrobble.Next(ctx, "fake", "userid") Expect(entry).To(BeNil()) Expect(err).ToNot(HaveOccurred()) }) @@ -85,7 +85,7 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("UserIds", func() { It("should return empty list with no data", func() { - ids, err := scrobble.UserIDs("service") + ids, err := scrobble.UserIDs(ctx, "service") Expect(err).ToNot(HaveOccurred()) Expect(ids).To(BeEmpty()) }) @@ -107,7 +107,7 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("Count", func() { It("Returns count when populated", func() { - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(4))) }) @@ -115,23 +115,23 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("Dequeue", func() { It("is a no-op when deleting a nonexistent item", func() { - err := scrobble.Dequeue(&model.ScrobbleEntry{ID: "fake"}) + err := scrobble.Dequeue(ctx, &model.ScrobbleEntry{ID: "fake"}) Expect(err).ToNot(HaveOccurred()) - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(4))) }) It("deletes an item when specified properly", func() { - err := scrobble.Dequeue(&model.ScrobbleEntry{ID: ids[3]}) + err := scrobble.Dequeue(ctx, &model.ScrobbleEntry{ID: ids[3]}) Expect(err).ToNot(HaveOccurred()) - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(3))) - entry, err := scrobble.Next("b", "2222") + entry, err := scrobble.Next(ctx, "b", "2222") Expect(err).ToNot(HaveOccurred()) Expect(entry).To(BeNil()) }) @@ -141,14 +141,14 @@ var _ = Describe("ScrobbleBufferRepository", func() { DescribeTable("enqueues an item properly", func(service, userId, fileId string, playTime time.Time) { now := time.Now() - err := scrobble.Enqueue(service, userId, fileId, playTime) + err := scrobble.Enqueue(ctx, service, userId, fileId, playTime) Expect(err).ToNot(HaveOccurred()) - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(5))) - entry, err := scrobble.Next(service, userId) + entry, err := scrobble.Next(ctx, service, userId) Expect(err).ToNot(HaveOccurred()) Expect(entry).ToNot(BeNil()) @@ -165,7 +165,7 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("Next", func() { DescribeTable("Returns the next item when populated", func(service, id string, playTime time.Time, fileId, artistId string) { - entry, err := scrobble.Next(service, id) + entry, err := scrobble.Next(ctx, service, id) Expect(err).ToNot(HaveOccurred()) Expect(entry).ToNot(BeNil()) @@ -193,21 +193,21 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("Discard", func() { It("deletes all entries for a service, keeping other services intact", func() { - Expect(scrobble.Discard("a")).To(Succeed()) + Expect(scrobble.Discard(ctx, "a")).To(Succeed()) - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(1))) - entry, err := scrobble.Next("b", "2222") + entry, err := scrobble.Next(ctx, "b", "2222") Expect(err).ToNot(HaveOccurred()) Expect(entry).ToNot(BeNil()) }) It("is a no-op for a service without entries", func() { - Expect(scrobble.Discard("nonexistent")).To(Succeed()) + Expect(scrobble.Discard(ctx, "nonexistent")).To(Succeed()) - count, err := scrobble.Length() + count, err := scrobble.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(4))) }) @@ -215,13 +215,13 @@ var _ = Describe("ScrobbleBufferRepository", func() { Describe("UserIds", func() { It("should return ordered list for services", func() { - ids, err := scrobble.UserIDs("a") + ids, err := scrobble.UserIDs(ctx, "a") Expect(err).ToNot(HaveOccurred()) Expect(ids).To(Equal([]string{"2222", "userid"})) }) It("should return for a different service", func() { - ids, err := scrobble.UserIDs("b") + ids, err := scrobble.UserIDs(ctx, "b") Expect(err).ToNot(HaveOccurred()) Expect(ids).To(Equal([]string{"2222"})) }) diff --git a/persistence/scrobble_repository.go b/persistence/scrobble_repository.go index 7cc60ae23..3d9882b99 100644 --- a/persistence/scrobble_repository.go +++ b/persistence/scrobble_repository.go @@ -22,17 +22,16 @@ func toTs(_ string, value any) Sqlizer { return LtOrEq{"scrobbles.submission_time": value} } -func (r *scrobbleRepository) baseQuery(options ...model.QueryOptions) SelectBuilder { - user := loggedUser(r.ctx) +func (r *scrobbleRepository) baseQuery(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + user := loggedUser(ctx) - return r.newSelect(options...). + return r.newSelect(ctx, options...). Columns("id", "media_file_id", "submission_time"). Where(Eq{"scrobbles.user_id": user.ID}) } -func NewScrobbleRepository(ctx context.Context, db dbx.Builder) model.ScrobbleRepository { +func NewScrobbleRepository(db dbx.Builder) model.ScrobbleRepository { r := &scrobbleRepository{} - r.ctx = ctx r.db = db r.tableName = "scrobbles" r.registerModel(&model.Scrobble{}, map[string]filterFunc{ @@ -45,55 +44,47 @@ func NewScrobbleRepository(ctx context.Context, db dbx.Builder) model.ScrobbleRe return r } -func (r *scrobbleRepository) RecordScrobble(mediaFileID string, submissionTime time.Time) error { - userID := loggedUser(r.ctx).ID +func (r *scrobbleRepository) RecordScrobble(ctx context.Context, mediaFileID string, submissionTime time.Time) error { + userID := loggedUser(ctx).ID values := map[string]any{ "media_file_id": mediaFileID, "user_id": userID, "submission_time": submissionTime.Unix(), } insert := Insert(r.tableName).SetMap(values) - _, err := r.executeSQL(insert) + _, err := r.executeSQL(ctx, insert) return err } -func (r *scrobbleRepository) CountAll(options ...model.QueryOptions) (int64, error) { - return r.count(r.baseQuery(), options...) +func (r *scrobbleRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + return r.count(ctx, r.baseQuery(ctx), options...) } -func (r *scrobbleRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *scrobbleRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *scrobbleRepository) Get(id string) (*model.Scrobble, error) { - sel := r.baseQuery().Where(Eq{"id": id}) +func (r *scrobbleRepository) Get(ctx context.Context, id string) (*model.Scrobble, error) { + sel := r.baseQuery(ctx).Where(Eq{"id": id}) var res model.Scrobble - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *scrobbleRepository) GetAll(options ...model.QueryOptions) (model.Scrobbles, error) { - sel := r.baseQuery(options...) +func (r *scrobbleRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Scrobbles, error) { + sel := r.baseQuery(ctx, options...) var scrobbles model.Scrobbles - err := r.queryAll(sel, &scrobbles) + err := r.queryAll(ctx, sel, &scrobbles) return scrobbles, err } -func (r *scrobbleRepository) Read(id string) (any, error) { - return r.Get(id) +func (r *scrobbleRepository) Read(ctx context.Context, id string) (*model.Scrobble, error) { + return r.Get(ctx, id) } -func (r *scrobbleRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - return r.GetAll(r.parseRestOptions(r.ctx, options...)) -} - -func (r *scrobbleRepository) EntityName() string { - return "scrobble" -} - -func (r *scrobbleRepository) NewInstance() any { - return &model.Scrobble{} +func (r *scrobbleRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Scrobble, error) { + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } var _ model.ScrobbleRepository = (*scrobbleRepository)(nil) -var _ model.ResourceRepository = (*scrobbleRepository)(nil) +var _ rest.Repository[model.Scrobble] = (*scrobbleRepository)(nil) diff --git a/persistence/scrobble_repository_test.go b/persistence/scrobble_repository_test.go index e9103b127..860b1ad02 100644 --- a/persistence/scrobble_repository_test.go +++ b/persistence/scrobble_repository_test.go @@ -28,10 +28,9 @@ var _ = Describe("ScrobbleRepository", func() { userID = id.NewRandom() ctx = request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: userID, UserName: "johndoe", IsAdmin: true}) db := GetDBXBuilder() - repo = NewScrobbleRepository(ctx, db) + repo = NewScrobbleRepository(db) rawRepo = sqlRepository{ - ctx: ctx, tableName: "scrobbles", db: db, } @@ -65,7 +64,7 @@ var _ = Describe("ScrobbleRepository", func() { }).Execute() Expect(err).ToNot(HaveOccurred()) - err = repo.RecordScrobble(fileID, submissionTime) + err = repo.RecordScrobble(ctx, fileID, submissionTime) Expect(err).ToNot(HaveOccurred()) // Verify insertion @@ -87,22 +86,22 @@ var _ = Describe("ScrobbleRepository", func() { Context("admin user (id userid)", func() { BeforeEach(func() { ctx = request.WithUser(log.NewContext(context.TODO()), adminUser) - repo = NewScrobbleRepository(ctx, GetDBXBuilder()) + repo = NewScrobbleRepository(GetDBXBuilder()) }) Describe("Count", func() { It("Returns the number of scrobbles in the DB for admin user", func() { - Expect(repo.CountAll()).To(Equal(int64(2))) + Expect(repo.CountAll(ctx)).To(Equal(int64(2))) }) It("returns scrobbles in a range", func() { - Expect(repo.CountAll(model.QueryOptions{Filters: squirrel.LtOrEq{"submission_time": 1}})).To(Equal(int64(1))) + Expect(repo.CountAll(ctx, model.QueryOptions{Filters: squirrel.LtOrEq{"submission_time": 1}})).To(Equal(int64(1))) }) }) Describe("Get", func() { It("returns an existing scrobble for the user", func() { - scrobble, err := repo.Get("1") + scrobble, err := repo.Get(ctx, "1") Expect(err).To(BeNil()) Expect(scrobble.ID).To(Equal(int64(1))) Expect(scrobble.MediaFileID).To(Equal("1001")) @@ -111,19 +110,19 @@ var _ = Describe("ScrobbleRepository", func() { }) It("does not return a scrobble that exists for another user", func() { - _, err := repo.Get("2") + _, err := repo.Get(ctx, "2") Expect(err).To(MatchError(model.ErrNotFound)) }) It("does not return a scrobble that does not exist", func() { - _, err := repo.Get("444") + _, err := repo.Get(ctx, "444") Expect(err).To(MatchError(model.ErrNotFound)) }) }) Describe("GetAll", func() { It("returns all scrobbles in reverse order", func() { - scrobbles, err := repo.GetAll(model.QueryOptions{ + scrobbles, err := repo.GetAll(ctx, model.QueryOptions{ Sort: "submission_time", Order: "DESC", }) @@ -140,7 +139,7 @@ var _ = Describe("ScrobbleRepository", func() { }) It("returns scrobbles in a range", func() { - scrobbles, err := repo.GetAll(model.QueryOptions{ + scrobbles, err := repo.GetAll(ctx, model.QueryOptions{ Filters: squirrel.GtOrEq{"submission_time": 1}}) Expect(err).To(BeNil()) @@ -156,22 +155,22 @@ var _ = Describe("ScrobbleRepository", func() { Context("non-admin user", func() { BeforeEach(func() { ctx = request.WithUser(log.NewContext(context.TODO()), regularUser) - repo = NewScrobbleRepository(ctx, GetDBXBuilder()) + repo = NewScrobbleRepository(GetDBXBuilder()) }) Describe("Count", func() { It("Returns the number of scrobbles in the DB for admin user", func() { - Expect(repo.CountAll()).To(Equal(int64(1))) + Expect(repo.CountAll(ctx)).To(Equal(int64(1))) }) It("returns scrobbles in a range", func() { - Expect(repo.CountAll(model.QueryOptions{Filters: squirrel.LtOrEq{"submission_time": 1}})).To(Equal(int64(0))) + Expect(repo.CountAll(ctx, model.QueryOptions{Filters: squirrel.LtOrEq{"submission_time": 1}})).To(Equal(int64(0))) }) }) Describe("Get", func() { It("returns an existing scrobble for the user", func() { - scrobble, err := repo.Get("2") + scrobble, err := repo.Get(ctx, "2") Expect(err).To(BeNil()) Expect(scrobble.ID).To(Equal(int64(2))) Expect(scrobble.MediaFileID).To(Equal("1003")) @@ -179,19 +178,19 @@ var _ = Describe("ScrobbleRepository", func() { }) It("does not return a scrobble that exists for another user", func() { - _, err := repo.Get("1") + _, err := repo.Get(ctx, "1") Expect(err).To(MatchError(model.ErrNotFound)) }) It("does not return a scrobble that does not exist", func() { - _, err := repo.Get("444") + _, err := repo.Get(ctx, "444") Expect(err).To(MatchError(model.ErrNotFound)) }) }) Describe("GetAll", func() { It("returns all scrobbles in reverse order", func() { - scrobbles, err := repo.GetAll(model.QueryOptions{ + scrobbles, err := repo.GetAll(ctx, model.QueryOptions{ Sort: "submission_time", Order: "DESC", }) @@ -204,7 +203,7 @@ var _ = Describe("ScrobbleRepository", func() { }) It("returns scrobbles in a range", func() { - scrobbles, err := repo.GetAll(model.QueryOptions{ + scrobbles, err := repo.GetAll(ctx, model.QueryOptions{ Filters: squirrel.GtOrEq{"submission_time": 1}}) Expect(err).To(BeNil()) diff --git a/persistence/share_repository.go b/persistence/share_repository.go index 0cf23068a..6dd9c3d85 100644 --- a/persistence/share_repository.go +++ b/persistence/share_repository.go @@ -18,9 +18,8 @@ type shareRepository struct { sqlRepository } -func NewShareRepository(ctx context.Context, db dbx.Builder) model.ShareRepository { +func NewShareRepository(db dbx.Builder) model.ShareRepository { r := &shareRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Share{}, nil) r.setSortMappings(map[string]string{ @@ -29,40 +28,40 @@ func NewShareRepository(ctx context.Context, db dbx.Builder) model.ShareReposito return r } -func (r *shareRepository) Delete(id string) error { - return r.deleteOwned(id) +func (r *shareRepository) Delete(ctx context.Context, ids ...string) error { + return r.deleteOwnedAll(ctx, ids...) } -func (r *shareRepository) selectShare(options ...model.QueryOptions) SelectBuilder { - return r.newSelect(options...).Join("user u on u.id = share.user_id"). +func (r *shareRepository) selectShare(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + return r.newSelect(ctx, options...).Join("user u on u.id = share.user_id"). Columns("share.*", "user_name as username"). - Where(r.addRestriction()) + Where(r.addRestriction(ctx)) } -func (r *shareRepository) Exists(id string) (bool, error) { - return r.exists(r.addRestriction(And{Eq{"id": id}})) +func (r *shareRepository) Exists(ctx context.Context, id string) (bool, error) { + return r.exists(ctx, r.addRestriction(ctx, And{Eq{"id": id}})) } -func (r *shareRepository) Get(id string) (*model.Share, error) { - sel := r.selectShare().Where(Eq{"share.id": id}) +func (r *shareRepository) Get(ctx context.Context, id string) (*model.Share, error) { + sel := r.selectShare(ctx).Where(Eq{"share.id": id}) var res model.Share - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) if err != nil { return nil, err } - err = r.loadMedia(&res) + err = r.loadMedia(ctx, &res) return &res, err } -func (r *shareRepository) GetAll(options ...model.QueryOptions) (model.Shares, error) { - sq := r.selectShare(options...) +func (r *shareRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Shares, error) { + sq := r.selectShare(ctx, options...) res := model.Shares{} - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) if err != nil { return nil, err } for i := range res { - err = r.loadMedia(&res[i]) + err = r.loadMedia(ctx, &res[i]) if err != nil { return nil, fmt.Errorf("error loading media for share %s: %w", res[i].ID, err) } @@ -70,7 +69,7 @@ func (r *shareRepository) GetAll(options ...model.QueryOptions) (model.Shares, e return res, err } -func (r *shareRepository) loadMedia(share *model.Share) error { +func (r *shareRepository) loadMedia(ctx context.Context, share *model.Share) error { ids := strings.Split(share.ResourceIDs, ",") if len(ids) == 0 { return nil @@ -79,7 +78,7 @@ func (r *shareRepository) loadMedia(share *model.Share) error { return And{cond, Eq{"missing": false}} } // Load as the share owner so their library access is applied, whoever renders the share. - ctx, err := r.ownerContext(share) + ownerCtx, err := r.ownerContext(ctx, share) if err != nil { return err } @@ -87,59 +86,59 @@ func (r *shareRepository) loadMedia(share *model.Share) error { case "artist": // Match by album-artist participation, not the deprecated album_artist_id // column (first album artist only), so co-album-artists are included too. - albumRepo := NewAlbumRepository(ctx, r.db) - share.Albums, err = albumRepo.GetAll(model.QueryOptions{Filters: noMissing(ParticipantIDFilter("album", ids, model.RoleAlbumArtist)), Sort: "artist"}) + albumRepo := NewAlbumRepository(r.db) + share.Albums, err = albumRepo.GetAll(ownerCtx, model.QueryOptions{Filters: noMissing(ParticipantIDFilter("album", ids, model.RoleAlbumArtist)), Sort: "artist"}) if err != nil { return err } - mfRepo := NewMediaFileRepository(ctx, r.db) - share.Tracks, err = mfRepo.GetAll(model.QueryOptions{Filters: noMissing(ParticipantIDFilter("media_file", ids, model.RoleAlbumArtist)), Sort: "artist"}) + mfRepo := NewMediaFileRepository(r.db) + share.Tracks, err = mfRepo.GetAll(ownerCtx, model.QueryOptions{Filters: noMissing(ParticipantIDFilter("media_file", ids, model.RoleAlbumArtist)), Sort: "artist"}) return err case "album": - albumRepo := NewAlbumRepository(ctx, r.db) - share.Albums, err = albumRepo.GetAll(model.QueryOptions{Filters: noMissing(Eq{"album.id": ids})}) + albumRepo := NewAlbumRepository(r.db) + share.Albums, err = albumRepo.GetAll(ownerCtx, model.QueryOptions{Filters: noMissing(Eq{"album.id": ids})}) if err != nil { return err } - mfRepo := NewMediaFileRepository(ctx, r.db) - share.Tracks, err = mfRepo.GetAll(model.QueryOptions{Filters: noMissing(Eq{"album_id": ids}), Sort: "album"}) + mfRepo := NewMediaFileRepository(r.db) + share.Tracks, err = mfRepo.GetAll(ownerCtx, model.QueryOptions{Filters: noMissing(Eq{"album_id": ids}), Sort: "album"}) return err case "playlist": - plsRepo := NewPlaylistRepository(ctx, r.db) + plsRepo := NewPlaylistRepository(r.db) // Tracks returns nil when the playlist is no longer visible to the owner // (e.g. it was made private after the share was created); leave the share // with no tracks rather than exposing it. - trackRepo := plsRepo.Tracks(ids[0], true) + trackRepo := plsRepo.Tracks(ownerCtx, ids[0], true) if trackRepo == nil { return nil } - tracks, err := trackRepo.GetAll(model.QueryOptions{Sort: "id", Filters: noMissing(Eq{})}) + tracks, err := trackRepo.GetAll(ownerCtx, model.QueryOptions{Sort: "id", Filters: noMissing(Eq{})}) if err != nil { return err } share.Tracks = tracks.MediaFiles() return nil case "media_file": - mfRepo := NewMediaFileRepository(ctx, r.db) - tracks, err := mfRepo.GetAll(model.QueryOptions{Filters: noMissing(Eq{"media_file.id": ids})}) + mfRepo := NewMediaFileRepository(r.db) + tracks, err := mfRepo.GetAll(ownerCtx, model.QueryOptions{Filters: noMissing(Eq{"media_file.id": ids})}) share.Tracks = sortByIdPosition(tracks, ids) return err } - log.Warn(r.ctx, "Unsupported Share ResourceType", "share", share.ID, "resourceType", share.ResourceType) + log.Warn(ctx, "Unsupported Share ResourceType", "share", share.ID, "resourceType", share.ResourceType) return nil } // ownerContext returns a context scoped to the share owner, so repository // queries apply the owner's library access when a public share is rendered. -func (r *shareRepository) ownerContext(share *model.Share) (context.Context, error) { - owner, err := NewUserRepository(r.ctx, r.db).Get(share.UserID) +func (r *shareRepository) ownerContext(ctx context.Context, share *model.Share) (context.Context, error) { + owner, err := NewUserRepository(r.db).Get(ctx, share.UserID) if err != nil { return nil, fmt.Errorf("loading share owner %q: %w", share.UserID, err) } if owner == nil { return nil, fmt.Errorf("share owner %q not found", share.UserID) } - return request.WithUser(r.ctx, *owner), nil + return request.WithUser(ctx, *owner), nil } func sortByIdPosition(mfs model.MediaFiles, ids []string) model.MediaFiles { @@ -156,60 +155,51 @@ func sortByIdPosition(mfs model.MediaFiles, ids []string) model.MediaFiles { return sorted } -func (r *shareRepository) Update(id string, entity any, cols ...string) error { - s := entity.(*model.Share) +func (r *shareRepository) Update(ctx context.Context, id string, entity model.Share, cols ...string) error { + s := &entity s.ID = id s.UpdatedAt = time.Now() if len(cols) > 0 { cols = append(cols, "updated_at") } - return r.updateOwned(id, s, cols...) + return r.updateOwned(ctx, id, s, cols...) } -func (r *shareRepository) Save(entity any) (string, error) { - s := entity.(*model.Share) +func (r *shareRepository) Save(ctx context.Context, s *model.Share) (string, error) { // TODO Validate record // Owner is server-managed: for an authenticated request, never trust a // client-supplied UserID, as it drives the share's library-access context. - u := loggedUser(r.ctx) + u := loggedUser(ctx) if u.ID != invalidUserId || s.UserID == "" { s.UserID = u.ID } s.CreatedAt = time.Now() s.UpdatedAt = time.Now() - return r.put(s.ID, s) + return r.put(ctx, s.ID, s) } -func (r *shareRepository) CountAll(options ...model.QueryOptions) (int64, error) { - return r.count(r.selectShare(), options...) +func (r *shareRepository) CountAll(ctx context.Context, options ...model.QueryOptions) (int64, error) { + return r.count(ctx, r.selectShare(ctx), options...) } -func (r *shareRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.CountAll(r.parseRestOptions(r.ctx, options...)) +func (r *shareRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *shareRepository) EntityName() string { - return "share" -} - -func (r *shareRepository) NewInstance() any { - return &model.Share{} -} - -func (r *shareRepository) Read(id string) (any, error) { - sel := r.selectShare().Where(Eq{"share.id": id}) +func (r *shareRepository) Read(ctx context.Context, id string) (*model.Share, error) { + sel := r.selectShare(ctx).Where(Eq{"share.id": id}) var res model.Share - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *shareRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - sq := r.selectShare(r.parseRestOptions(r.ctx, options...)) +func (r *shareRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Share, error) { + sq := r.selectShare(ctx, r.parseRestOptions(ctx, options...)) res := model.Shares{} - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) return res, err } var _ model.ShareRepository = (*shareRepository)(nil) -var _ rest.Repository = (*shareRepository)(nil) -var _ rest.Persistable = (*shareRepository)(nil) +var _ rest.Repository[model.Share] = (*shareRepository)(nil) +var _ rest.Persistable[model.Share] = (*shareRepository)(nil) diff --git a/persistence/share_repository_test.go b/persistence/share_repository_test.go index af8cf2ae9..33b5e2110 100644 --- a/persistence/share_repository_test.go +++ b/persistence/share_repository_test.go @@ -22,11 +22,11 @@ var _ = Describe("ShareRepository", func() { BeforeEach(func() { DeferCleanup(configtest.SetupConfig()) ctx = request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - repo = NewShareRepository(ctx, GetDBXBuilder()) + repo = NewShareRepository(GetDBXBuilder()) // Insert the admin user into the database (required for foreign key constraint) - ur := NewUserRepository(ctx, GetDBXBuilder()) - err := ur.Put(&adminUser) + ur := NewUserRepository(GetDBXBuilder()) + err := ur.Put(ctx, &adminUser) Expect(err).ToNot(HaveOccurred()) // Clean up shares @@ -37,9 +37,15 @@ var _ = Describe("ShareRepository", func() { Describe("Headless Access", func() { Context("Repository creation and basic operations", func() { + var headlessCtx context.Context + + BeforeEach(func() { + headlessCtx = GinkgoT().Context() + }) + It("should create repository successfully with no user context", func() { // Create repository with no user context (headless) - headlessRepo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) + headlessRepo := NewShareRepository(GetDBXBuilder()) Expect(headlessRepo).ToNot(BeNil()) }) @@ -61,8 +67,8 @@ var _ = Describe("ShareRepository", func() { Expect(err).ToNot(HaveOccurred()) // Headless process should see all shares - headlessRepo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - shares, err := headlessRepo.GetAll() + headlessRepo := NewShareRepository(GetDBXBuilder()) + shares, err := headlessRepo.GetAll(headlessCtx) Expect(err).ToNot(HaveOccurred()) found := false @@ -93,8 +99,8 @@ var _ = Describe("ShareRepository", func() { Expect(err).ToNot(HaveOccurred()) // Headless process should be able to get the share - headlessRepo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - share, err := headlessRepo.Get(shareID) + headlessRepo := NewShareRepository(GetDBXBuilder()) + share, err := headlessRepo.Get(headlessCtx, shareID) Expect(err).ToNot(HaveOccurred()) Expect(share.ID).To(Equal(shareID)) Expect(share.Description).To(Equal("Headless Get Share")) @@ -125,7 +131,7 @@ var _ = Describe("ShareRepository", func() { // The Get operation should work without SQL ambiguity errors // even if no albums are found - share, err := repo.Get(shareID) + share, err := repo.Get(ctx, shareID) Expect(err).ToNot(HaveOccurred()) Expect(share.ID).To(Equal(shareID)) // Albums array should be empty since we used non-existent album ID @@ -142,26 +148,26 @@ var _ = Describe("ShareRepository", func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) // A second library the owner has no access to, plus a track in it - lr := NewLibraryRepository(adminCtx, GetDBXBuilder()) + lr := NewLibraryRepository(GetDBXBuilder()) otherLib = model.Library{ID: 0, Name: "Share Other Library", Path: "/share/other/lib"} - Expect(lr.Put(&otherLib)).To(Succeed()) - mr := NewMediaFileRepository(adminCtx, GetDBXBuilder()) - Expect(mr.Put(&model.MediaFile{ID: "share-other", LibraryID: otherLib.ID, Path: "s/other.mp3", Title: "ShareOther"})).To(Succeed()) - Expect(mr.Put(&model.MediaFile{ID: "share-ok", LibraryID: 1, Path: "s/ok.mp3", Title: "ShareOK"})).To(Succeed()) + Expect(lr.Put(adminCtx, &otherLib)).To(Succeed()) + mr := NewMediaFileRepository(GetDBXBuilder()) + Expect(mr.Put(adminCtx, &model.MediaFile{ID: "share-other", LibraryID: otherLib.ID, Path: "s/other.mp3", Title: "ShareOther"})).To(Succeed()) + Expect(mr.Put(adminCtx, &model.MediaFile{ID: "share-ok", LibraryID: 1, Path: "s/ok.mp3", Title: "ShareOK"})).To(Succeed()) // Non-admin owner with access to library 1 only owner = createUserWithLibraries("share-owner", []int{1}) - ur := NewUserRepository(adminCtx, GetDBXBuilder()) - Expect(ur.Put(&owner)).To(Succeed()) - Expect(ur.SetUserLibraries(owner.ID, []int{1})).To(Succeed()) + ur := NewUserRepository(GetDBXBuilder()) + Expect(ur.Put(adminCtx, &owner)).To(Succeed()) + Expect(ur.SetUserLibraries(adminCtx, owner.ID, []int{1})).To(Succeed()) // Owner-owned playlist containing tracks from both libraries plsID = "share-scope-pls" ownerCtx := request.WithUser(log.NewContext(GinkgoT().Context()), owner) - pr := NewPlaylistRepository(ownerCtx, GetDBXBuilder()) + pr := NewPlaylistRepository(GetDBXBuilder()) pls := &model.Playlist{ID: plsID, Name: "Scope Test", OwnerID: owner.ID} pls.AddMediaFiles(model.MediaFiles{{ID: "share-ok"}, {ID: "share-other"}}) - Expect(pr.Put(pls)).To(Succeed()) + Expect(pr.Put(ownerCtx, pls)).To(Succeed()) // Share row owned by the non-admin owner _, err := GetDBXBuilder().NewQuery(` @@ -178,20 +184,21 @@ var _ = Describe("ShareRepository", func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) b := GetDBXBuilder() _, _ = b.NewQuery(`DELETE FROM share WHERE id = 'share-scope'`).Execute() - pr := NewPlaylistRepository(adminCtx, b) - _ = pr.Delete(plsID) - mr := NewMediaFileRepository(adminCtx, b).(*mediaFileRepository) - _, _ = mr.executeSQL(squirrel.Delete("media_file").Where(squirrel.Eq{"id": []string{"share-other", "share-ok"}})) - lr := NewLibraryRepository(adminCtx, b).(*libraryRepository) - _ = lr.delete(squirrel.Eq{"id": otherLib.ID}) - _ = NewUserRepository(adminCtx, b).Delete(owner.ID) + pr := NewPlaylistRepository(b) + _ = pr.Delete(adminCtx, plsID) + mr := NewMediaFileRepository(b).(*mediaFileRepository) + _, _ = mr.executeSQL(adminCtx, squirrel.Delete("media_file").Where(squirrel.Eq{"id": []string{"share-other", "share-ok"}})) + lr := NewLibraryRepository(b).(*libraryRepository) + _ = lr.delete(adminCtx, squirrel.Eq{"id": otherLib.ID}) + _ = NewUserRepository(b).Delete(adminCtx, owner.ID) }) It("excludes tracks the owner cannot access from the shared playlist", func() { // Read the share as admin (mimics the public-share render path, which uses // the share repository's own context). loadMedia must scope to the owner. - adminRepo := NewShareRepository(request.WithUser(log.NewContext(GinkgoT().Context()), adminUser), GetDBXBuilder()) - share, err := adminRepo.Get("share-scope") + adminCtx := request.WithUser(ctx, adminUser) + adminRepo := NewShareRepository(GetDBXBuilder()) + share, err := adminRepo.Get(adminCtx, "share-scope") Expect(err).ToNot(HaveOccurred()) Expect(share.Tracks).To(ContainElement(HaveField("ID", "share-ok"))) @@ -205,11 +212,11 @@ var _ = Describe("ShareRepository", func() { // instead of panicking. privatePlsID := "private-pls" adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - pr := NewPlaylistRepository(adminCtx, GetDBXBuilder()) + pr := NewPlaylistRepository(GetDBXBuilder()) privatePls := &model.Playlist{ID: privatePlsID, Name: "Private", OwnerID: adminUser.ID, Public: false} privatePls.AddMediaFiles(model.MediaFiles{{ID: "share-ok"}}) - Expect(pr.Put(privatePls)).To(Succeed()) - DeferCleanup(func() { _ = pr.Delete(privatePlsID) }) + Expect(pr.Put(adminCtx, privatePls)).To(Succeed()) + DeferCleanup(func() { _ = pr.Delete(adminCtx, privatePlsID) }) _, err := GetDBXBuilder().NewQuery(` INSERT INTO share (id, user_id, description, resource_type, resource_ids, created_at, updated_at) @@ -221,8 +228,8 @@ var _ = Describe("ShareRepository", func() { Expect(err).ToNot(HaveOccurred()) DeferCleanup(func() { _, _ = GetDBXBuilder().NewQuery(`DELETE FROM share WHERE id = 'share-private'`).Execute() }) - adminRepo := NewShareRepository(adminCtx, GetDBXBuilder()) - share, err := adminRepo.Get("share-private") + adminRepo := NewShareRepository(GetDBXBuilder()) + share, err := adminRepo.Get(adminCtx, "share-private") Expect(err).ToNot(HaveOccurred()) Expect(share.Tracks).To(BeEmpty()) }) @@ -239,13 +246,13 @@ var _ = Describe("ShareRepository", func() { b := GetDBXBuilder() // A second library the owner has no access to - lr := NewLibraryRepository(adminCtx, b) + lr := NewLibraryRepository(b) otherLib = model.Library{ID: 0, Name: "Artist Share Other Library", Path: "/share/artist/other"} - Expect(lr.Put(&otherLib)).To(Succeed()) + Expect(lr.Put(adminCtx, &otherLib)).To(Succeed()) - ar := NewArtistRepository(adminCtx, b) - Expect(createArtistWithLibrary(ar, &model.Artist{ID: primaryID, Name: "AA Primary", OrderArtistName: "aa primary"}, 1)).To(Succeed()) - Expect(createArtistWithLibrary(ar, &model.Artist{ID: secondaryID, Name: "AA Secondary", OrderArtistName: "aa secondary"}, 1)).To(Succeed()) + ar := NewArtistRepository(b) + Expect(createArtistWithLibrary(adminCtx, ar, &model.Artist{ID: primaryID, Name: "AA Primary", OrderArtistName: "aa primary"}, 1)).To(Succeed()) + Expect(createArtistWithLibrary(adminCtx, ar, &model.Artist{ID: secondaryID, Name: "AA Secondary", OrderArtistName: "aa secondary"}, 1)).To(Succeed()) // Secondary is a co-album-artist (not the first): album_artist_id points at // primary, so the legacy-column filter would miss both tracks. @@ -253,19 +260,19 @@ var _ = Describe("ShareRepository", func() { {Artist: model.Artist{ID: primaryID, Name: "AA Primary"}}, {Artist: model.Artist{ID: secondaryID, Name: "AA Secondary"}}, }} - alr := NewAlbumRepository(adminCtx, b) - Expect(alr.Put(&model.Album{ID: "art-album-ok", LibraryID: 1, Name: "Art Album OK", AlbumArtistID: primaryID, AlbumArtist: "AA Primary", Participants: aaParticipants})).To(Succeed()) - Expect(alr.Put(&model.Album{ID: "art-album-other", LibraryID: otherLib.ID, Name: "Art Album Other", AlbumArtistID: primaryID, AlbumArtist: "AA Primary", Participants: aaParticipants})).To(Succeed()) + alr := NewAlbumRepository(b) + Expect(alr.Put(ctx, &model.Album{ID: "art-album-ok", LibraryID: 1, Name: "Art Album OK", AlbumArtistID: primaryID, AlbumArtist: "AA Primary", Participants: aaParticipants})).To(Succeed()) + Expect(alr.Put(ctx, &model.Album{ID: "art-album-other", LibraryID: otherLib.ID, Name: "Art Album Other", AlbumArtistID: primaryID, AlbumArtist: "AA Primary", Participants: aaParticipants})).To(Succeed()) - mr := NewMediaFileRepository(adminCtx, b) - Expect(mr.Put(&model.MediaFile{ID: "art-ok", LibraryID: 1, AlbumID: "art-album-ok", Path: "a/ok.mp3", Title: "ArtOK", AlbumArtistID: primaryID, Participants: aaParticipants})).To(Succeed()) - Expect(mr.Put(&model.MediaFile{ID: "art-other", LibraryID: otherLib.ID, AlbumID: "art-album-other", Path: "a/other.mp3", Title: "ArtOther", AlbumArtistID: primaryID, Participants: aaParticipants})).To(Succeed()) + mr := NewMediaFileRepository(b) + Expect(mr.Put(adminCtx, &model.MediaFile{ID: "art-ok", LibraryID: 1, AlbumID: "art-album-ok", Path: "a/ok.mp3", Title: "ArtOK", AlbumArtistID: primaryID, Participants: aaParticipants})).To(Succeed()) + Expect(mr.Put(adminCtx, &model.MediaFile{ID: "art-other", LibraryID: otherLib.ID, AlbumID: "art-album-other", Path: "a/other.mp3", Title: "ArtOther", AlbumArtistID: primaryID, Participants: aaParticipants})).To(Succeed()) // Non-admin owner with access to library 1 only owner = createUserWithLibraries("artist-share-owner", []int{1}) - ur := NewUserRepository(adminCtx, b) - Expect(ur.Put(&owner)).To(Succeed()) - Expect(ur.SetUserLibraries(owner.ID, []int{1})).To(Succeed()) + ur := NewUserRepository(b) + Expect(ur.Put(adminCtx, &owner)).To(Succeed()) + Expect(ur.SetUserLibraries(adminCtx, owner.ID, []int{1})).To(Succeed()) for _, s := range []struct{ id, typ, ids string }{ {"art-share", "artist", secondaryID}, @@ -287,22 +294,23 @@ var _ = Describe("ShareRepository", func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) b := GetDBXBuilder() _, _ = b.NewQuery(`DELETE FROM share WHERE id IN ('art-share', 'art-album-share', 'art-mf-share')`).Execute() - mr := NewMediaFileRepository(adminCtx, b).(*mediaFileRepository) - _, _ = mr.executeSQL(squirrel.Delete("media_file").Where(squirrel.Eq{"id": []string{"art-ok", "art-other"}})) - alr := NewAlbumRepository(adminCtx, b).(*albumRepository) - _, _ = alr.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"art-album-ok", "art-album-other"}})) - ar := NewArtistRepository(adminCtx, b).(*artistRepository) - _, _ = ar.executeSQL(squirrel.Delete("artist").Where(squirrel.Eq{"id": []string{primaryID, secondaryID}})) - lr := NewLibraryRepository(adminCtx, b).(*libraryRepository) - _ = lr.delete(squirrel.Eq{"id": otherLib.ID}) - _ = NewUserRepository(adminCtx, b).Delete(owner.ID) + mr := NewMediaFileRepository(b).(*mediaFileRepository) + _, _ = mr.executeSQL(adminCtx, squirrel.Delete("media_file").Where(squirrel.Eq{"id": []string{"art-ok", "art-other"}})) + alr := NewAlbumRepository(b).(*albumRepository) + _, _ = alr.executeSQL(adminCtx, squirrel.Delete("album").Where(squirrel.Eq{"id": []string{"art-album-ok", "art-album-other"}})) + ar := NewArtistRepository(b).(*artistRepository) + _, _ = ar.executeSQL(adminCtx, squirrel.Delete("artist").Where(squirrel.Eq{"id": []string{primaryID, secondaryID}})) + lr := NewLibraryRepository(b).(*libraryRepository) + _ = lr.delete(adminCtx, squirrel.Eq{"id": otherLib.ID}) + _ = NewUserRepository(b).Delete(adminCtx, owner.ID) }) It("includes co-album-artist tracks the owner can access and excludes those they cannot", func() { // Read as admin (mimics the public-share render path); loadMedia must still // scope to the owner's libraries. - adminRepo := NewShareRepository(request.WithUser(log.NewContext(GinkgoT().Context()), adminUser), GetDBXBuilder()) - share, err := adminRepo.Get("art-share") + adminCtx := request.WithUser(ctx, adminUser) + adminRepo := NewShareRepository(GetDBXBuilder()) + share, err := adminRepo.Get(adminCtx, "art-share") Expect(err).ToNot(HaveOccurred()) Expect(share.Tracks).To(ContainElement(HaveField("ID", "art-ok")), @@ -318,7 +326,7 @@ var _ = Describe("ShareRepository", func() { It("excludes albums and their tracks outside the owner's libraries from an album share", func() { // Public share rendering has no user in the context. - share, err := NewShareRepository(log.NewContext(GinkgoT().Context()), GetDBXBuilder()).Get("art-album-share") + share, err := NewShareRepository(GetDBXBuilder()).Get(log.NewContext(GinkgoT().Context()), "art-album-share") Expect(err).ToNot(HaveOccurred()) Expect(share.Albums).To(ContainElement(HaveField("ID", "art-album-ok"))) Expect(share.Albums).ToNot(ContainElement(HaveField("ID", "art-album-other"))) @@ -327,7 +335,7 @@ var _ = Describe("ShareRepository", func() { }) It("excludes tracks outside the owner's libraries from a media file share", func() { - share, err := NewShareRepository(log.NewContext(GinkgoT().Context()), GetDBXBuilder()).Get("art-mf-share") + share, err := NewShareRepository(GetDBXBuilder()).Get(log.NewContext(GinkgoT().Context()), "art-mf-share") Expect(err).ToNot(HaveOccurred()) Expect(share.Tracks).To(ContainElement(HaveField("ID", "art-ok"))) Expect(share.Tracks).ToNot(ContainElement(HaveField("ID", "art-other"))) @@ -358,58 +366,59 @@ var _ = Describe("ShareRepository", func() { It("allows a non-admin user to delete their own share", func() { insertShare("own-share-del", ownerUser.ID) ctx := request.WithUser(log.NewContext(GinkgoT().Context()), ownerUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Delete("own-share-del") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Delete(ctx, "own-share-del") Expect(err).ToNot(HaveOccurred()) }) It("denies a non-admin user from deleting another user's share", func() { insertShare("other-share-del", ownerUser.ID) ctx := request.WithUser(log.NewContext(GinkgoT().Context()), otherUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Delete("other-share-del") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Delete(ctx, "other-share-del") Expect(err).To(Equal(rest.ErrPermissionDenied)) // The share was not deleted: the owner can still read it. ownerCtx := request.WithUser(log.NewContext(GinkgoT().Context()), ownerUser) - ownerRepo := NewShareRepository(ownerCtx, GetDBXBuilder()) - _, err = ownerRepo.(rest.Repository).Read("other-share-del") + ownerRepo := NewShareRepository(GetDBXBuilder()) + _, err = ownerRepo.Read(ownerCtx, "other-share-del") Expect(err).ToNot(HaveOccurred()) }) It("allows an admin to delete any user's share", func() { insertShare("admin-del-share", ownerUser.ID) ctx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Delete("admin-del-share") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Delete(ctx, "admin-del-share") Expect(err).ToNot(HaveOccurred()) }) It("allows headless context (no user) to delete a share", func() { insertShare("headless-del-share", ownerUser.ID) - repo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - err := repo.(rest.Persistable).Delete("headless-del-share") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Delete(GinkgoT().Context(), "headless-del-share") Expect(err).ToNot(HaveOccurred()) }) }) Describe("Save", func() { It("assigns the logged-in user as owner, ignoring a client-supplied UserID", func() { - ur := NewUserRepository(ctx, GetDBXBuilder()) - Expect(ur.Put(&ownerUser)).To(Succeed()) - Expect(ur.Put(&otherUser)).To(Succeed()) + ur := NewUserRepository(GetDBXBuilder()) + Expect(ur.Put(ctx, &ownerUser)).To(Succeed()) + Expect(ur.Put(ctx, &otherUser)).To(Succeed()) attackerCtx := request.WithUser(log.NewContext(GinkgoT().Context()), ownerUser) - attackerRepo := NewShareRepository(attackerCtx, GetDBXBuilder()).(rest.Persistable) + attackerRepo := NewShareRepository(GetDBXBuilder()) - id, err := attackerRepo.Save(&model.Share{ + id, err := attackerRepo.Save(attackerCtx, &model.Share{ ID: "spoof-save-share", UserID: otherUser.ID, ResourceType: "media_file", ResourceIDs: "1001", }) Expect(err).ToNot(HaveOccurred()) - adminRepo := NewShareRepository(request.WithUser(log.NewContext(GinkgoT().Context()), adminUser), GetDBXBuilder()) - got, err := adminRepo.Get(id) + adminCtx := request.WithUser(ctx, adminUser) + adminRepo := NewShareRepository(GetDBXBuilder()) + got, err := adminRepo.Get(adminCtx, id) Expect(err).ToNot(HaveOccurred()) Expect(got.UserID).To(Equal(ownerUser.ID)) }) @@ -419,53 +428,52 @@ var _ = Describe("ShareRepository", func() { It("allows a non-admin user to update their own share", func() { insertShare("own-share-upd", ownerUser.ID) ctx := request.WithUser(log.NewContext(GinkgoT().Context()), ownerUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Update("own-share-upd", &model.Share{Description: "Updated"}, "description") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Update(ctx, "own-share-upd", model.Share{Description: "Updated"}, "description") Expect(err).ToNot(HaveOccurred()) }) It("denies a non-admin user from updating another user's share", func() { insertShare("other-share-upd", ownerUser.ID) ctx := request.WithUser(log.NewContext(GinkgoT().Context()), otherUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Update("other-share-upd", &model.Share{Description: "Hacked"}, "description") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Update(ctx, "other-share-upd", model.Share{Description: "Hacked"}, "description") Expect(err).To(Equal(rest.ErrPermissionDenied)) }) It("allows an admin to update any user's share", func() { insertShare("admin-upd-share", ownerUser.ID) ctx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Update("admin-upd-share", &model.Share{Description: "Admin Updated"}, "description") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Update(ctx, "admin-upd-share", model.Share{Description: "Admin Updated"}, "description") Expect(err).ToNot(HaveOccurred()) }) It("allows headless context (no user) to update a share", func() { insertShare("headless-upd-share", ownerUser.ID) - repo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - err := repo.(rest.Persistable).Update("headless-upd-share", &model.Share{Description: "Headless"}, "description") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Update(GinkgoT().Context(), "headless-upd-share", model.Share{Description: "Headless"}, "description") Expect(err).ToNot(HaveOccurred()) }) It("returns not found when updating a nonexistent share", func() { ctx := request.WithUser(log.NewContext(context.TODO()), ownerUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Update("does-not-exist", &model.Share{Description: "Ghost"}, "description") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Update(ctx, "does-not-exist", model.Share{Description: "Ghost"}, "description") Expect(err).To(Equal(rest.ErrNotFound)) }) It("updates all columns when no specific columns are given", func() { insertShare("all-cols-share", ownerUser.ID) ctx := request.WithUser(log.NewContext(context.TODO()), ownerUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) + repo := NewShareRepository(GetDBXBuilder()) // No cols: the update must write every column, not just updated_at. - err := repo.(rest.Persistable).Update("all-cols-share", - &model.Share{Description: "All Updated", MaxBitRate: 192, ResourceType: "album", ResourceIDs: "2002"}) + err := repo.Update(ctx, "all-cols-share", + model.Share{Description: "All Updated", MaxBitRate: 192, ResourceType: "album", ResourceIDs: "2002"}) Expect(err).ToNot(HaveOccurred()) - got, err := repo.(rest.Repository).Read("all-cols-share") + share, err := repo.Read(ctx, "all-cols-share") Expect(err).ToNot(HaveOccurred()) - share := got.(*model.Share) Expect(share.Description).To(Equal("All Updated")) Expect(share.MaxBitRate).To(Equal(192)) Expect(share.ResourceType).To(Equal("album")) @@ -474,24 +482,24 @@ var _ = Describe("ShareRepository", func() { It("does not let an owner reassign their share to another user", func() { insertShare("reassign-share", ownerUser.ID) ctx := request.WithUser(log.NewContext(context.TODO()), ownerUser) - repo := NewShareRepository(ctx, GetDBXBuilder()) - err := repo.(rest.Persistable).Update("reassign-share", - &model.Share{UserID: otherUser.ID, Description: "Given away"}, "user_id", "description") + repo := NewShareRepository(GetDBXBuilder()) + err := repo.Update(ctx, "reassign-share", + model.Share{UserID: otherUser.ID, Description: "Given away"}, "user_id", "description") Expect(err).ToNot(HaveOccurred()) // Ownership must not have moved, even though user_id was passed in the body and cols. - got, err := repo.(rest.Repository).Read("reassign-share") + got, err := repo.Read(ctx, "reassign-share") Expect(err).ToNot(HaveOccurred()) - Expect(got.(*model.Share).UserID).To(Equal(ownerUser.ID)) + Expect(got.UserID).To(Equal(ownerUser.ID)) }) }) Describe("Read scoping", func() { BeforeEach(func() { // Persist owner/other users so the JOIN in selectShare resolves. - ur := NewUserRepository(ctx, GetDBXBuilder()) - Expect(ur.Put(&ownerUser)).To(Succeed()) - Expect(ur.Put(&otherUser)).To(Succeed()) + ur := NewUserRepository(GetDBXBuilder()) + Expect(ur.Put(ctx, &ownerUser)).To(Succeed()) + Expect(ur.Put(ctx, &otherUser)).To(Succeed()) insertShare("share-owner-1", ownerUser.ID) insertShare("share-owner-2", ownerUser.ID) @@ -500,16 +508,15 @@ var _ = Describe("ShareRepository", func() { Context("non-admin user", func() { var nonAdminRepo model.ShareRepository - var nonAdminRest rest.Repository + var nonAdminCtx context.Context BeforeEach(func() { - nonAdminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), ownerUser) - nonAdminRepo = NewShareRepository(nonAdminCtx, GetDBXBuilder()) - nonAdminRest = nonAdminRepo.(rest.Repository) + nonAdminCtx = request.WithUser(ctx, ownerUser) + nonAdminRepo = NewShareRepository(GetDBXBuilder()) }) It("GetAll returns only own shares", func() { - shares, err := nonAdminRepo.GetAll() + shares, err := nonAdminRepo.GetAll(nonAdminCtx) Expect(err).ToNot(HaveOccurred()) ids := make([]string, len(shares)) for i, s := range shares { @@ -519,9 +526,8 @@ var _ = Describe("ShareRepository", func() { }) It("ReadAll returns only own shares", func() { - res, err := nonAdminRest.ReadAll() + shares, err := nonAdminRepo.ReadAll(nonAdminCtx) Expect(err).ToNot(HaveOccurred()) - shares := res.(model.Shares) ids := make([]string, len(shares)) for i, s := range shares { ids[i] = s.ID @@ -530,41 +536,41 @@ var _ = Describe("ShareRepository", func() { }) It("Get returns own share", func() { - s, err := nonAdminRepo.Get("share-owner-1") + s, err := nonAdminRepo.Get(nonAdminCtx, "share-owner-1") Expect(err).ToNot(HaveOccurred()) Expect(s.ID).To(Equal("share-owner-1")) }) It("Get returns ErrNotFound for another user's share", func() { - _, err := nonAdminRepo.Get("share-other-1") + _, err := nonAdminRepo.Get(nonAdminCtx, "share-other-1") Expect(err).To(MatchError(model.ErrNotFound)) }) It("Read returns ErrNotFound for another user's share", func() { - _, err := nonAdminRest.Read("share-other-1") + _, err := nonAdminRepo.Read(nonAdminCtx, "share-other-1") Expect(err).To(MatchError(model.ErrNotFound)) }) It("Exists returns true for own share", func() { - exists, err := nonAdminRepo.Exists("share-owner-1") + exists, err := nonAdminRepo.Exists(nonAdminCtx, "share-owner-1") Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeTrue()) }) It("Exists returns false for another user's share", func() { - exists, err := nonAdminRepo.Exists("share-other-1") + exists, err := nonAdminRepo.Exists(nonAdminCtx, "share-other-1") Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeFalse()) }) It("CountAll counts only own shares", func() { - count, err := nonAdminRepo.CountAll() + count, err := nonAdminRepo.CountAll(nonAdminCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeNumerically("==", 2)) }) It("Count (rest) counts only own shares", func() { - count, err := nonAdminRest.Count() + count, err := nonAdminRepo.Count(nonAdminCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeNumerically("==", 2)) }) @@ -573,8 +579,8 @@ var _ = Describe("ShareRepository", func() { Context("admin user", func() { It("GetAll returns all shares", func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - adminRepo := NewShareRepository(adminCtx, GetDBXBuilder()) - shares, err := adminRepo.GetAll() + adminRepo := NewShareRepository(GetDBXBuilder()) + shares, err := adminRepo.GetAll(adminCtx) Expect(err).ToNot(HaveOccurred()) ids := make([]string, len(shares)) for i, s := range shares { @@ -585,31 +591,37 @@ var _ = Describe("ShareRepository", func() { It("CountAll counts all shares", func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - adminRepo := NewShareRepository(adminCtx, GetDBXBuilder()) - count, err := adminRepo.CountAll() + adminRepo := NewShareRepository(GetDBXBuilder()) + count, err := adminRepo.CountAll(adminCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeNumerically("==", 3)) }) }) Context("headless context (public share route)", func() { + var headlessCtx context.Context + + BeforeEach(func() { + headlessCtx = GinkgoT().Context() + }) + It("GetAll returns all shares", func() { - headlessRepo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - shares, err := headlessRepo.GetAll() + headlessRepo := NewShareRepository(GetDBXBuilder()) + shares, err := headlessRepo.GetAll(headlessCtx) Expect(err).ToNot(HaveOccurred()) Expect(shares).To(HaveLen(3)) }) It("Get returns another user's share", func() { - headlessRepo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - s, err := headlessRepo.Get("share-other-1") + headlessRepo := NewShareRepository(GetDBXBuilder()) + s, err := headlessRepo.Get(headlessCtx, "share-other-1") Expect(err).ToNot(HaveOccurred()) Expect(s.ID).To(Equal("share-other-1")) }) It("Exists returns true for any share", func() { - headlessRepo := NewShareRepository(GinkgoT().Context(), GetDBXBuilder()) - exists, err := headlessRepo.Exists("share-other-1") + headlessRepo := NewShareRepository(GetDBXBuilder()) + exists, err := headlessRepo.Exists(headlessCtx, "share-other-1") Expect(err).ToNot(HaveOccurred()) Expect(exists).To(BeTrue()) }) diff --git a/persistence/smart_playlist_repository.go b/persistence/smart_playlist_repository.go index 24c6f5fc5..871f53dd2 100644 --- a/persistence/smart_playlist_repository.go +++ b/persistence/smart_playlist_repository.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "slices" "time" @@ -18,76 +19,76 @@ import ( // configured refresh delay. // refreshSmartPlaylist evaluates the criteria of a smart playlist and updates its tracks accordingly. -func (r *playlistRepository) refreshSmartPlaylist(pls *model.Playlist) bool { - return r.refreshSmartPlaylistTree(pls, map[string]struct{}{}) +func (r *playlistRepository) refreshSmartPlaylist(ctx context.Context, pls *model.Playlist) bool { + return r.refreshSmartPlaylistTree(ctx, pls, map[string]struct{}{}) } // The visited set stops playlists that reference each other from recursing forever. -func (r *playlistRepository) refreshSmartPlaylistTree(pls *model.Playlist, visited map[string]struct{}) bool { +func (r *playlistRepository) refreshSmartPlaylistTree(ctx context.Context, pls *model.Playlist, visited map[string]struct{}) bool { if _, seen := visited[pls.ID]; seen { - log.Trace(r.ctx, "Skipping already visited smart playlist", "playlist", pls.Name, "id", pls.ID) + log.Trace(ctx, "Skipping already visited smart playlist", "playlist", pls.Name, "id", pls.ID) return false } visited[pls.ID] = struct{}{} - usr := loggedUser(r.ctx) - if !r.shouldRefreshSmartPlaylist(pls, usr) { + usr := loggedUser(ctx) + if !r.shouldRefreshSmartPlaylist(ctx, pls, usr) { return false } - log.Debug(r.ctx, "Refreshing smart playlist", "playlist", pls.Name, "id", pls.ID) + log.Debug(ctx, "Refreshing smart playlist", "playlist", pls.Name, "id", pls.ID) start := time.Now() del := Delete("playlist_tracks").Where(Eq{"playlist_id": pls.ID}) - if _, err := r.executeSQL(del); err != nil { - log.Error(r.ctx, "Error deleting old smart playlist tracks", "playlist", pls.Name, "id", pls.ID, err) + if _, err := r.executeSQL(ctx, del); err != nil { + log.Error(ctx, "Error deleting old smart playlist tracks", "playlist", pls.Name, "id", pls.ID, err) return false } rulesSQL := newSmartPlaylistCriteria(*pls.NormalizedRules(), withSmartPlaylistOwner(*usr)) - if !r.refreshChildPlaylists(pls, rulesSQL, visited) { + if !r.refreshChildPlaylists(ctx, pls, rulesSQL, visited) { return false } - if err := r.resolvePercentageLimit(pls, &rulesSQL, usr.ID); err != nil { + if err := r.resolvePercentageLimit(ctx, pls, &rulesSQL, usr.ID); err != nil { return false } - sq := r.buildSmartPlaylistQuery(pls, rulesSQL, usr.ID) + sq := r.buildSmartPlaylistQuery(ctx, pls, rulesSQL, usr.ID) sq, err := r.addCriteria(sq, rulesSQL) if err != nil { - log.Error(r.ctx, "Error building smart playlist criteria", "playlist", pls.Name, "id", pls.ID, err) + log.Error(ctx, "Error building smart playlist criteria", "playlist", pls.Name, "id", pls.ID, err) return false } insSql := Insert("playlist_tracks").Columns("id", "playlist_id", "media_file_id").Select(sq) - if _, err = r.executeSQL(insSql); err != nil { - log.Error(r.ctx, "Error refreshing smart playlist tracks", "playlist", pls.Name, "id", pls.ID, err) + if _, err = r.executeSQL(ctx, insSql); err != nil { + log.Error(ctx, "Error refreshing smart playlist tracks", "playlist", pls.Name, "id", pls.ID, err) return false } - if err = r.refreshCounters(pls); err != nil { - log.Error(r.ctx, "Error updating smart playlist stats", "playlist", pls.Name, "id", pls.ID, err) + if err = r.refreshCounters(ctx, pls); err != nil { + log.Error(ctx, "Error updating smart playlist stats", "playlist", pls.Name, "id", pls.ID, err) return false } // Reuse the stamp refreshCounters just wrote, so evaluated_at and updated_at agree now := pls.UpdatedAt updSql := Update(r.tableName).Set("evaluated_at", now).Where(Eq{"id": pls.ID}) - if _, err = r.executeSQL(updSql); err != nil { - log.Error(r.ctx, "Error updating smart playlist", "playlist", pls.Name, "id", pls.ID, err) + if _, err = r.executeSQL(ctx, updSql); err != nil { + log.Error(ctx, "Error updating smart playlist", "playlist", pls.Name, "id", pls.ID, err) return false } pls.EvaluatedAt = &now - log.Debug(r.ctx, "Refreshed playlist", "playlist", pls.Name, "id", pls.ID, "numTracks", pls.SongCount, "elapsed", time.Since(start)) + log.Debug(ctx, "Refreshed playlist", "playlist", pls.Name, "id", pls.ID, "numTracks", pls.SongCount, "elapsed", time.Since(start)) return true } // shouldRefreshSmartPlaylist determines if a smart playlist needs to be refreshed based on its type, last evaluated // time, and ownership. -func (r *playlistRepository) shouldRefreshSmartPlaylist(pls *model.Playlist, usr *model.User) bool { +func (r *playlistRepository) shouldRefreshSmartPlaylist(ctx context.Context, pls *model.Playlist, usr *model.User) bool { if !pls.IsSmartPlaylist() { return false } @@ -95,7 +96,7 @@ func (r *playlistRepository) shouldRefreshSmartPlaylist(pls *model.Playlist, usr return false } if pls.OwnerID != usr.ID { - log.Trace(r.ctx, "Not refreshing smart playlist from other user", "playlist", pls.Name, "id", pls.ID) + log.Trace(ctx, "Not refreshing smart playlist from other user", "playlist", pls.Name, "id", pls.ID) return false } return true @@ -103,7 +104,7 @@ func (r *playlistRepository) shouldRefreshSmartPlaylist(pls *model.Playlist, usr // refreshChildPlaylists handles refreshing any child playlists that are referenced in the smart playlist criteria. // Returns false if child playlists could not be loaded (DB error), signaling the parent refresh should abort. -func (r *playlistRepository) refreshChildPlaylists(pls *model.Playlist, rulesSQL smartPlaylistCriteria, visited map[string]struct{}) bool { +func (r *playlistRepository) refreshChildPlaylists(ctx context.Context, pls *model.Playlist, rulesSQL smartPlaylistCriteria, visited map[string]struct{}) bool { childPlaylistIds := rulesSQL.ChildPlaylistIds() childPlaylistPaths := rulesSQL.ChildPlaylistPaths() if len(childPlaylistIds) == 0 && len(childPlaylistPaths) == 0 { @@ -119,9 +120,9 @@ func (r *playlistRepository) refreshChildPlaylists(pls *model.Playlist, rulesSQL conditions = append(conditions, Eq{"playlist.path": lookupPaths}) } - childPlaylists, err := r.GetAll(model.QueryOptions{Filters: conditions}) + childPlaylists, err := r.GetAll(ctx, model.QueryOptions{Filters: conditions}) if err != nil { - log.Error(r.ctx, "Error loading child playlists for smart playlist refresh", "playlist", pls.Name, "id", pls.ID, "childIds", childPlaylistIds, "childPaths", childPlaylistPaths, err) + log.Error(ctx, "Error loading child playlists for smart playlist refresh", "playlist", pls.Name, "id", pls.ID, "childIds", childPlaylistIds, "childPaths", childPlaylistPaths, err) return false } @@ -131,58 +132,58 @@ func (r *playlistRepository) refreshChildPlaylists(pls *model.Playlist, rulesSQL if childPlaylists[i].Path != "" { found[norm.NFC.String(childPlaylists[i].Path)] = struct{}{} } - r.refreshSmartPlaylistTree(&childPlaylists[i], visited) + r.refreshSmartPlaylistTree(ctx, &childPlaylists[i], visited) } for _, id := range childPlaylistIds { if _, ok := found[id]; !ok { - log.Warn(r.ctx, "Referenced playlist is not accessible to smart playlist owner", "playlist", pls.Name, "id", pls.ID, "childId", id, "ownerId", pls.OwnerID) + log.Warn(ctx, "Referenced playlist is not accessible to smart playlist owner", "playlist", pls.Name, "id", pls.ID, "childId", id, "ownerId", pls.OwnerID) } } for _, path := range childPlaylistPaths { if _, ok := found[norm.NFC.String(path)]; !ok { - log.Warn(r.ctx, "Referenced playlist is not accessible to smart playlist owner", "playlist", pls.Name, "id", pls.ID, "path", path, "ownerId", pls.OwnerID) + log.Warn(ctx, "Referenced playlist is not accessible to smart playlist owner", "playlist", pls.Name, "id", pls.ID, "path", path, "ownerId", pls.OwnerID) } } return true } // resolvePercentageLimit calculates the actual limit for a smart playlist criteria that uses a percentage-based limit. -func (r *playlistRepository) resolvePercentageLimit(pls *model.Playlist, rulesSQL *smartPlaylistCriteria, userID string) error { +func (r *playlistRepository) resolvePercentageLimit(ctx context.Context, pls *model.Playlist, rulesSQL *smartPlaylistCriteria, userID string) error { if !rulesSQL.IsPercentageLimit() { return nil } countSq := Select("count(*) as count").From("media_file") countSq = rulesSQL.applyExpressionJoins(countSq, userID) - countSq = r.applyLibraryFilter(countSq, "media_file") + countSq = r.applyLibraryFilter(ctx, countSq, "media_file") cond, err := rulesSQL.where() if err != nil { - log.Error(r.ctx, "Error building smart playlist criteria", "playlist", pls.Name, "id", pls.ID, err) + log.Error(ctx, "Error building smart playlist criteria", "playlist", pls.Name, "id", pls.ID, err) return err } countSq = countSq.Where(cond) var res struct{ Count int64 } - if err = r.queryOne(countSq, &res); err != nil { - log.Error(r.ctx, "Error counting matching tracks for percentage limit", "playlist", pls.Name, "id", pls.ID, err) + if err = r.queryOne(ctx, countSq, &res); err != nil { + log.Error(ctx, "Error counting matching tracks for percentage limit", "playlist", pls.Name, "id", pls.ID, err) return err } rulesSQL.ResolveLimit(res.Count) - log.Debug(r.ctx, "Resolved percentage limit", "playlist", pls.Name, "percent", rulesSQL.LimitPercent, "totalMatching", res.Count, "resolvedLimit", rulesSQL.Limit) + log.Debug(ctx, "Resolved percentage limit", "playlist", pls.Name, "percent", rulesSQL.LimitPercent, "totalMatching", res.Count, "resolvedLimit", rulesSQL.Limit) return nil } // buildSmartPlaylistQuery constructs the SQL query to select media files matching the smart playlist criteria, // including the joins its fields require and library filtering. -func (r *playlistRepository) buildSmartPlaylistQuery(pls *model.Playlist, rulesSQL smartPlaylistCriteria, userID string) SelectBuilder { +func (r *playlistRepository) buildSmartPlaylistQuery(ctx context.Context, pls *model.Playlist, rulesSQL smartPlaylistCriteria, userID string) SelectBuilder { orderBy := rulesSQL.orderBy() sq := Select("row_number() over (order by "+orderBy+") as id", "'"+pls.ID+"' as playlist_id", "media_file.id as media_file_id"). From("media_file") sq = rulesSQL.applyRequiredJoins(sq, userID) - sq = r.applyLibraryFilter(sq, "media_file") + sq = r.applyLibraryFilter(ctx, sq, "media_file") return sq } diff --git a/persistence/smart_playlist_repository_test.go b/persistence/smart_playlist_repository_test.go index 33da0c1b7..4d4c5edb0 100644 --- a/persistence/smart_playlist_repository_test.go +++ b/persistence/smart_playlist_repository_test.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "path/filepath" "time" @@ -18,11 +19,11 @@ import ( var _ = Describe("PlaylistRepository - Smart Playlists", func() { var repo model.PlaylistRepository + var ctx context.Context BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - repo = NewPlaylistRepository(ctx, GetDBXBuilder()) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid", UserName: "userid", IsAdmin: true}) + repo = NewPlaylistRepository(GetDBXBuilder()) }) Context("Smart Playlists", func() { @@ -37,10 +38,10 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { Context("valid rules", func() { Specify("Put/Get", func() { newPls := model.Playlist{Name: "Great!", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(newPls.ID) }) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, newPls.ID) }) - savedPls, err := repo.Get(newPls.ID) + savedPls, err := repo.Get(ctx, newPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(savedPls.Rules).To(Equal(rules)) }) @@ -49,13 +50,13 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { Context("after an evaluation", func() { It("stamps updated_at and evaluated_at with the same instant", func() { newPls := model.Playlist{Name: "Evaluated", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(newPls.ID) }) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, newPls.ID) }) - refreshed, err := repo.GetWithTracks(newPls.ID, true, false) + refreshed, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) - stored, err := repo.Get(newPls.ID) + stored, err := repo.Get(ctx, newPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(stored.EvaluatedAt).ToNot(BeNil()) Expect(stored.UpdatedAt).To(BeTemporally("==", *stored.EvaluatedAt)) @@ -72,7 +73,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "Great!", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(MatchError(ContainSubstring("invalid criteria expression"))) + Expect(repo.Put(ctx, &newPls)).To(MatchError(ContainSubstring("invalid criteria expression"))) }) }) @@ -86,14 +87,14 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } pls := model.Playlist{Name: "Smart", OwnerID: "userid", Rules: rules, Path: "/music/smart.nsp", Sync: true} - Expect(repo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(pls.ID) }) + Expect(repo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, pls.ID) }) - evaluated, err := repo.GetWithTracks(pls.ID, true, false) + evaluated, err := repo.GetWithTracks(ctx, pls.ID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(evaluated.SongCount).To(BeNumerically(">", 0)) - stored, err := repo.Get(pls.ID) + stored, err := repo.Get(ctx, pls.ID) Expect(err).ToNot(HaveOccurred()) Expect(stored.SongCount).To(Equal(evaluated.SongCount)) @@ -101,9 +102,9 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { ID: pls.ID, Name: pls.Name, OwnerID: "userid", Rules: rules, Path: pls.Path, Sync: true, } - Expect(repo.Put(&reimported)).To(Succeed()) + Expect(repo.Put(ctx, &reimported)).To(Succeed()) - afterImport, err := repo.Get(pls.ID) + afterImport, err := repo.Get(ctx, pls.ID) Expect(err).ToNot(HaveOccurred()) Expect(afterImport.SongCount).To(Equal(stored.SongCount)) Expect(afterImport.Duration).To(Equal(stored.Duration)) @@ -126,8 +127,8 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } nestedPls := model.Playlist{Name: "Nested [ID]", OwnerID: "userid", Public: true, Rules: childRules} - Expect(repo.Put(&nestedPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(nestedPls.ID) }) + Expect(repo.Put(ctx, &nestedPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, nestedPls.ID) }) childRules = &criteria.Criteria{ Expression: criteria.All{ @@ -135,8 +136,8 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } nestedPathPls := model.Playlist{Name: "Nested [Path]", OwnerID: "userid", Path: "test.nsp", Public: true, Rules: childRules} - Expect(repo.Put(&nestedPathPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(nestedPathPls.ID) }) + Expect(repo.Put(ctx, &nestedPathPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, nestedPathPls.ID) }) parentPls := model.Playlist{Name: "Parent", OwnerID: "userid", Rules: &criteria.Criteria{ Expression: criteria.Any{ @@ -144,16 +145,16 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { criteria.InPlaylist{"path": nestedPathPls.Path}, }, }} - Expect(repo.Put(&parentPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(parentPls.ID) }) + Expect(repo.Put(ctx, &parentPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, parentPls.ID) }) // Nested playlist has not been evaluated yet - nestedPlsRead, err := repo.Get(nestedPls.ID) + nestedPlsRead, err := repo.Get(ctx, nestedPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(nestedPlsRead.EvaluatedAt).To(BeNil()) // Getting parent with refresh should recursively refresh the nested playlist - pls, err := repo.GetWithTracks(parentPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, parentPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(pls.EvaluatedAt).ToNot(BeNil()) Expect(*pls.EvaluatedAt).To(BeTemporally("~", time.Now(), 2*time.Second)) @@ -163,12 +164,12 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { Expect(pls.Tracks[0].MediaFileID).To(Equal(songDayInALife.ID)) // Nested playlists should now have been refreshed (EvaluatedAt set) - nestedPlsAfterParentGet, err := repo.Get(nestedPls.ID) + nestedPlsAfterParentGet, err := repo.Get(ctx, nestedPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(nestedPlsAfterParentGet.EvaluatedAt).ToNot(BeNil()) Expect(*nestedPlsAfterParentGet.EvaluatedAt).To(BeTemporally("~", time.Now(), 2*time.Second)) - nestedPlsAfterParentGet, err = repo.Get(nestedPathPls.ID) + nestedPlsAfterParentGet, err = repo.Get(ctx, nestedPathPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(nestedPlsAfterParentGet.EvaluatedAt).ToNot(BeNil()) Expect(*nestedPlsAfterParentGet.EvaluatedAt).To(BeTemporally("~", time.Now(), 2*time.Second)) @@ -181,19 +182,19 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { plsA := model.Playlist{Name: "Cycle A", OwnerID: "userid", Public: true, Rules: &criteria.Criteria{ Expression: criteria.All{criteria.Contains{"title": "Day"}}, }} - Expect(repo.Put(&plsA)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(plsA.ID) }) + Expect(repo.Put(ctx, &plsA)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, plsA.ID) }) plsB := model.Playlist{Name: "Cycle B", OwnerID: "userid", Public: true, Rules: &criteria.Criteria{ Expression: criteria.All{criteria.InPlaylist{"id": plsA.ID}}, }} - Expect(repo.Put(&plsB)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(plsB.ID) }) + Expect(repo.Put(ctx, &plsB)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, plsB.ID) }) plsA.Rules = &criteria.Criteria{Expression: criteria.All{criteria.InPlaylist{"id": plsB.ID}}} - Expect(repo.Put(&plsA)).To(Succeed()) + Expect(repo.Put(ctx, &plsA)).To(Succeed()) - _, err := repo.GetWithTracks(plsA.ID, true, false) + _, err := repo.GetWithTracks(ctx, plsA.ID, true, false) Expect(err).ToNot(HaveOccurred()) }) @@ -203,19 +204,19 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { bystander := model.Playlist{Name: "Bystander", OwnerID: "userid", Public: true, Rules: &criteria.Criteria{ Expression: criteria.All{criteria.Contains{"title": "Day"}}, }} - Expect(repo.Put(&bystander)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(bystander.ID) }) + Expect(repo.Put(ctx, &bystander)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, bystander.ID) }) parent := model.Playlist{Name: "Empty Path", OwnerID: "userid", Public: true, Rules: &criteria.Criteria{ Expression: criteria.All{criteria.InPlaylist{"path": ""}}, }} - Expect(repo.Put(&parent)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(parent.ID) }) + Expect(repo.Put(ctx, &parent)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, parent.ID) }) - _, err := repo.GetWithTracks(parent.ID, true, false) + _, err := repo.GetWithTracks(ctx, parent.ID, true, false) Expect(err).ToNot(HaveOccurred()) - reloaded, err := repo.Get(bystander.ID) + reloaded, err := repo.Get(ctx, bystander.ID) Expect(err).ToNot(HaveOccurred()) Expect(reloaded.EvaluatedAt).To(BeNil()) }) @@ -226,16 +227,16 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { child := model.Playlist{Name: "NFD Child", OwnerID: "userid", Public: true, Path: filepath.FromSlash("/mu\u0301sica/child.nsp"), Rules: &criteria.Criteria{ Expression: criteria.All{criteria.Contains{"title": "Day"}}, }} - Expect(repo.Put(&child)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(child.ID) }) + Expect(repo.Put(ctx, &child)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, child.ID) }) parent := model.Playlist{Name: "NFC Parent", OwnerID: "userid", Rules: &criteria.Criteria{ Expression: criteria.All{criteria.InPlaylist{"path": "/m\u00fasica/child.nsp"}}, }} - Expect(repo.Put(&parent)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(parent.ID) }) + Expect(repo.Put(ctx, &parent)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, parent.ID) }) - pls, err := repo.GetWithTracks(parent.ID, true, false) + pls, err := repo.GetWithTracks(ctx, parent.ID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(pls.Tracks).To(HaveLen(1)) Expect(pls.Tracks[0].MediaFileID).To(Equal(songDayInALife.ID)) @@ -252,8 +253,8 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } nestedPls := model.Playlist{Name: "Nested", OwnerID: "userid", Public: true, Rules: childRules, EvaluatedAt: &childEvaluatedAt} - Expect(repo.Put(&nestedPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(nestedPls.ID) }) + Expect(repo.Put(ctx, &nestedPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, nestedPls.ID) }) // Parent has no EvaluatedAt, so it WILL refresh, but the child should not parentPls := model.Playlist{Name: "Parent", OwnerID: "userid", Rules: &criteria.Criteria{ @@ -261,14 +262,14 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { criteria.InPlaylist{"id": nestedPls.ID}, }, }} - Expect(repo.Put(&parentPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(parentPls.ID) }) + Expect(repo.Put(ctx, &parentPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, parentPls.ID) }) - nestedPlsRead, err := repo.Get(nestedPls.ID) + nestedPlsRead, err := repo.Get(ctx, nestedPls.ID) Expect(err).ToNot(HaveOccurred()) // Getting parent with refresh should NOT recursively refresh the nested playlist - parent, err := repo.GetWithTracks(parentPls.ID, true, false) + parent, err := repo.GetWithTracks(ctx, parentPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) // Parent should have been refreshed (its EvaluatedAt was nil) @@ -276,7 +277,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { Expect(*parent.EvaluatedAt).To(BeTemporally("~", time.Now(), 2*time.Second)) // Nested playlist should NOT have been refreshed (still within delay window) - nestedPlsAfterParentGet, err := repo.Get(nestedPls.ID) + nestedPlsAfterParentGet, err := repo.Get(ctx, nestedPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(*nestedPlsAfterParentGet.EvaluatedAt).To(BeTemporally("~", childEvaluatedAt, time.Second)) Expect(*nestedPlsAfterParentGet.EvaluatedAt).To(Equal(*nestedPlsRead.EvaluatedAt)) @@ -297,10 +298,10 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { RefreshDelay: 24 * time.Hour, } pls := model.Playlist{Name: "Frozen Daily", OwnerID: "userid", Rules: rules, EvaluatedAt: &evaluatedAt} - Expect(repo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(pls.ID) }) + Expect(repo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, pls.ID) }) - got, err := repo.GetWithTracks(pls.ID, true, false) + got, err := repo.GetWithTracks(ctx, pls.ID, true, false) Expect(err).ToNot(HaveOccurred()) // Not re-evaluated: EvaluatedAt unchanged, no tracks materialized Expect(*got.EvaluatedAt).To(BeTemporally("~", evaluatedAt, time.Second)) @@ -316,10 +317,10 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { RefreshDelay: 5 * time.Minute, } pls := model.Playlist{Name: "Fast Refresh", OwnerID: "userid", Rules: rules, EvaluatedAt: &evaluatedAt} - Expect(repo.Put(&pls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(pls.ID) }) + Expect(repo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, pls.ID) }) - got, err := repo.GetWithTracks(pls.ID, true, false) + got, err := repo.GetWithTracks(ctx, pls.ID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(*got.EvaluatedAt).To(BeTemporally("~", time.Now(), 2*time.Second)) Expect(got.Tracks).To(HaveLen(1)) @@ -334,7 +335,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { AfterEach(func() { if testPlaylistID != "" { - Expect(repo.Delete(testPlaylistID)).To(BeNil()) + Expect(repo.Delete(ctx, testPlaylistID)).To(BeNil()) testPlaylistID = "" } }) @@ -344,12 +345,12 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { newPls := model.Playlist{Name: "Multi-Disc Test", OwnerID: "userid"} // Add tracks in intentionally scrambled order newPls.AddMediaFilesByID([]string{"2001", "2002", "2003", "2004"}) - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("retrieving tracks sorted by album") - tracksRepo := repo.Tracks(newPls.ID, false) - tracks, err := tracksRepo.GetAll(model.QueryOptions{Sort: "album", Order: "asc"}) + tracksRepo := repo.Tracks(ctx, newPls.ID, false) + tracks, err := tracksRepo.GetAll(ctx, model.QueryOptions{Sort: "album", Order: "asc"}) Expect(err).ToNot(HaveOccurred()) By("verifying tracks are sorted by disc number then track number") @@ -367,7 +368,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { AfterEach(func() { if testPlaylistID != "" { - _ = repo.Delete(testPlaylistID) + _ = repo.Delete(ctx, testPlaylistID) testPlaylistID = "" } }) @@ -381,11 +382,11 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "Starred Album Songs", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) trackIDs := make([]string, len(pls.Tracks)) @@ -404,11 +405,11 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "Starred Artist Songs", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) trackIDs := make([]string, len(pls.Tracks)) @@ -429,11 +430,11 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "Combined Album+Artist", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) trackIDs := make([]string, len(pls.Tracks)) @@ -451,11 +452,11 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "No Match", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) Expect(pls.Tracks).To(BeEmpty()) @@ -471,11 +472,11 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "String Loved Nested", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) trackIDs := make([]string, len(pls.Tracks)) @@ -495,8 +496,8 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } boolPls := model.Playlist{Name: "Bool Loved", OwnerID: "userid", Rules: boolRules} - Expect(repo.Put(&boolPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(boolPls.ID) }) + Expect(repo.Put(ctx, &boolPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, boolPls.ID) }) stringRules := &criteria.Criteria{ Expression: criteria.All{ @@ -506,13 +507,13 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } stringPls := model.Playlist{Name: "String Loved", OwnerID: "userid", Rules: stringRules} - Expect(repo.Put(&stringPls)).To(Succeed()) + Expect(repo.Put(ctx, &stringPls)).To(Succeed()) testPlaylistID = stringPls.ID conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second - boolResult, err := repo.GetWithTracks(boolPls.ID, true, false) + boolResult, err := repo.GetWithTracks(ctx, boolPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) - stringResult, err := repo.GetWithTracks(stringPls.ID, true, false) + stringResult, err := repo.GetWithTracks(ctx, stringPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) boolIDs := make([]string, len(boolResult.Tracks)) @@ -535,10 +536,10 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { trackIDsOf := func(rules *criteria.Criteria) []string { newPls := model.Playlist{Name: "Album Aggregates", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(newPls.ID) }) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, newPls.ID) }) - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) return slice.Map(pls.Tracks, func(t model.PlaylistTrack) string { return t.MediaFileID }) } @@ -569,7 +570,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { BeforeEach(func() { ctx := log.NewContext(GinkgoT().Context()) ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - mfRepo = NewMediaFileRepository(ctx, GetDBXBuilder()) + mfRepo = NewMediaFileRepository(GetDBXBuilder()) // Register 'grouping' as a valid tag for smart playlists criteria.AddTagNames([]string{"grouping"}) @@ -590,7 +591,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { LibraryID: 1, Lyrics: "[]", } - Expect(mfRepo.Put(&songWithGrouping)).To(Succeed()) + Expect(mfRepo.Put(ctx, &songWithGrouping)).To(Succeed()) // Create a song without the grouping tag songWithoutGrouping = model.MediaFile{ @@ -606,12 +607,12 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { LibraryID: 1, Lyrics: "[]", } - Expect(mfRepo.Put(&songWithoutGrouping)).To(Succeed()) + Expect(mfRepo.Put(ctx, &songWithoutGrouping)).To(Succeed()) }) AfterEach(func() { if testPlaylistID != "" { - _ = repo.Delete(testPlaylistID) + _ = repo.Delete(ctx, testPlaylistID) testPlaylistID = "" } // Clean up test media files @@ -629,12 +630,12 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "Tracks with Grouping", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("refreshing the smart playlist") conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second // Force refresh - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) By("verifying only the track with grouping tag is matched") @@ -650,12 +651,12 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "Tracks without Grouping", OwnerID: "userid", Rules: rules} - Expect(repo.Put(&newPls)).To(Succeed()) + Expect(repo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("refreshing the smart playlist") conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second // Force refresh - pls, err := repo.GetWithTracks(newPls.ID, true, false) + pls, err := repo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) By("verifying the track with grouping is NOT in the playlist") @@ -705,7 +706,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { // Create test media files in each library ctx := log.NewContext(GinkgoT().Context()) ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "userid", IsAdmin: true}) - mfRepo = NewMediaFileRepository(ctx, db) + mfRepo = NewMediaFileRepository(db) // Song in library 1 (accessible by restricted user) songLib1 := model.MediaFile{ @@ -721,7 +722,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { Tags: model.Tags{}, Lyrics: "[]", } - Expect(mfRepo.Put(&songLib1)).To(Succeed()) + Expect(mfRepo.Put(ctx, &songLib1)).To(Succeed()) // Song in library 2 (NOT accessible by restricted user) songLib2 := model.MediaFile{ @@ -737,13 +738,13 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { Tags: model.Tags{}, Lyrics: "[]", } - Expect(mfRepo.Put(&songLib2)).To(Succeed()) + Expect(mfRepo.Put(ctx, &songLib2)).To(Succeed()) }) AfterEach(func() { db := GetDBXBuilder() if testPlaylistID != "" { - _ = repo.Delete(testPlaylistID) + _ = repo.Delete(ctx, testPlaylistID) testPlaylistID = "" } // Clean up test data @@ -761,7 +762,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { // Create the smart playlist as the restricted user restrictedUser := model.User{ID: restrictedUserID, UserName: restrictedUserID, IsAdmin: false} ctx = request.WithUser(ctx, restrictedUser) - restrictedRepo := NewPlaylistRepository(ctx, db) + restrictedRepo := NewPlaylistRepository(db) // Create a smart playlist that matches all songs rules := &criteria.Criteria{ @@ -770,12 +771,12 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { }, } newPls := model.Playlist{Name: "All Songs", OwnerID: restrictedUserID, Rules: rules} - Expect(restrictedRepo.Put(&newPls)).To(Succeed()) + Expect(restrictedRepo.Put(ctx, &newPls)).To(Succeed()) testPlaylistID = newPls.ID By("refreshing the smart playlist") conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second // Force refresh - pls, err := restrictedRepo.GetWithTracks(newPls.ID, true, false) + pls, err := restrictedRepo.GetWithTracks(ctx, newPls.ID, true, false) Expect(err).ToNot(HaveOccurred()) By("verifying only the track from library 1 is in the playlist") diff --git a/persistence/sort_index_coverage_test.go b/persistence/sort_index_coverage_test.go index b5dea231d..138cb32c8 100644 --- a/persistence/sort_index_coverage_test.go +++ b/persistence/sort_index_coverage_test.go @@ -52,7 +52,7 @@ var _ = Describe("Sort index coverage", func() { { table: "media_file", newRepo: func(ctx context.Context) *sqlRepository { - return &NewMediaFileRepository(ctx, GetDBXBuilder()).(*mediaFileRepository).sqlRepository + return &NewMediaFileRepository(GetDBXBuilder()).(*mediaFileRepository).sqlRepository }, exceptions: map[string]string{ "random": "not a column sort", @@ -67,7 +67,7 @@ var _ = Describe("Sort index coverage", func() { { table: "album", newRepo: func(ctx context.Context) *sqlRepository { - return &NewAlbumRepository(ctx, GetDBXBuilder()).(*albumRepository).sqlRepository + return &NewAlbumRepository(GetDBXBuilder()).(*albumRepository).sqlRepository }, exceptions: map[string]string{ "random": "not a column sort", @@ -79,7 +79,7 @@ var _ = Describe("Sort index coverage", func() { { table: "artist", newRepo: func(ctx context.Context) *sqlRepository { - return &NewArtistRepository(ctx, GetDBXBuilder()).(*artistRepository).sqlRepository + return &NewArtistRepository(GetDBXBuilder()).(*artistRepository).sqlRepository }, exceptions: map[string]string{ //nolint:gosec // G101 false positive, same as the artist sortMappings "starred_at": "sorts on annotation join columns", diff --git a/persistence/sql_annotations.go b/persistence/sql_annotations.go index 27445b886..51c7dc146 100644 --- a/persistence/sql_annotations.go +++ b/persistence/sql_annotations.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "database/sql" "errors" "fmt" @@ -59,8 +60,8 @@ func filtersNeedAnnotation(query SelectBuilder) bool { return annotationColumnRE().MatchString(sql) } -func (r sqlRepository) withAnnotation(query SelectBuilder, idField string) SelectBuilder { - userID := loggedUser(r.ctx).ID +func (r sqlRepository) withAnnotation(ctx context.Context, query SelectBuilder, idField string) SelectBuilder { + userID := loggedUser(ctx).ID if userID == invalidUserId { return query.Columns(fmt.Sprintf("%s.average_rating", r.tableName)) } @@ -102,8 +103,8 @@ func annotationBoolFilter(field string) func(string, any) Sqlizer { } } -func (r sqlRepository) annId(itemID ...string) And { - userID := loggedUser(r.ctx).ID +func (r sqlRepository) annId(ctx context.Context, itemID ...string) And { + userID := loggedUser(ctx).ID return And{ Eq{annotationTable + ".user_id": userID}, Eq{annotationTable + ".item_type": r.tableName}, @@ -111,20 +112,20 @@ func (r sqlRepository) annId(itemID ...string) And { } } -func (r sqlRepository) annUpsert(values map[string]any, itemIDs ...string) error { - upd := Update(annotationTable).Where(r.annId(itemIDs...)) +func (r sqlRepository) annUpsert(ctx context.Context, values map[string]any, itemIDs ...string) error { + upd := Update(annotationTable).Where(r.annId(ctx, itemIDs...)) for f, v := range values { upd = upd.Set(f, v) } - c, err := r.executeSQL(upd) + c, err := r.executeSQL(ctx, upd) if c == 0 || errors.Is(err, sql.ErrNoRows) { - userID := loggedUser(r.ctx).ID + userID := loggedUser(ctx).ID for _, itemID := range itemIDs { values["user_id"] = userID values["item_type"] = r.tableName values["item_id"] = itemID ins := Insert(annotationTable).SetMap(values) - _, err = r.executeSQL(ins) + _, err = r.executeSQL(ctx, ins) if err != nil { return err } @@ -133,39 +134,39 @@ func (r sqlRepository) annUpsert(values map[string]any, itemIDs ...string) error return err } -func (r sqlRepository) SetStar(starred bool, ids ...string) error { +func (r sqlRepository) SetStar(ctx context.Context, starred bool, ids ...string) error { starredAt := time.Now() - return r.annUpsert(map[string]any{"starred": starred, "starred_at": starredAt}, ids...) + return r.annUpsert(ctx, map[string]any{"starred": starred, "starred_at": starredAt}, ids...) } -func (r sqlRepository) SetRating(rating int, itemID string) error { +func (r sqlRepository) SetRating(ctx context.Context, rating int, itemID string) error { ratedAt := time.Now() - err := r.annUpsert(map[string]any{"rating": rating, "rated_at": ratedAt}, itemID) + err := r.annUpsert(ctx, map[string]any{"rating": rating, "rated_at": ratedAt}, itemID) if err != nil { return err } - return r.updateAvgRating(itemID) + return r.updateAvgRating(ctx, itemID) } -func (r sqlRepository) updateAvgRating(itemID string) error { +func (r sqlRepository) updateAvgRating(ctx context.Context, itemID string) error { upd := Update(r.tableName). Where(Eq{"id": itemID}). Set("average_rating", Expr( "coalesce((select round(avg(rating), 2) from annotation where item_id = ? and item_type = ? and rating > 0), 0)", itemID, r.tableName, )) - _, err := r.executeSQL(upd) + _, err := r.executeSQL(ctx, upd) return err } -func (r sqlRepository) IncPlayCount(itemID string, ts time.Time) error { - upd := Update(annotationTable).Where(r.annId(itemID)). +func (r sqlRepository) IncPlayCount(ctx context.Context, itemID string, ts time.Time) error { + upd := Update(annotationTable).Where(r.annId(ctx, itemID)). Set("play_count", Expr("play_count+1")). Set("play_date", Expr("max(ifnull(play_date,''),?)", ts)) - c, err := r.executeSQL(upd) + c, err := r.executeSQL(ctx, upd) if c == 0 || errors.Is(err, sql.ErrNoRows) { - userID := loggedUser(r.ctx).ID + userID := loggedUser(ctx).ID values := map[string]any{} values["user_id"] = userID values["item_type"] = r.tableName @@ -173,7 +174,7 @@ func (r sqlRepository) IncPlayCount(itemID string, ts time.Time) error { values["play_count"] = 1 values["play_date"] = ts ins := Insert(annotationTable).SetMap(values) - _, err = r.executeSQL(ins) + _, err = r.executeSQL(ctx, ins) if err != nil { return err } @@ -181,28 +182,28 @@ func (r sqlRepository) IncPlayCount(itemID string, ts time.Time) error { return err } -func (r sqlRepository) ReassignAnnotation(prevID string, newID string) error { +func (r sqlRepository) ReassignAnnotation(ctx context.Context, prevID string, newID string) error { if prevID == newID || prevID == "" || newID == "" { return nil } // OR IGNORE keeps newID's own row where a user annotated both, instead of aborting the whole statement upd := Expr("update or ignore "+annotationTable+" set item_id = ? where item_type = ? and item_id = ?", newID, r.tableName, prevID) - if _, err := r.executeSQL(upd); err != nil { + if _, err := r.executeSQL(ctx, upd); err != nil { return err } // The moved rows change newID's rating population, so its cached average no longer matches - return r.updateAvgRating(newID) + return r.updateAvgRating(ctx, newID) } -func (r sqlRepository) cleanAnnotations() error { +func (r sqlRepository) cleanAnnotations(ctx context.Context) error { del := Delete(annotationTable).Where(Eq{"item_type": r.tableName}).Where("item_id not in (select id from " + r.tableName + ")") - c, err := r.executeSQL(del) + c, err := r.executeSQL(ctx, del) if err != nil { return fmt.Errorf("error cleaning up %s annotations: %w", r.tableName, err) } if c > 0 { - log.Debug(r.ctx, "Clean-up annotations", "table", r.tableName, "totalDeleted", c) + log.Debug(ctx, "Clean-up annotations", "table", r.tableName, "totalDeleted", c) } return nil } diff --git a/persistence/sql_annotations_test.go b/persistence/sql_annotations_test.go index 79ad3c152..ec28c610a 100644 --- a/persistence/sql_annotations_test.go +++ b/persistence/sql_annotations_test.go @@ -15,19 +15,20 @@ var _ = Describe("Annotation Filters", func() { var ( albumRepo *albumRepository albumWithoutAnnotation model.Album + ctx context.Context ) BeforeEach(func() { - ctx := request.WithUser(context.Background(), model.User{ID: "userid", UserName: "johndoe"}) - albumRepo = NewAlbumRepository(ctx, GetDBXBuilder()).(*albumRepository) + ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "userid", UserName: "johndoe"}) + albumRepo = NewAlbumRepository(GetDBXBuilder()).(*albumRepository) // Create album without any annotation (no star, no rating) albumWithoutAnnotation = model.Album{ID: "no-annotation-album", Name: "No Annotation", LibraryID: 1} - Expect(albumRepo.Put(&albumWithoutAnnotation)).To(Succeed()) + Expect(albumRepo.Put(ctx, &albumWithoutAnnotation)).To(Succeed()) }) AfterEach(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": albumWithoutAnnotation.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": albumWithoutAnnotation.ID})) }) Describe("ReassignAnnotation", func() { @@ -36,42 +37,42 @@ var _ = Describe("Annotation Filters", func() { BeforeEach(func() { prev = model.Album{ID: "reassign-prev", Name: "Prev", LibraryID: 1} next = model.Album{ID: "reassign-next", Name: "Next", LibraryID: 1} - Expect(albumRepo.Put(&prev)).To(Succeed()) - Expect(albumRepo.Put(&next)).To(Succeed()) + Expect(albumRepo.Put(ctx, &prev)).To(Succeed()) + Expect(albumRepo.Put(ctx, &next)).To(Succeed()) }) AfterEach(func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": []string{prev.ID, next.ID}})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": []string{prev.ID, next.ID}})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": []string{prev.ID, next.ID}})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": []string{prev.ID, next.ID}})) }) It("moves the annotation when the new item has none", func() { - Expect(albumRepo.SetRating(4, prev.ID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 4, prev.ID)).To(Succeed()) - Expect(albumRepo.ReassignAnnotation(prev.ID, next.ID)).To(Succeed()) + Expect(albumRepo.ReassignAnnotation(ctx, prev.ID, next.ID)).To(Succeed()) - got, err := albumRepo.Get(next.ID) + got, err := albumRepo.Get(ctx, next.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.Rating).To(Equal(4)) }) It("recomputes the new item's cached average rating", func() { - Expect(albumRepo.SetRating(4, prev.ID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 4, prev.ID)).To(Succeed()) - Expect(albumRepo.ReassignAnnotation(prev.ID, next.ID)).To(Succeed()) + Expect(albumRepo.ReassignAnnotation(ctx, prev.ID, next.ID)).To(Succeed()) - got, err := albumRepo.Get(next.ID) + got, err := albumRepo.Get(ctx, next.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.AverageRating).To(Equal(4.0)) }) It("keeps the new item's annotation when both exist", func() { - Expect(albumRepo.SetRating(4, prev.ID)).To(Succeed()) - Expect(albumRepo.SetRating(2, next.ID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 4, prev.ID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 2, next.ID)).To(Succeed()) - Expect(albumRepo.ReassignAnnotation(prev.ID, next.ID)).To(Succeed()) + Expect(albumRepo.ReassignAnnotation(ctx, prev.ID, next.ID)).To(Succeed()) - got, err := albumRepo.Get(next.ID) + got, err := albumRepo.Get(ctx, next.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.Rating).To(Equal(2)) }) @@ -100,7 +101,7 @@ var _ = Describe("Annotation Filters", func() { Describe("starredFilter", func() { It("false includes items without annotations", func() { - albums, err := albumRepo.GetAll(model.QueryOptions{ + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("starred")("starred", "false"), }) Expect(err).ToNot(HaveOccurred()) @@ -116,7 +117,7 @@ var _ = Describe("Annotation Filters", func() { }) It("true excludes items without annotations", func() { - albums, err := albumRepo.GetAll(model.QueryOptions{ + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("starred")("starred", "true"), }) Expect(err).ToNot(HaveOccurred()) @@ -129,7 +130,7 @@ var _ = Describe("Annotation Filters", func() { Describe("hasRatingFilter", func() { It("false includes items without annotations", func() { - albums, err := albumRepo.GetAll(model.QueryOptions{ + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("rating")("rating", "false"), }) Expect(err).ToNot(HaveOccurred()) @@ -145,7 +146,7 @@ var _ = Describe("Annotation Filters", func() { }) It("true excludes items without annotations", func() { - albums, err := albumRepo.GetAll(model.QueryOptions{ + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("rating")("rating", "true"), }) Expect(err).ToNot(HaveOccurred()) @@ -158,14 +159,14 @@ var _ = Describe("Annotation Filters", func() { It("true includes items with rating > 0", func() { // Create album with rating 1 ratedAlbum := model.Album{ID: "rated-album", Name: "Rated Album", LibraryID: 1} - Expect(albumRepo.Put(&ratedAlbum)).To(Succeed()) - Expect(albumRepo.SetRating(1, ratedAlbum.ID)).To(Succeed()) + Expect(albumRepo.Put(ctx, &ratedAlbum)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 1, ratedAlbum.ID)).To(Succeed()) defer func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": ratedAlbum.ID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": ratedAlbum.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": ratedAlbum.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": ratedAlbum.ID})) }() - albums, err := albumRepo.GetAll(model.QueryOptions{ + albums, err := albumRepo.GetAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("rating")("rating", "true"), }) Expect(err).ToNot(HaveOccurred()) @@ -182,11 +183,11 @@ var _ = Describe("Annotation Filters", func() { }) It("ignores invalid filter values (not strings)", func() { - res, err := albumRepo.ReadAll(rest.QueryOptions{ + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ Filters: map[string]any{"starred": 123}, }) Expect(err).ToNot(HaveOccurred()) - albums := res.(model.Albums) + albums := res var found bool for _, a := range albums { @@ -254,11 +255,11 @@ var _ = Describe("Annotation Filters", func() { Describe("CountAll annotation-join gating", func() { It("counts all items unfiltered (join dropped)", func() { - total, err := albumRepo.CountAll() + total, err := albumRepo.CountAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(total).To(BeNumerically(">=", int64(1))) - filtered, err := albumRepo.CountAll(model.QueryOptions{ + filtered, err := albumRepo.CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.id": albumWithoutAnnotation.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -267,16 +268,16 @@ var _ = Describe("Annotation Filters", func() { It("counts starred items correctly (named annotation filter keeps the join)", func() { starredAlbum := model.Album{ID: "counted-starred-album", Name: "Counted Starred", LibraryID: 1} - Expect(albumRepo.Put(&starredAlbum)).To(Succeed()) - Expect(albumRepo.SetStar(true, starredAlbum.ID)).To(Succeed()) + Expect(albumRepo.Put(ctx, &starredAlbum)).To(Succeed()) + Expect(albumRepo.SetStar(ctx, true, starredAlbum.ID)).To(Succeed()) defer func() { - _, _ = albumRepo.executeSQL(squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": starredAlbum.ID})) - _, _ = albumRepo.executeSQL(squirrel.Delete("album").Where(squirrel.Eq{"id": starredAlbum.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": starredAlbum.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": starredAlbum.ID})) }() // Exactly two albums are starred for this user: the one created above and // albumRadioactivity (id 103) from the seed data. - count, err := albumRepo.CountAll(model.QueryOptions{ + count, err := albumRepo.CountAll(ctx, model.QueryOptions{ Filters: annotationBoolFilter("starred")("starred", "true"), }) Expect(err).ToNot(HaveOccurred()) @@ -284,7 +285,7 @@ var _ = Describe("Annotation Filters", func() { }) It("counts via a raw annotation filter without a 'no such column' error", func() { - count, err := albumRepo.CountAll(model.QueryOptions{ + count, err := albumRepo.CountAll(ctx, model.QueryOptions{ Filters: squirrel.Expr("COALESCE(rating, 0) > 0"), }) Expect(err).ToNot(HaveOccurred()) diff --git a/persistence/sql_base_repository.go b/persistence/sql_base_repository.go index 026e42b03..be88156d8 100644 --- a/persistence/sql_base_repository.go +++ b/persistence/sql_base_repository.go @@ -29,7 +29,7 @@ import ( // When creating a new repository using this base, you must: // // - Embed this struct. -// - Set ctx and db fields. ctx should be the context passed to the constructor method, usually obtained from the request +// - Set the db field. // - Call registerModel with the model instance and any possible filters. // - If the model has a different table name than the default (lowercase of the model name), it should be set manually // using the tableName field. @@ -38,7 +38,6 @@ import ( // All fields in filters and sortMappings must be in snake_case. Only sorts and filters based on real field names or // defined in the mappings will be allowed. type sqlRepository struct { - ctx context.Context tableName string db dbx.Builder @@ -65,8 +64,8 @@ func loggedUser(ctx context.Context) *model.User { // // The predicate uses an unqualified user_id, so it only works on queries where that column is // unambiguous (no join introducing a second user_id). -func (r sqlRepository) ownerFilter() Sqlizer { - if usr := loggedUser(r.ctx); !usr.IsAdmin && usr.ID != invalidUserId { +func (r sqlRepository) ownerFilter(ctx context.Context) Sqlizer { + if usr := loggedUser(ctx); !usr.IsAdmin && usr.ID != invalidUserId { return Eq{"user_id": usr.ID} } return nil @@ -75,12 +74,12 @@ func (r sqlRepository) ownerFilter() Sqlizer { // addRestriction combines an optional caller predicate with the ownership filter, producing the // WHERE clause for owner-scoped reads. For admins and headless contexts ownerFilter() is nil and // only the caller's predicate (if any) remains. -func (r sqlRepository) addRestriction(sql ...Sqlizer) Sqlizer { +func (r sqlRepository) addRestriction(ctx context.Context, sql ...Sqlizer) Sqlizer { s := And{} if len(sql) > 0 { s = append(s, sql[0]) } - if owner := r.ownerFilter(); owner != nil { + if owner := r.ownerFilter(ctx); owner != nil { s = append(s, owner) } return s @@ -121,10 +120,10 @@ func (r *sqlRepository) setSortMappings(mappings map[string]string, tableName .. r.sortMappings = mappings } -func (r sqlRepository) newSelect(options ...model.QueryOptions) SelectBuilder { +func (r sqlRepository) newSelect(ctx context.Context, options ...model.QueryOptions) SelectBuilder { sq := Select().From(r.tableName) if len(options) > 0 { - r.resetSeededRandom(options) + r.resetSeededRandom(ctx, options) sq = r.applyOptions(sq, options...) sq = r.applyFilters(sq, options...) } @@ -243,8 +242,8 @@ func libraryIdFilter(_ string, value any) Sqlizer { // applyLibraryFilter adds library filtering to queries for tables that have a library_id column // This ensures users only see content from libraries they have access to -func (r sqlRepository) applyLibraryFilter(sq SelectBuilder, tableName ...string) SelectBuilder { - user := loggedUser(r.ctx) +func (r sqlRepository) applyLibraryFilter(ctx context.Context, sq SelectBuilder, tableName ...string) SelectBuilder { + user := loggedUser(ctx) // If the user is an admin, or the user ID is invalid (e.g., when no user is logged in), skip the library filter if user.IsAdmin || user.ID == invalidUserId { @@ -253,7 +252,7 @@ func (r sqlRepository) applyLibraryFilter(sq SelectBuilder, tableName ...string) // A non-admin granted every library sees everything the subquery would return, so applying it is // pure overhead. Skip it in that case (same fast path admins get). - if visible, err := r.visibleLibraryIDs(); err == nil && r.userSeesAllLibraries(visible) { + if visible, err := r.visibleLibraryIDs(ctx); err == nil && r.userSeesAllLibraries(ctx, visible) { return sq } @@ -270,65 +269,66 @@ func (r sqlRepository) applyLibraryFilter(sq SelectBuilder, tableName ...string) // userSeesAllLibraries reports whether the visible set already covers every library, so a // library filter would exclude nothing. -func (r sqlRepository) userSeesAllLibraries(visible []int) bool { - user := loggedUser(r.ctx) +func (r sqlRepository) userSeesAllLibraries(ctx context.Context, visible []int) bool { + user := loggedUser(ctx) if user.IsAdmin || user.ID == invalidUserId { return true // visible is the whole library table } - total, err := NewLibraryRepository(r.ctx, r.db).CountAll() - if err != nil || total == 0 { + var res struct{ Count int64 } + err := r.queryOne(ctx, Select("count(*) as count").From("library"), &res) + if err != nil || res.Count == 0 { return false } - return int64(len(visible)) == total + return int64(len(visible)) == res.Count } // visibleLibraryIDs returns the libraries the current user can see: all libraries for admin and // headless processes, otherwise the user's granted libraries. -func (r sqlRepository) visibleLibraryIDs() ([]int, error) { - user := loggedUser(r.ctx) +func (r sqlRepository) visibleLibraryIDs(ctx context.Context) ([]int, error) { + user := loggedUser(ctx) if user.IsAdmin || user.ID == invalidUserId { var ids []int - err := r.queryAllSlice(Select("id").From("library"), &ids) + err := r.queryAllSlice(ctx, Select("id").From("library"), &ids) return ids, err } return slice.Map(user.Libraries, func(lib model.Library) int { return lib.ID }), nil } -func (r sqlRepository) seedKey() string { +func (r sqlRepository) seedKey(ctx context.Context) string { // Seed keys must be all lowercase, or else SQLite3 will encode it, making it not match the seed // used in the query. Hashing the user ID and converting it to a hex string will do the trick - userIDHash := xxh3.Hash([]byte(loggedUser(r.ctx).ID)) + userIDHash := xxh3.Hash([]byte(loggedUser(ctx).ID)) return fmt.Sprintf("%s|%016x", r.tableName, userIDHash) } -func (r sqlRepository) resetSeededRandom(options []model.QueryOptions) { +func (r sqlRepository) resetSeededRandom(ctx context.Context, options []model.QueryOptions) { if len(options) == 0 || options[0].Sort != "random" { return } // CAST: playlist_tracks.id is an INTEGER (unlike other tables' TEXT ids); passing it to // SEEDEDRAND's string param uncast silently drops every row (go-sqlite3 binding gotcha). - options[0].Sort = fmt.Sprintf("SEEDEDRAND('%s', CAST(%s.id AS TEXT))", r.seedKey(), r.tableName) + options[0].Sort = fmt.Sprintf("SEEDEDRAND('%s', CAST(%s.id AS TEXT))", r.seedKey(ctx), r.tableName) if options[0].Seed != "" { - hasher.SetSeed(r.seedKey(), options[0].Seed) + hasher.SetSeed(r.seedKey(ctx), options[0].Seed) return } if options[0].Offset == 0 { - hasher.Reseed(r.seedKey()) + hasher.Reseed(r.seedKey(ctx)) } } -func (r sqlRepository) executeSQL(sq Sqlizer) (int64, error) { +func (r sqlRepository) executeSQL(ctx context.Context, sq Sqlizer) (int64, error) { query, args, err := r.toSQL(sq) if err != nil { return 0, err } start := time.Now() var c int64 - res, err := r.db.NewQuery(query).Bind(args).WithContext(r.ctx).Execute() + res, err := r.db.NewQuery(query).Bind(args).WithContext(ctx).Execute() if res != nil { c, _ = res.RowsAffected() } - r.logSQL(query, args, err, c, start) + r.logSQL(ctx, query, args, err, c, start) if err != nil { if err.Error() != "LastInsertId is not supported by this driver" { return 0, err @@ -356,18 +356,18 @@ func (r sqlRepository) toSQL(sq Sqlizer) (string, dbx.Params, error) { return result, params, nil } -func (r sqlRepository) queryOne(sq Sqlizer, response any) error { +func (r sqlRepository) queryOne(ctx context.Context, sq Sqlizer, response any) error { query, args, err := r.toSQL(sq) if err != nil { return err } start := time.Now() - err = r.db.NewQuery(query).Bind(args).WithContext(r.ctx).One(response) + err = r.db.NewQuery(query).Bind(args).WithContext(ctx).One(response) if errors.Is(err, sql.ErrNoRows) { - r.logSQL(query, args, nil, 0, start) + r.logSQL(ctx, query, args, nil, 0, start) return model.ErrNotFound } - r.logSQL(query, args, err, 1, start) + r.logSQL(ctx, query, args, err, 1, start) return err } @@ -392,7 +392,7 @@ func wrapCursor[D, T any](cursor iter.Seq2[D, error], toModel func(D) *T) iter.S // queryWithStableResults is a helper function to execute a query and return an iterator that will yield its results // from a cursor, guaranteeing that the results will be stable, even if the underlying data changes. -func queryWithStableResults[T any](r sqlRepository, sq SelectBuilder, options ...model.QueryOptions) (iter.Seq2[T, error], error) { +func queryWithStableResults[T any](ctx context.Context, r sqlRepository, sq SelectBuilder, options ...model.QueryOptions) (iter.Seq2[T, error], error) { if len(options) > 0 && options[0].Offset > 0 { sq = r.optimizePagination(sq, options[0]) } @@ -401,8 +401,8 @@ func queryWithStableResults[T any](r sqlRepository, sq SelectBuilder, options .. return nil, err } start := time.Now() - rows, err := r.db.NewQuery(query).Bind(args).WithContext(r.ctx).Rows() - r.logSQL(query, args, err, -1, start) + rows, err := r.db.NewQuery(query).Bind(args).WithContext(ctx).Rows() + r.logSQL(ctx, query, args, err, -1, start) if err != nil { return nil, err } @@ -422,7 +422,7 @@ func queryWithStableResults[T any](r sqlRepository, sq SelectBuilder, options .. }, nil } -func (r sqlRepository) queryAll(sq SelectBuilder, response any, options ...model.QueryOptions) error { +func (r sqlRepository) queryAll(ctx context.Context, sq SelectBuilder, response any, options ...model.QueryOptions) error { if len(options) > 0 && options[0].Offset > 0 { sq = r.optimizePagination(sq, options[0]) } @@ -431,28 +431,28 @@ func (r sqlRepository) queryAll(sq SelectBuilder, response any, options ...model return err } start := time.Now() - err = r.db.NewQuery(query).Bind(args).WithContext(r.ctx).All(response) + err = r.db.NewQuery(query).Bind(args).WithContext(ctx).All(response) if errors.Is(err, sql.ErrNoRows) { - r.logSQL(query, args, nil, -1, start) + r.logSQL(ctx, query, args, nil, -1, start) return model.ErrNotFound } - r.logSQL(query, args, err, int64(reflect.ValueOf(response).Elem().Len()), start) + r.logSQL(ctx, query, args, err, int64(reflect.ValueOf(response).Elem().Len()), start) return err } // queryAllSlice is a helper function to query a single column and return the result in a slice -func (r sqlRepository) queryAllSlice(sq SelectBuilder, response any) error { +func (r sqlRepository) queryAllSlice(ctx context.Context, sq SelectBuilder, response any) error { query, args, err := r.toSQL(sq) if err != nil { return err } start := time.Now() - err = r.db.NewQuery(query).Bind(args).WithContext(r.ctx).Column(response) + err = r.db.NewQuery(query).Bind(args).WithContext(ctx).Column(response) if errors.Is(err, sql.ErrNoRows) { - r.logSQL(query, args, nil, -1, start) + r.logSQL(ctx, query, args, nil, -1, start) return model.ErrNotFound } - r.logSQL(query, args, err, int64(reflect.ValueOf(response).Elem().Len()), start) + r.logSQL(ctx, query, args, err, int64(reflect.ValueOf(response).Elem().Len()), start) return err } @@ -469,10 +469,10 @@ func (r sqlRepository) optimizePagination(sq SelectBuilder, options model.QueryO return sq } -func (r sqlRepository) exists(cond Sqlizer) (bool, error) { +func (r sqlRepository) exists(ctx context.Context, cond Sqlizer) (bool, error) { existsQuery := Select("count(*) as exist").From(r.tableName).Where(cond) var res struct{ Exist int64 } - err := r.queryOne(existsQuery, &res) + err := r.queryOne(ctx, existsQuery, &res) return res.Exist > 0, err } @@ -487,20 +487,20 @@ func (r sqlRepository) exists(cond Sqlizer) (bool, error) { // another user it returns rest.ErrPermissionDenied, otherwise rest.ErrNotFound. The write itself is // still atomic; the extra lookup happens only on the failure path (count == 0), where no write // occurred, so there is no TOCTOU on the update. -func (r sqlRepository) updateOwned(id string, m any, colsToUpdate ...string) error { +func (r sqlRepository) updateOwned(ctx context.Context, id string, m any, colsToUpdate ...string) error { values, err := toSQLArgs(m) if err != nil { return fmt.Errorf("error preparing values to write to DB: %w", err) } updateValues := filterUpdateValues(values, id, colsToUpdate...) delete(updateValues, "user_id") // ownership is immutable on update - update := Update(r.tableName).Where(r.addRestriction(Eq{"id": id})).SetMap(updateValues) - count, err := r.executeSQL(update) + update := Update(r.tableName).Where(r.addRestriction(ctx, Eq{"id": id})).SetMap(updateValues) + count, err := r.executeSQL(ctx, update) if err != nil { return err } if count == 0 { - return r.classifyOwnedWriteMiss(id) + return r.classifyOwnedWriteMiss(ctx, id) } return nil } @@ -510,13 +510,22 @@ func (r sqlRepository) updateOwned(id string, m any, colsToUpdate ...string) err // ownership predicate is part of the DELETE's WHERE clause, so a row owned by another user simply // does not match and is left untouched. The failure path mirrors updateOwned (see // classifyOwnedWriteMiss), so there is no TOCTOU on the delete. -func (r sqlRepository) deleteOwned(id string) error { - count, err := r.executeSQL(Delete(r.tableName).Where(r.addRestriction(Eq{"id": id}))) +func (r sqlRepository) deleteOwned(ctx context.Context, id string) error { + count, err := r.executeSQL(ctx, Delete(r.tableName).Where(r.addRestriction(ctx, Eq{"id": id}))) if err != nil { return err } if count == 0 { - return r.classifyOwnedWriteMiss(id) + return r.classifyOwnedWriteMiss(ctx, id) + } + return nil +} + +func (r sqlRepository) deleteOwnedAll(ctx context.Context, ids ...string) error { + for _, id := range ids { + if err := r.deleteOwned(ctx, id); err != nil { + return err + } } return nil } @@ -524,8 +533,8 @@ func (r sqlRepository) deleteOwned(id string) error { // classifyOwnedWriteMiss explains why an ownership-filtered write (updateOwned/deleteOwned) matched // no row: rest.ErrPermissionDenied if the row exists but is owned by another user, otherwise // rest.ErrNotFound. It runs only on the failure path (count == 0), where no write occurred. -func (r sqlRepository) classifyOwnedWriteMiss(id string) error { - exists, err := r.exists(Eq{"id": id}) +func (r sqlRepository) classifyOwnedWriteMiss(ctx context.Context, id string) error { + exists, err := r.exists(ctx, Eq{"id": id}) if err != nil { return err } @@ -535,7 +544,7 @@ func (r sqlRepository) classifyOwnedWriteMiss(id string) error { return rest.ErrNotFound } -func (r sqlRepository) count(countQuery SelectBuilder, options ...model.QueryOptions) (int64, error) { +func (r sqlRepository) count(ctx context.Context, countQuery SelectBuilder, options ...model.QueryOptions) (int64, error) { countQuery = countQuery. RemoveColumns().Columns("count(distinct " + r.tableName + ".id) as count"). RemoveOffset().RemoveLimit(). @@ -543,22 +552,22 @@ func (r sqlRepository) count(countQuery SelectBuilder, options ...model.QueryOpt From(r.tableName) countQuery = r.applyFilters(countQuery, options...) var res struct{ Count int64 } - err := r.queryOne(countQuery, &res) + err := r.queryOne(ctx, countQuery, &res) return res.Count, err } -func (r sqlRepository) putByMatch(filter Sqlizer, id string, m any, colsToUpdate ...string) (string, error) { +func (r sqlRepository) putByMatch(ctx context.Context, filter Sqlizer, id string, m any, colsToUpdate ...string) (string, error) { if id != "" { - return r.put(id, m, colsToUpdate...) + return r.put(ctx, id, m, colsToUpdate...) } - existsQuery := r.newSelect().Columns("id").From(r.tableName).Where(filter) + existsQuery := r.newSelect(ctx).Columns("id").From(r.tableName).Where(filter) var res struct{ ID string } - err := r.queryOne(existsQuery, &res) + err := r.queryOne(ctx, existsQuery, &res) if err != nil && !errors.Is(err, model.ErrNotFound) { return "", err } - return r.put(res.ID, m, colsToUpdate...) + return r.put(ctx, res.ID, m, colsToUpdate...) } // selectUpdateColumns keeps only the requested colsToUpdate (or all columns when none are @@ -588,7 +597,7 @@ func filterUpdateValues(values map[string]any, id string, colsToUpdate ...string return updateValues } -func (r sqlRepository) put(id string, m any, colsToUpdate ...string) (newId string, err error) { +func (r sqlRepository) put(ctx context.Context, id string, m any, colsToUpdate ...string) (newId string, err error) { values, err := toSQLArgs(m) if err != nil { return "", fmt.Errorf("error preparing values to write to DB: %w", err) @@ -596,7 +605,7 @@ func (r sqlRepository) put(id string, m any, colsToUpdate ...string) (newId stri // If there's an ID, try to update first if id != "" { update := Update(r.tableName).Where(Eq{"id": id}).SetMap(filterUpdateValues(values, id, colsToUpdate...)) - count, err := r.executeSQL(update) + count, err := r.executeSQL(ctx, update) if err != nil { return "", err } @@ -610,18 +619,18 @@ func (r sqlRepository) put(id string, m any, colsToUpdate ...string) (newId stri values["id"] = id } insert := Insert(r.tableName).SetMap(values) - _, err = r.executeSQL(insert) + _, err = r.executeSQL(ctx, insert) return id, err } -func (r sqlRepository) delete(cond Sqlizer) error { - _, err := r.executeSQL(Delete(r.tableName).Where(cond)) +func (r sqlRepository) delete(ctx context.Context, cond Sqlizer) error { + _, err := r.executeSQL(ctx, Delete(r.tableName).Where(cond)) return err } // deleteByID is for single-item deletes that must report a missing row; delete succeeds silently. -func (r sqlRepository) deleteByID(id string) error { - count, err := r.executeSQL(Delete(r.tableName).Where(Eq{"id": id})) +func (r sqlRepository) deleteByID(ctx context.Context, id string) error { + count, err := r.executeSQL(ctx, Delete(r.tableName).Where(Eq{"id": id})) if err != nil { return err } @@ -631,9 +640,9 @@ func (r sqlRepository) deleteByID(id string) error { return nil } -func (r sqlRepository) logSQL(sql string, args dbx.Params, err error, rowsAffected int64, start time.Time) { +func (r sqlRepository) logSQL(ctx context.Context, sql string, args dbx.Params, err error, rowsAffected int64, start time.Time) { elapsed := time.Since(start) - fields := []any{r.ctx, "SQL: `" + sql + "`", "args", args, "rowsAffected", rowsAffected, "elapsedTime", elapsed} + fields := []any{ctx, "SQL: `" + sql + "`", "args", args, "rowsAffected", rowsAffected, "elapsedTime", elapsed} if err == nil || errors.Is(err, context.Canceled) { log.Trace(append(fields, err)...) return @@ -643,7 +652,7 @@ func (r sqlRepository) logSQL(sql string, args dbx.Params, err error, rowsAffect if code, extended, ok := db.ErrorCodes(err); ok { fields = append(fields, "sqliteCode", code, "sqliteExtended", extended) } - if db.IsBusy(err) && hasBusyRetry(r.ctx) { + if db.IsBusy(err) && hasBusyRetry(ctx) { log.Warn(append(fields, err)...) return } diff --git a/persistence/sql_base_repository_test.go b/persistence/sql_base_repository_test.go index 0f76eb6ab..33a8140f8 100644 --- a/persistence/sql_base_repository_test.go +++ b/persistence/sql_base_repository_test.go @@ -13,8 +13,9 @@ import ( var _ = Describe("sqlRepository", func() { var r sqlRepository + var ctx context.Context BeforeEach(func() { - r.ctx = request.WithUser(context.Background(), model.User{ID: "user-id"}) + ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "user-id"}) r.tableName = "table" }) @@ -88,19 +89,19 @@ var _ = Describe("sqlRepository", func() { When("sanitizing sort", func() { It("returns empty if the sort key is not found in the model nor in the mappings", func() { - sort, _ := r.sanitizeSort("unknown", "") + sort, _ := r.sanitizeSort(ctx, "unknown", "") Expect(sort).To(BeEmpty()) }) // Validation only: buildSortOrder resolves the mapping, so mapping here too would hand // sortMapping its own output and re-map values whose parts are themselves keys. It("accepts a known sort key without resolving it", func() { - sort, _ := r.sanitizeSort("sort1", "") + sort, _ := r.sanitizeSort(ctx, "sort1", "") Expect(sort).To(Equal("sort1")) }) It("is case insensitive", func() { - sort, _ := r.sanitizeSort("Sort1", "") + sort, _ := r.sanitizeSort(ctx, "Sort1", "") Expect(sort).To(Equal("sort1")) }) @@ -112,38 +113,38 @@ var _ = Describe("sqlRepository", func() { // must survive the round trip through sanitizeSort and buildSortOrder unduplicated. It("does not re-map a value whose parts are also keys", func() { r.sortMappings = map[string]string{"rating": "rating", "rated_at": "rating, rated_at"} - sort, _ := r.sanitizeSort("rated_at", "") + sort, _ := r.sanitizeSort(ctx, "rated_at", "") Expect(r.buildSortOrder(sort, "asc")).To(Equal("rating asc, rated_at asc")) }) It("returns the field if it is a valid field", func() { - sort, _ := r.sanitizeSort("field", "") + sort, _ := r.sanitizeSort(ctx, "field", "") Expect(sort).To(Equal("field")) }) It("is case insensitive for fields", func() { - sort, _ := r.sanitizeSort("FIELD", "") + sort, _ := r.sanitizeSort(ctx, "FIELD", "") Expect(sort).To(Equal("field")) }) }) When("sanitizing order", func() { It("returns 'asc' if order is empty", func() { - _, order := r.sanitizeSort("", "") + _, order := r.sanitizeSort(ctx, "", "") Expect(order).To(Equal("")) }) It("returns 'asc' if order is 'asc'", func() { - _, order := r.sanitizeSort("", "ASC") + _, order := r.sanitizeSort(ctx, "", "ASC") Expect(order).To(Equal("asc")) }) It("returns 'desc' if order is 'desc'", func() { - _, order := r.sanitizeSort("", "desc") + _, order := r.sanitizeSort(ctx, "", "desc") Expect(order).To(Equal("desc")) }) It("returns 'asc' if order is unknown", func() { - _, order := r.sanitizeSort("", "something") + _, order := r.sanitizeSort(ctx, "", "something") Expect(order).To(Equal("asc")) }) }) @@ -248,31 +249,31 @@ var _ = Describe("sqlRepository", func() { Describe("resetSeededRandom", func() { var id string BeforeEach(func() { - id = r.seedKey() + id = r.seedKey(ctx) hasher.SetSeed(id, "") }) It("does not reset seed if sort is not random", func() { var options []model.QueryOptions - r.resetSeededRandom(options) + r.resetSeededRandom(ctx, options) Expect(hasher.CurrentSeed(id)).To(BeEmpty()) }) It("resets seed if sort is random", func() { options := []model.QueryOptions{{Sort: "random"}} - r.resetSeededRandom(options) + r.resetSeededRandom(ctx, options) Expect(hasher.CurrentSeed(id)).NotTo(BeEmpty()) }) It("resets seed if sort is random and seed is provided", func() { options := []model.QueryOptions{{Sort: "random", Seed: "seed"}} - r.resetSeededRandom(options) + r.resetSeededRandom(ctx, options) Expect(hasher.CurrentSeed(id)).To(Equal("seed")) }) It("keeps seed when paginating", func() { options := []model.QueryOptions{{Sort: "random", Seed: "seed", Offset: 0}} - r.resetSeededRandom(options) + r.resetSeededRandom(ctx, options) Expect(hasher.CurrentSeed(id)).To(Equal("seed")) options = []model.QueryOptions{{Sort: "random", Offset: 1}} - r.resetSeededRandom(options) + r.resetSeededRandom(ctx, options) Expect(hasher.CurrentSeed(id)).To(Equal("seed")) }) }) @@ -298,11 +299,11 @@ var _ = Describe("sqlRepository", func() { Context("Admin User", func() { BeforeEach(func() { - r.ctx = request.WithUser(context.Background(), model.User{ID: "admin", IsAdmin: true}) + ctx = request.WithUser(ctx, model.User{ID: "admin", IsAdmin: true}) }) It("should not apply library filter for admin users", func() { - result := r.applyLibraryFilter(sq) + result := r.applyLibraryFilter(ctx, sq) sql, _, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(Equal("SELECT * FROM test_table")) @@ -312,13 +313,13 @@ var _ = Describe("sqlRepository", func() { Context("Regular User with a subset of libraries", func() { BeforeEach(func() { // Strict subset: granted lib 1, DB has libs 1 and 2, so the filter must apply. - r.ctx = request.WithUser(context.Background(), model.User{ + ctx = request.WithUser(ctx, model.User{ ID: "user123", IsAdmin: false, Libraries: model.Libraries{{ID: 1}}, }) }) It("should apply library filter for regular users", func() { - result := r.applyLibraryFilter(sq) + result := r.applyLibraryFilter(ctx, sq) sql, args, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(ContainSubstring("IN (SELECT ul.library_id FROM user_library ul WHERE ul.user_id = ?)")) @@ -326,7 +327,7 @@ var _ = Describe("sqlRepository", func() { }) It("should use custom table name when provided", func() { - result := r.applyLibraryFilter(sq, "custom_table") + result := r.applyLibraryFilter(ctx, sq, "custom_table") sql, args, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(ContainSubstring("custom_table.library_id IN")) @@ -336,11 +337,11 @@ var _ = Describe("sqlRepository", func() { Context("Regular User with no libraries", func() { BeforeEach(func() { - r.ctx = request.WithUser(context.Background(), model.User{ID: "empty", IsAdmin: false}) + ctx = request.WithUser(ctx, model.User{ID: "empty", IsAdmin: false}) }) It("should apply the library filter (never skip on empty)", func() { - result := r.applyLibraryFilter(sq) + result := r.applyLibraryFilter(ctx, sq) sql, _, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(ContainSubstring("IN (SELECT ul.library_id FROM user_library ul WHERE ul.user_id = ?)")) @@ -359,20 +360,20 @@ var _ = Describe("sqlRepository", func() { for _, id := range ids { libs = append(libs, model.Library{ID: id}) } - r.ctx = request.WithUser(context.Background(), model.User{ + ctx = request.WithUser(ctx, model.User{ ID: "alllibs", IsAdmin: false, Libraries: libs, }) }) It("should not apply the library filter (subquery would filter nothing)", func() { - result := r.applyLibraryFilter(sq) + result := r.applyLibraryFilter(ctx, sq) sql, _, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(Equal("SELECT * FROM test_table")) }) It("should not apply the filter even with a custom table name", func() { - result := r.applyLibraryFilter(sq, "custom_table") + result := r.applyLibraryFilter(ctx, sq, "custom_table") sql, _, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(Equal("SELECT * FROM test_table")) @@ -381,18 +382,18 @@ var _ = Describe("sqlRepository", func() { Context("Headless Process (No User Context)", func() { BeforeEach(func() { - r.ctx = context.Background() // No user context + ctx = GinkgoT().Context() // No user context }) It("should not apply library filter for headless processes", func() { - result := r.applyLibraryFilter(sq) + result := r.applyLibraryFilter(ctx, sq) sql, _, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(Equal("SELECT * FROM test_table")) }) It("should not apply library filter even with custom table name", func() { - result := r.applyLibraryFilter(sq, "custom_table") + result := r.applyLibraryFilter(ctx, sq, "custom_table") sql, _, err := result.ToSql() Expect(err).ToNot(HaveOccurred()) Expect(sql).To(Equal("SELECT * FROM test_table")) diff --git a/persistence/sql_bookmarks.go b/persistence/sql_bookmarks.go index cff57dc9d..ddb63019d 100644 --- a/persistence/sql_bookmarks.go +++ b/persistence/sql_bookmarks.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "database/sql" "errors" "fmt" @@ -14,8 +15,8 @@ import ( const bookmarkTable = "bookmark" -func (r sqlRepository) withBookmark(query SelectBuilder, idField string) SelectBuilder { - userID := loggedUser(r.ctx).ID +func (r sqlRepository) withBookmark(ctx context.Context, query SelectBuilder, idField string) SelectBuilder { + userID := loggedUser(ctx).ID if userID == invalidUserId { return query } @@ -26,17 +27,17 @@ func (r sqlRepository) withBookmark(query SelectBuilder, idField string) SelectB Columns("coalesce(position, 0) as bookmark_position") } -func (r sqlRepository) bmkID(itemID ...string) And { +func (r sqlRepository) bmkID(ctx context.Context, itemID ...string) And { return And{ - Eq{bookmarkTable + ".user_id": loggedUser(r.ctx).ID}, + Eq{bookmarkTable + ".user_id": loggedUser(ctx).ID}, Eq{bookmarkTable + ".item_type": r.tableName}, Eq{bookmarkTable + ".item_id": itemID}, } } -func (r sqlRepository) bmkUpsert(itemID, comment string, position int64) error { - client, _ := request.ClientFrom(r.ctx) - user, _ := request.UserFrom(r.ctx) +func (r sqlRepository) bmkUpsert(ctx context.Context, itemID, comment string, position int64) error { + client, _ := request.ClientFrom(ctx) + user, _ := request.UserFrom(ctx) values := map[string]any{ "comment": comment, "position": position, @@ -44,10 +45,10 @@ func (r sqlRepository) bmkUpsert(itemID, comment string, position int64) error { "changed_by": client, } - upd := Update(bookmarkTable).Where(r.bmkID(itemID)).SetMap(values) - c, err := r.executeSQL(upd) + upd := Update(bookmarkTable).Where(r.bmkID(ctx, itemID)).SetMap(values) + c, err := r.executeSQL(ctx, upd) if err == nil { - log.Debug(r.ctx, "Updated bookmark", "id", itemID, "user", user.UserName, "position", position, "comment", comment) + log.Debug(ctx, "Updated bookmark", "id", itemID, "user", user.UserName, "position", position, "comment", comment) } if c == 0 || errors.Is(err, sql.ErrNoRows) { values["user_id"] = user.ID @@ -56,31 +57,31 @@ func (r sqlRepository) bmkUpsert(itemID, comment string, position int64) error { values["created_at"] = time.Now() values["updated_at"] = time.Now() ins := Insert(bookmarkTable).SetMap(values) - _, err = r.executeSQL(ins) + _, err = r.executeSQL(ctx, ins) if err != nil { return err } - log.Debug(r.ctx, "Added bookmark", "id", itemID, "user", user.UserName, "position", position, "comment", comment) + log.Debug(ctx, "Added bookmark", "id", itemID, "user", user.UserName, "position", position, "comment", comment) } return err } -func (r sqlRepository) AddBookmark(id, comment string, position int64) error { - user, _ := request.UserFrom(r.ctx) - err := r.bmkUpsert(id, comment, position) +func (r sqlRepository) AddBookmark(ctx context.Context, id, comment string, position int64) error { + user, _ := request.UserFrom(ctx) + err := r.bmkUpsert(ctx, id, comment, position) if err != nil { - log.Error(r.ctx, "Error adding bookmark", "id", id, "user", user.UserName, "position", position, "comment", comment) + log.Error(ctx, "Error adding bookmark", "id", id, "user", user.UserName, "position", position, "comment", comment) } return err } -func (r sqlRepository) DeleteBookmark(id string) error { - user, _ := request.UserFrom(r.ctx) - del := Delete(bookmarkTable).Where(r.bmkID(id)) - _, err := r.executeSQL(del) +func (r sqlRepository) DeleteBookmark(ctx context.Context, id string) error { + user, _ := request.UserFrom(ctx) + del := Delete(bookmarkTable).Where(r.bmkID(ctx, id)) + _, err := r.executeSQL(ctx, del) if err != nil { - log.Error(r.ctx, "Error removing bookmark", "id", id, "user", user.UserName) + log.Error(ctx, "Error removing bookmark", "id", id, "user", user.UserName) } return err } @@ -96,18 +97,18 @@ type bookmark struct { UpdatedAt time.Time `json:"updatedAt"` } -func (r sqlRepository) GetBookmarks() (model.Bookmarks, error) { - user, _ := request.UserFrom(r.ctx) +func (r sqlRepository) GetBookmarks(ctx context.Context) (model.Bookmarks, error) { + user, _ := request.UserFrom(ctx) idField := r.tableName + ".id" - sq := r.newSelect().Columns(r.tableName + ".*") - sq = r.withAnnotation(sq, idField) - sq = r.withBookmark(sq, idField).Where(NotEq{bookmarkTable + ".item_id": nil}) - sq = r.applyLibraryFilter(sq) + sq := r.newSelect(ctx).Columns(r.tableName + ".*") + sq = r.withAnnotation(ctx, sq, idField) + sq = r.withBookmark(ctx, sq, idField).Where(NotEq{bookmarkTable + ".item_id": nil}) + sq = r.applyLibraryFilter(ctx, sq) var mfs dbMediaFiles // TODO Decouple from media_file - err := r.queryAll(sq, &mfs) + err := r.queryAll(ctx, sq, &mfs) if err != nil { - log.Error(r.ctx, "Error getting mediafiles with bookmarks", "user", user.UserName, err) + log.Error(ctx, "Error getting mediafiles with bookmarks", "user", user.UserName, err) return nil, err } @@ -118,18 +119,18 @@ func (r sqlRepository) GetBookmarks() (model.Bookmarks, error) { mfMap[mf.ID] = i } - sq = Select("*").From(bookmarkTable).Where(r.bmkID(ids...)) + sq = Select("*").From(bookmarkTable).Where(r.bmkID(ctx, ids...)) var bmks []bookmark - err = r.queryAll(sq, &bmks) + err = r.queryAll(ctx, sq, &bmks) if err != nil { - log.Error(r.ctx, "Error getting bookmarks", "user", user.UserName, "ids", ids, err) + log.Error(ctx, "Error getting bookmarks", "user", user.UserName, "ids", ids, err) return nil, err } resp := make(model.Bookmarks, len(bmks)) for i, bmk := range bmks { if itemIdx, ok := mfMap[bmk.ItemID]; !ok { - log.Debug(r.ctx, "Invalid bookmark", "id", bmk.ItemID, "user", user.UserName) + log.Debug(ctx, "Invalid bookmark", "id", bmk.ItemID, "user", user.UserName) continue } else { resp[i] = model.Bookmark{ @@ -145,21 +146,21 @@ func (r sqlRepository) GetBookmarks() (model.Bookmarks, error) { return resp, nil } -func (r sqlRepository) reassignBookmark(prevID, newID string) error { +func (r sqlRepository) reassignBookmark(ctx context.Context, prevID, newID string) error { upd := Expr("update or ignore "+bookmarkTable+" set item_id = ? where item_type = ? and item_id = ?", newID, r.tableName, prevID) - _, err := r.executeSQL(upd) + _, err := r.executeSQL(ctx, upd) return err } -func (r sqlRepository) cleanBookmarks() error { +func (r sqlRepository) cleanBookmarks(ctx context.Context) error { del := Delete(bookmarkTable).Where(Eq{"item_type": r.tableName}).Where("item_id not in (select id from " + r.tableName + ")") - c, err := r.executeSQL(del) + c, err := r.executeSQL(ctx, del) if err != nil { return fmt.Errorf("error cleaning up %s bookmarks: %w", r.tableName, err) } if c > 0 { - log.Debug(r.ctx, "Clean-up bookmarks", "totalDeleted", c, "itemType", r.tableName) + log.Debug(ctx, "Clean-up bookmarks", "totalDeleted", c, "itemType", r.tableName) } return nil } diff --git a/persistence/sql_bookmarks_test.go b/persistence/sql_bookmarks_test.go index ae01a0e35..b50a35d5d 100644 --- a/persistence/sql_bookmarks_test.go +++ b/persistence/sql_bookmarks_test.go @@ -12,23 +12,23 @@ import ( var _ = Describe("sqlBookmarks", func() { var mr model.MediaFileRepository + var ctx context.Context BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, model.User{ID: "userid"}) - mr = NewMediaFileRepository(ctx, GetDBXBuilder()) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), model.User{ID: "userid"}) + mr = NewMediaFileRepository(GetDBXBuilder()) }) Describe("Bookmarks", func() { It("returns an empty collection if there are no bookmarks", func() { - Expect(mr.GetBookmarks()).To(BeEmpty()) + Expect(mr.GetBookmarks(ctx)).To(BeEmpty()) }) It("saves and overrides bookmarks", func() { By("Saving the bookmark") - Expect(mr.AddBookmark(songAntenna.ID, "this is a comment", 123)).To(BeNil()) + Expect(mr.AddBookmark(ctx, songAntenna.ID, "this is a comment", 123)).To(BeNil()) - bms, err := mr.GetBookmarks() + bms, err := mr.GetBookmarks(ctx) Expect(err).ToNot(HaveOccurred()) Expect(bms).To(HaveLen(1)) @@ -42,9 +42,9 @@ var _ = Describe("sqlBookmarks", func() { Expect(updated).To(BeTemporally(">=", created)) By("Overriding the bookmark") - Expect(mr.AddBookmark(songAntenna.ID, "another comment", 333)).To(BeNil()) + Expect(mr.AddBookmark(ctx, songAntenna.ID, "another comment", 333)).To(BeNil()) - bms, err = mr.GetBookmarks() + bms, err = mr.GetBookmarks(ctx) Expect(err).ToNot(HaveOccurred()) Expect(bms[0].Item.ID).To(Equal(songAntenna.ID)) @@ -54,66 +54,66 @@ var _ = Describe("sqlBookmarks", func() { Expect(bms[0].UpdatedAt).To(BeTemporally(">=", updated)) By("Saving another bookmark") - Expect(mr.AddBookmark(songComeTogether.ID, "one more comment", 444)).To(BeNil()) - bms, err = mr.GetBookmarks() + Expect(mr.AddBookmark(ctx, songComeTogether.ID, "one more comment", 444)).To(BeNil()) + bms, err = mr.GetBookmarks(ctx) Expect(err).ToNot(HaveOccurred()) Expect(bms).To(HaveLen(2)) By("Delete bookmark") - Expect(mr.DeleteBookmark(songAntenna.ID)).To(Succeed()) - bms, err = mr.GetBookmarks() + Expect(mr.DeleteBookmark(ctx, songAntenna.ID)).To(Succeed()) + bms, err = mr.GetBookmarks(ctx) Expect(err).ToNot(HaveOccurred()) Expect(bms).To(HaveLen(1)) Expect(bms[0].Item.ID).To(Equal(songComeTogether.ID)) Expect(bms[0].Item.Title).To(Equal(songComeTogether.Title)) - Expect(mr.DeleteBookmark(songComeTogether.ID)).To(Succeed()) - Expect(mr.GetBookmarks()).To(BeEmpty()) + Expect(mr.DeleteBookmark(ctx, songComeTogether.ID)).To(Succeed()) + Expect(mr.GetBookmarks(ctx)).To(BeEmpty()) }) }) Describe("library access", func() { var otherLib model.Library var restrictedUser model.User - var adminCtx context.Context + var adminCtx, userCtx context.Context var userMr model.MediaFileRepository BeforeEach(func() { adminCtx, otherLib, restrictedUser = restrictedFixture("bmk") - adminMr := NewMediaFileRepository(adminCtx, GetDBXBuilder()) - Expect(adminMr.Put(&model.MediaFile{ + adminMr := NewMediaFileRepository(GetDBXBuilder()) + Expect(adminMr.Put(adminCtx, &model.MediaFile{ ID: "bmk-otherlib-track", LibraryID: otherLib.ID, Path: "hidden/bookmarked.mp3", Title: "Hidden Bookmarked", })).To(Succeed()) - DeferCleanup(func() { _ = adminMr.Delete("bmk-otherlib-track") }) + DeferCleanup(func() { _ = adminMr.Delete(adminCtx, "bmk-otherlib-track") }) - userCtx := request.WithUser(log.NewContext(GinkgoT().Context()), restrictedUser) - userMr = NewMediaFileRepository(userCtx, GetDBXBuilder()) + userCtx = request.WithUser(ctx, restrictedUser) + userMr = NewMediaFileRepository(GetDBXBuilder()) }) It("does not return bookmarks for tracks outside the user's libraries", func() { - Expect(userMr.AddBookmark("bmk-otherlib-track", "sneaky", 1)).To(Succeed()) + Expect(userMr.AddBookmark(userCtx, "bmk-otherlib-track", "sneaky", 1)).To(Succeed()) - Expect(userMr.GetBookmarks()).To(BeEmpty()) + Expect(userMr.GetBookmarks(userCtx)).To(BeEmpty()) }) It("still returns the bookmark for an admin", func() { - adminMr := NewMediaFileRepository(adminCtx, GetDBXBuilder()) - Expect(adminMr.AddBookmark("bmk-otherlib-track", "mine", 1)).To(Succeed()) - DeferCleanup(func() { _ = adminMr.DeleteBookmark("bmk-otherlib-track") }) + adminMr := NewMediaFileRepository(GetDBXBuilder()) + Expect(adminMr.AddBookmark(adminCtx, "bmk-otherlib-track", "mine", 1)).To(Succeed()) + DeferCleanup(func() { _ = adminMr.DeleteBookmark(adminCtx, "bmk-otherlib-track") }) - bms, err := adminMr.GetBookmarks() + bms, err := adminMr.GetBookmarks(adminCtx) Expect(err).ToNot(HaveOccurred()) Expect(bms).To(HaveLen(1)) Expect(bms[0].Item.ID).To(Equal("bmk-otherlib-track")) }) It("keeps returning bookmarks for tracks inside the user's libraries", func() { - Expect(userMr.AddBookmark(songAntenna.ID, "allowed", 5)).To(Succeed()) - DeferCleanup(func() { _ = userMr.DeleteBookmark(songAntenna.ID) }) + Expect(userMr.AddBookmark(userCtx, songAntenna.ID, "allowed", 5)).To(Succeed()) + DeferCleanup(func() { _ = userMr.DeleteBookmark(userCtx, songAntenna.ID) }) - bms, err := userMr.GetBookmarks() + bms, err := userMr.GetBookmarks(userCtx) Expect(err).ToNot(HaveOccurred()) Expect(bms).To(HaveLen(1)) Expect(bms[0].Item.ID).To(Equal(songAntenna.ID)) diff --git a/persistence/sql_participations.go b/persistence/sql_participations.go index 746abed01..bc2af39e2 100644 --- a/persistence/sql_participations.go +++ b/persistence/sql_participations.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "encoding/json" "fmt" @@ -74,12 +75,12 @@ func unmarshalParticipants(data string) (model.Participants, error) { return participants, nil } -func (r sqlRepository) updateParticipants(itemID string, participants model.Participants) error { +func (r sqlRepository) updateParticipants(ctx context.Context, itemID string, participants model.Participants) error { // Delete all existing participant entries for this item. // This ensures stale role associations are removed when an artist's role changes // (e.g., an artist was both albumartist and composer, but is now only composer). sqd := Delete(r.tableName + "_artists").Where(Eq{r.tableName + "_id": itemID}) - _, err := r.executeSQL(sqd) + _, err := r.executeSQL(ctx, sqd) if err != nil { return err } @@ -119,14 +120,14 @@ func (r sqlRepository) updateParticipants(itemID string, participants model.Part ON CONFLICT (artist_id, %[1]s_id, role, sub_role) DO NOTHING -- Ignore duplicates `, r.tableName) - _, err = r.executeSQL(Expr(query, itemID, string(participantsJSON))) + _, err = r.executeSQL(ctx, Expr(query, itemID, string(participantsJSON))) return err } -func (r *sqlRepository) getParticipants(m *model.MediaFile) (model.Participants, error) { - ar := NewArtistRepository(r.ctx, r.db) +func (r *sqlRepository) getParticipants(ctx context.Context, m *model.MediaFile) (model.Participants, error) { + ar := NewArtistRepository(r.db) ids := m.Participants.AllIDs() - artists, err := ar.GetAll(model.QueryOptions{Filters: Eq{"artist.id": ids}}) + artists, err := ar.GetAll(ctx, model.QueryOptions{Filters: Eq{"artist.id": ids}}) if err != nil { return nil, fmt.Errorf("getting participants: %w", err) } diff --git a/persistence/sql_restful.go b/persistence/sql_restful.go index b1cfd2379..a758cb847 100644 --- a/persistence/sql_restful.go +++ b/persistence/sql_restful.go @@ -54,7 +54,7 @@ func (r *sqlRepository) parseRestFilters(ctx context.Context, options rest.Query func (r *sqlRepository) parseRestOptions(ctx context.Context, options ...rest.QueryOptions) model.QueryOptions { qo := model.QueryOptions{} if len(options) > 0 { - qo.Sort, qo.Order = r.sanitizeSort(options[0].Sort, options[0].Order) + qo.Sort, qo.Order = r.sanitizeSort(ctx, options[0].Sort, options[0].Order) qo.Max = options[0].Max qo.Offset = options[0].Offset if seed, ok := options[0].Filters["seed"].(string); ok { @@ -66,13 +66,13 @@ func (r *sqlRepository) parseRestOptions(ctx context.Context, options ...rest.Qu return qo } -func (r sqlRepository) sanitizeSort(sort, order string) (string, string) { +func (r sqlRepository) sanitizeSort(ctx context.Context, sort, order string) (string, string) { if sort != "" { sort = toSnakeCase(sort) // Validate only: buildSortOrder resolves the mapping later, and mapping here as well would // feed sortMapping its own output. if _, _, known := r.lookupSortMapping(sort); !known && !r.isFieldWhiteListed(sort) { - log.Warn(r.ctx, "Ignoring sort not whitelisted", "sort", sort, "table", r.tableName) + log.Warn(ctx, "Ignoring sort not whitelisted", "sort", sort, "table", r.tableName) sort = "" } } @@ -129,11 +129,9 @@ func idFilter(tableName string) func(string, any) Sqlizer { return func(field string, value any) Sqlizer { return Eq{tableName + ".id": value} } } -func invalidFilter(ctx context.Context) func(string, any) Sqlizer { - return func(field string, value any) Sqlizer { - log.Warn(ctx, "Invalid filter", "fieldName", field, "value", value) - return Eq{"1": "0"} - } +func invalidFilter(field string, value any) Sqlizer { + log.Warn("Invalid filter", "fieldName", field, "value", value) + return Eq{"1": "0"} } var ( diff --git a/persistence/sql_search.go b/persistence/sql_search.go index 3049baae7..aa8707247 100644 --- a/persistence/sql_search.go +++ b/persistence/sql_search.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "fmt" "strings" @@ -32,7 +33,7 @@ type searchConfig struct { // including FTS Phase 1 which builds its own query outside sq. type searchStrategy interface { Sqlizer - execute(r sqlRepository, sq SelectBuilder, dest any, cfg searchConfig, options model.QueryOptions) error + execute(ctx context.Context, r sqlRepository, sq SelectBuilder, dest any, cfg searchConfig, options model.QueryOptions) error } // getSearchStrategy returns the appropriate search strategy based on config and query content. @@ -51,7 +52,7 @@ func getSearchStrategy(tableName, query string) searchStrategy { // otherwise delegates to getSearchStrategy. sq must already have LIMIT/OFFSET set // via newSelect(options...). options is forwarded so FTS Phase 1 can apply the same // filters and pagination independently. -func (r sqlRepository) doSearch(sq SelectBuilder, q string, results any, cfg searchConfig, options model.QueryOptions) error { +func (r sqlRepository) doSearch(ctx context.Context, sq SelectBuilder, q string, results any, cfg searchConfig, options model.QueryOptions) error { q = strings.TrimSpace(q) q = strings.TrimSuffix(q, "*") @@ -60,13 +61,13 @@ func (r sqlRepository) doSearch(sq SelectBuilder, q string, results any, cfg sea // Empty query (OpenSubsonic `search3?query=""`) — return all in natural order. if q == "" || q == `""` { rowidCore := Select(r.tableName + ".rowid").From(r.tableName).OrderBy(cfg.NaturalOrder) - return r.executeTwoPhase(sq, results, rowidCore, cfg, options) + return r.executeTwoPhase(ctx, sq, results, rowidCore, cfg, options) } // MBID search: if query is a valid UUID, search by MBID fields instead if uuid.Validate(q) == nil && len(cfg.MBIDFields) > 0 { sq = sq.Where(mbidExpr(r.tableName, q, cfg.MBIDFields...)) - return r.queryAll(sq, results) + return r.queryAll(ctx, sq, results) } // Min-length guard: single-character queries are too broad for search3. @@ -81,7 +82,7 @@ func (r sqlRepository) doSearch(sq SelectBuilder, q string, results any, cfg sea return nil } - return strategy.execute(r, sq, results, cfg, options) + return strategy.execute(ctx, r, sq, results, cfg, options) } // executeTwoPhase runs a search in two phases: @@ -91,7 +92,7 @@ func (r sqlRepository) doSearch(sq SelectBuilder, q string, results any, cfg sea // covering index; with those JOINs, large offsets degrade to O(offset) join probes — // multi-second responses on 100k+ libraries. // - Phase 2: full SELECT with all JOINs, scoped to Phase 1's rowid page. -func (r sqlRepository) executeTwoPhase(sq SelectBuilder, results any, rowidCore SelectBuilder, cfg searchConfig, options model.QueryOptions) error { +func (r sqlRepository) executeTwoPhase(ctx context.Context, sq SelectBuilder, results any, rowidCore SelectBuilder, cfg searchConfig, options model.QueryOptions) error { rowidQuery := rowidCore. Where(Eq{r.tableName + ".missing": false}) if options.Max > 0 { @@ -103,17 +104,17 @@ func (r sqlRepository) executeTwoPhase(sq SelectBuilder, results any, rowidCore if cfg.LibraryFilter != nil { rowidQuery = cfg.LibraryFilter(rowidQuery) } else { - rowidQuery = r.applyLibraryFilter(rowidQuery) + rowidQuery = r.applyLibraryFilter(ctx, rowidQuery) } if options.Filters != nil { rowidQuery = rowidQuery.Where(options.Filters) } - return r.hydrateRowidPage(sq, rowidQuery, results) + return r.hydrateRowidPage(ctx, sq, rowidQuery, results) } // hydrateRowidPage joins sq to the ordered rowid set produced by rowidQuery, preserving its // ordering. rowidQuery must handle pagination itself; sq's LIMIT/OFFSET are stripped. -func (r sqlRepository) hydrateRowidPage(sq SelectBuilder, rowidQuery SelectBuilder, results any) error { +func (r sqlRepository) hydrateRowidPage(ctx context.Context, sq SelectBuilder, rowidQuery SelectBuilder, results any) error { rowidSQL, rowidArgs, err := rowidQuery.ToSql() if err != nil { return fmt.Errorf("building rowid query: %w", err) @@ -125,7 +126,7 @@ func (r sqlRepository) hydrateRowidPage(sq SelectBuilder, rowidQuery SelectBuild ) sq = sq.Join(rankedSubquery+" ON "+r.tableName+".rowid = _ranked._rid", rowidArgs...) sq = sq.OrderBy("_ranked._rn") - return r.queryAll(sq, results) + return r.queryAll(ctx, sq, results) } func mbidExpr(tableName, mbid string, mbidFields ...string) Sqlizer { diff --git a/persistence/sql_search_fts.go b/persistence/sql_search_fts.go index fce77afbb..ee43680e6 100644 --- a/persistence/sql_search_fts.go +++ b/persistence/sql_search_fts.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "fmt" "regexp" "strings" @@ -270,7 +271,7 @@ func (s *ftsSearch) ToSql() (string, []any, error) { // execute runs a two-phase FTS5 search (see executeTwoPhase): Phase 1 here contributes the // FTS MATCH join and BM25 rank ordering. Complex ORDER BY (function calls, aggregations) are // dropped from Phase 1. -func (s *ftsSearch) execute(r sqlRepository, sq SelectBuilder, dest any, cfg searchConfig, options model.QueryOptions) error { +func (s *ftsSearch) execute(ctx context.Context, r sqlRepository, sq SelectBuilder, dest any, cfg searchConfig, options model.QueryOptions) error { qualifiedOrderBys := []string{s.rankExpr} for _, ob := range cfg.OrderBy { if qualified := qualifyOrderBy(s.tableName, ob); qualified != "" { @@ -282,7 +283,7 @@ func (s *ftsSearch) execute(r sqlRepository, sq SelectBuilder, dest any, cfg sea From(s.tableName). Join(s.ftsTable+" ON "+s.ftsTable+".rowid = "+s.tableName+".rowid AND "+s.ftsTable+" MATCH ?", s.matchExpr). OrderBy(qualifiedOrderBys...) - return r.executeTwoPhase(sq, dest, rowidCore, cfg, options) + return r.executeTwoPhase(ctx, sq, dest, rowidCore, cfg, options) } // qualifyOrderBy prepends tableName to a simple column name. Returns empty string for diff --git a/persistence/sql_search_fts_test.go b/persistence/sql_search_fts_test.go index d0b26e8d5..5d52cb1e8 100644 --- a/persistence/sql_search_fts_test.go +++ b/persistence/sql_search_fts_test.go @@ -317,20 +317,20 @@ var _ = Describe("FTS5 Integration Search", func() { mr model.MediaFileRepository alr model.AlbumRepository arr model.ArtistRepository + ctx context.Context ) BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, adminUser) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) conn := GetDBXBuilder() - mr = NewMediaFileRepository(ctx, conn) - alr = NewAlbumRepository(ctx, conn) - arr = NewArtistRepository(ctx, conn) + mr = NewMediaFileRepository(conn) + alr = NewAlbumRepository(conn) + arr = NewArtistRepository(conn) }) Describe("MediaFile search", func() { It("finds media files by title", func() { - results, err := mr.Search("Radioactivity", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "Radioactivity", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Title).To(Equal("Radioactivity")) @@ -338,7 +338,7 @@ var _ = Describe("FTS5 Integration Search", func() { }) It("finds media files by artist name", func() { - results, err := mr.Search("Beatles", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "Beatles", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(3)) for _, r := range results { @@ -349,7 +349,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("Album search", func() { It("finds albums by name", func() { - results, err := alr.Search("Sgt Peppers", model.QueryOptions{Max: 10}) + results, err := alr.Search(ctx, "Sgt Peppers", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Name).To(Equal("Sgt Peppers")) @@ -357,7 +357,7 @@ var _ = Describe("FTS5 Integration Search", func() { }) It("finds albums with multi-word search", func() { - results, err := alr.Search("Abbey Road", model.QueryOptions{Max: 10}) + results, err := alr.Search(ctx, "Abbey Road", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(2)) }) @@ -365,7 +365,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("Artist search", func() { It("finds artists by name", func() { - results, err := arr.Search("Kraftwerk", model.QueryOptions{Max: 10}) + results, err := arr.Search(ctx, "Kraftwerk", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Name).To(Equal("Kraftwerk")) @@ -375,7 +375,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("CJK search", func() { It("finds media files by CJK title", func() { - results, err := mr.Search("プラチナ", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "プラチナ", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Title).To(Equal("プラチナ・ジェット")) @@ -383,14 +383,14 @@ var _ = Describe("FTS5 Integration Search", func() { }) It("finds media files by CJK artist name", func() { - results, err := mr.Search("シートベルツ", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "シートベルツ", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Artist).To(Equal("シートベルツ")) }) It("finds albums by CJK artist name", func() { - results, err := alr.Search("シートベルツ", model.QueryOptions{Max: 10}) + results, err := alr.Search(ctx, "シートベルツ", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Name).To(Equal("COWBOY BEBOP")) @@ -398,7 +398,7 @@ var _ = Describe("FTS5 Integration Search", func() { }) It("finds artists by CJK name", func() { - results, err := arr.Search("シートベルツ", model.QueryOptions{Max: 10}) + results, err := arr.Search(ctx, "シートベルツ", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Name).To(Equal("シートベルツ")) @@ -408,7 +408,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("Album version search", func() { It("finds albums by version tag via FTS", func() { - results, err := alr.Search("Deluxe", model.QueryOptions{Max: 10}) + results, err := alr.Search(ctx, "Deluxe", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].ID).To(Equal(albumWithVersion.ID)) @@ -417,7 +417,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("Punctuation-only search", func() { It("finds media files with punctuation-only title", func() { - results, err := mr.Search("!!!!!!!", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "!!!!!!!", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Title).To(Equal("!!!!!!!")) @@ -427,7 +427,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("Single-character search (doSearch min-length guard)", func() { It("returns empty results for single-char query via Search", func() { - results, err := mr.Search("a", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "a", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty(), "doSearch should reject single-char queries") }) @@ -435,7 +435,7 @@ var _ = Describe("FTS5 Integration Search", func() { Describe("Max=0 means no limit (regression: must not produce LIMIT 0)", func() { It("returns results with Max=0", func() { - results, err := mr.Search("Beatles", model.QueryOptions{Max: 0}) + results, err := mr.Search(ctx, "Beatles", model.QueryOptions{Max: 0}) Expect(err).ToNot(HaveOccurred()) Expect(results).ToNot(BeEmpty(), "Max=0 should mean no limit, not LIMIT 0") }) @@ -456,19 +456,19 @@ var _ = Describe("FTS5 Integration Search", func() { {ID: "fts-rank-2", Name: "Modest Mouse", OrderArtistName: "modest mouse"}, {ID: "fts-rank-3", Name: "Morrissey", OrderArtistName: "morrissey"}, } { - Expect(createArtistWithLibrary(arr, &a, 1)).To(Succeed()) + Expect(createArtistWithLibrary(ctx, arr, &a, 1)).To(Succeed()) } }) It("ranks the exact transliterated match first for 'MO'", func() { - results, err := arr.Search("MO", model.QueryOptions{Max: 10}) + results, err := arr.Search(ctx, "MO", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(3)) Expect(results[0].Name).To(Equal("MØ"), "exact match via search_normalized must outrank prefix matches") }) It("ranks the exact match first for the accented query 'MØ'", func() { - results, err := arr.Search("MØ", model.QueryOptions{Max: 10}) + results, err := arr.Search(ctx, "MØ", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).ToNot(BeEmpty()) Expect(results[0].Name).To(Equal("MØ")) diff --git a/persistence/sql_search_like.go b/persistence/sql_search_like.go index 972545ac5..c7d94cc04 100644 --- a/persistence/sql_search_like.go +++ b/persistence/sql_search_like.go @@ -1,6 +1,7 @@ package persistence import ( + "context" "strings" . "github.com/Masterminds/squirrel" @@ -20,10 +21,10 @@ func (s *likeSearch) ToSql() (string, []any, error) { return s.filter.ToSql() } -func (s *likeSearch) execute(r sqlRepository, sq SelectBuilder, dest any, cfg searchConfig, options model.QueryOptions) error { +func (s *likeSearch) execute(ctx context.Context, r sqlRepository, sq SelectBuilder, dest any, cfg searchConfig, options model.QueryOptions) error { sq = sq.Where(s.filter) sq = sq.OrderBy(cfg.OrderBy...) - return r.queryAll(sq, dest, options) + return r.queryAll(ctx, sq, dest, options) } // newLegacySearch creates a LIKE search against the full_text column. diff --git a/persistence/sql_search_like_test.go b/persistence/sql_search_like_test.go index 8ee4ef93c..2bfeba29b 100644 --- a/persistence/sql_search_like_test.go +++ b/persistence/sql_search_like_test.go @@ -102,32 +102,32 @@ var _ = Describe("likeSearchExpr", func() { var _ = Describe("Legacy Integration Search", func() { var mr model.MediaFileRepository + var ctx context.Context BeforeEach(func() { + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) DeferCleanup(configtest.SetupConfig()) conf.Server.Search.Backend = "legacy" - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, adminUser) conn := GetDBXBuilder() - mr = NewMediaFileRepository(ctx, conn) + mr = NewMediaFileRepository(conn) }) It("returns results using legacy LIKE-based search", func() { - results, err := mr.Search("Radioactivity", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "Radioactivity", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(HaveLen(1)) Expect(results[0].Title).To(Equal("Radioactivity")) }) It("returns empty results for single-char query (doSearch min-length guard)", func() { - results, err := mr.Search("a", model.QueryOptions{Max: 10}) + results, err := mr.Search(ctx, "a", model.QueryOptions{Max: 10}) Expect(err).ToNot(HaveOccurred()) Expect(results).To(BeEmpty(), "doSearch should reject single-char queries") }) It("returns results with Max=0 (regression: must not produce LIMIT 0)", func() { - results, err := mr.Search("Beatles", model.QueryOptions{Max: 0}) + results, err := mr.Search(ctx, "Beatles", model.QueryOptions{Max: 0}) Expect(err).ToNot(HaveOccurred()) Expect(results).ToNot(BeEmpty(), "Max=0 should mean no limit, not LIMIT 0") }) diff --git a/persistence/sql_tags.go b/persistence/sql_tags.go index 11f9fe00e..7e1978453 100644 --- a/persistence/sql_tags.go +++ b/persistence/sql_tags.go @@ -54,9 +54,9 @@ var indexedTagNames = []model.TagName{model.TagGenre} // updateTags rewrites this item's _tags rows from its in-memory tags, mirroring // updateParticipants (delete-then-insert in the same Put; JOIN to tag skips not-yet-saved ids). -func (r sqlRepository) updateTags(itemID string, tags model.Tags) error { +func (r sqlRepository) updateTags(ctx context.Context, itemID string, tags model.Tags) error { del := Delete(r.tableName + "_tags").Where(Eq{r.tableName + "_id": itemID}) - if _, err := r.executeSQL(del); err != nil { + if _, err := r.executeSQL(ctx, del); err != nil { return err } var tagIDs []string @@ -77,7 +77,7 @@ func (r sqlRepository) updateTags(itemID string, tags model.Tags) error { SELECT ?, value FROM json_each(?) JOIN tag ON tag.id = value ON CONFLICT (%[1]s_id, tag_id) DO NOTHING`, r.tableName) - _, err = r.executeSQL(Expr(query, itemID, string(idsJSON))) + _, err = r.executeSQL(ctx, Expr(query, itemID, string(idsJSON))) return err } @@ -146,11 +146,10 @@ type baseTagRepository struct { // newBaseTagRepository creates a new base tag repository with optional tag filtering. // If tagFilter is nil, the repository will work with all tags. // If tagFilter is provided, the repository will only work with tags of that specific name. -func newBaseTagRepository(ctx context.Context, db dbx.Builder, tagFilter *model.TagName) *baseTagRepository { +func newBaseTagRepository(db dbx.Builder, tagFilter *model.TagName) *baseTagRepository { r := &baseTagRepository{ tagFilter: tagFilter, } - r.ctx = ctx r.db = db r.tableName = "tag" r.registerModel(&model.Tag{}, map[string]filterFunc{ @@ -164,12 +163,12 @@ func newBaseTagRepository(ctx context.Context, db dbx.Builder, tagFilter *model. } // applyLibraryFiltering adds the appropriate library joins based on user context -func (r *baseTagRepository) applyLibraryFiltering(sq SelectBuilder) SelectBuilder { +func (r *baseTagRepository) applyLibraryFiltering(ctx context.Context, sq SelectBuilder) SelectBuilder { // Add library_tag join sq = sq.LeftJoin("library_tag on library_tag.tag_id = tag.id") // For authenticated users, also join with user_library to filter by accessible libraries - user := loggedUser(r.ctx) + user := loggedUser(ctx) if user.ID != invalidUserId { sq = sq.Join("user_library on user_library.library_id = library_tag.library_id AND user_library.user_id = ?", user.ID) } @@ -178,8 +177,8 @@ func (r *baseTagRepository) applyLibraryFiltering(sq SelectBuilder) SelectBuilde } // newSelect overrides the base implementation to apply tag name filtering and library filtering. -func (r *baseTagRepository) newSelect(options ...model.QueryOptions) SelectBuilder { - sq := r.sqlRepository.newSelect(options...) +func (r *baseTagRepository) newSelect(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + sq := r.sqlRepository.newSelect(ctx, options...) // Apply tag name filtering if specified if r.tagFilter != nil { @@ -187,7 +186,7 @@ func (r *baseTagRepository) newSelect(options ...model.QueryOptions) SelectBuild } // Apply library filtering and set up aggregation columns - sq = r.applyLibraryFiltering(sq).Columns( + sq = r.applyLibraryFiltering(ctx, sq).Columns( "tag.id", "tag.tag_name", "tag.tag_value", @@ -198,9 +197,9 @@ func (r *baseTagRepository) newSelect(options ...model.QueryOptions) SelectBuild return sq } -// ResourceRepository interface implementation +// REST interface methods -func (r *baseTagRepository) Count(options ...rest.QueryOptions) (int64, error) { +func (r *baseTagRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { sq := Select("COUNT(DISTINCT tag.id)").From("tag") // Apply tag name filtering if specified @@ -209,32 +208,24 @@ func (r *baseTagRepository) Count(options ...rest.QueryOptions) (int64, error) { } // Apply library filtering - sq = r.applyLibraryFiltering(sq) + sq = r.applyLibraryFiltering(ctx, sq) - return r.count(sq, r.parseRestOptions(r.ctx, options...)) + return r.count(ctx, sq, r.parseRestOptions(ctx, options...)) } -func (r *baseTagRepository) Read(id string) (any, error) { - query := r.newSelect().Where(Eq{"id": id}) +func (r *baseTagRepository) Read(ctx context.Context, id string) (*model.Tag, error) { + query := r.newSelect(ctx).Where(Eq{"id": id}) var res model.Tag - err := r.queryOne(query, &res) + err := r.queryOne(ctx, query, &res) return &res, err } -func (r *baseTagRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - query := r.newSelect(r.parseRestOptions(r.ctx, options...)) +func (r *baseTagRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Tag, error) { + query := r.newSelect(ctx, r.parseRestOptions(ctx, options...)) var res model.TagList - err := r.queryAll(query, &res) + err := r.queryAll(ctx, query, &res) return res, err } -func (r *baseTagRepository) EntityName() string { - return "tag" -} - -func (r *baseTagRepository) NewInstance() any { - return model.Tag{} -} - // Interface compliance check -var _ model.ResourceRepository = (*baseTagRepository)(nil) +var _ rest.Repository[model.Tag] = (*baseTagRepository)(nil) diff --git a/persistence/tag_library_filtering_test.go b/persistence/tag_library_filtering_test.go index ddd897165..a4382dccf 100644 --- a/persistence/tag_library_filtering_test.go +++ b/persistence/tag_library_filtering_test.go @@ -78,11 +78,11 @@ var _ = Describe("Tag Library Filtering", func() { // Create test tags adminCtx := request.WithUser(log.NewContext(context.TODO()), adminUser) - tagRepo := NewTagRepository(adminCtx, GetDBXBuilder()) + tagRepo := NewTagRepository(GetDBXBuilder()) createTag := func(libraryID int, name, value string) { tag := model.Tag{ID: id.NewTagID(name, value), TagName: model.TagName(name), TagValue: value} - err := tagRepo.Add(libraryID, tag) + err := tagRepo.Add(adminCtx, libraryID, tag) Expect(err).ToNot(HaveOccurred()) } @@ -119,17 +119,16 @@ var _ = Describe("Tag Library Filtering", func() { ctx = context.Background() // Headless context } - tagRepo := NewTagRepository(ctx, GetDBXBuilder()) - repo := tagRepo.(model.ResourceRepository) + repo := NewTagRepository(GetDBXBuilder()) var opts rest.QueryOptions if len(filters) > 0 { opts = filters[0] } - tags, err := repo.ReadAll(opts) + tags, err := repo.ReadAll(ctx, opts) Expect(err).ToNot(HaveOccurred()) - return tags.(model.TagList) + return tags } // Helper to count tags @@ -141,10 +140,9 @@ var _ = Describe("Tag Library Filtering", func() { ctx = context.Background() } - tagRepo := NewTagRepository(ctx, GetDBXBuilder()) - repo := tagRepo.(model.ResourceRepository) + repo := NewTagRepository(GetDBXBuilder()) - count, err := repo.Count() + count, err := repo.Count(ctx) Expect(err).ToNot(HaveOccurred()) return count } diff --git a/persistence/tag_repository.go b/persistence/tag_repository.go index f2093c9d8..d13dece66 100644 --- a/persistence/tag_repository.go +++ b/persistence/tag_repository.go @@ -7,6 +7,7 @@ import ( "time" . "github.com/Masterminds/squirrel" + "github.com/deluan/rest" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/pocketbase/dbx" @@ -16,20 +17,20 @@ type tagRepository struct { *baseTagRepository } -func NewTagRepository(ctx context.Context, db dbx.Builder) model.TagRepository { +func NewTagRepository(db dbx.Builder) model.TagRepository { return &tagRepository{ - baseTagRepository: newBaseTagRepository(ctx, db, nil), // nil = no filter, works with all tags + baseTagRepository: newBaseTagRepository(db, nil), // nil = no filter, works with all tags } } -func (r *tagRepository) Add(libraryID int, tags ...model.Tag) error { +func (r *tagRepository) Add(ctx context.Context, libraryID int, tags ...model.Tag) error { for chunk := range slices.Chunk(tags, 200) { sq := Insert(r.tableName).Columns("id", "tag_name", "tag_value"). Suffix("on conflict (id) do nothing") for _, t := range chunk { sq = sq.Values(t.ID, t.TagName, t.TagValue) } - _, err := r.executeSQL(sq) + _, err := r.executeSQL(ctx, sq) if err != nil { return err } @@ -40,7 +41,7 @@ func (r *tagRepository) Add(libraryID int, tags ...model.Tag) error { for _, t := range chunk { libSq = libSq.Values(t.ID, libraryID, 0, 0) } - _, err = r.executeSQL(libSq) + _, err = r.executeSQL(ctx, libSq) if err != nil { return fmt.Errorf("adding library_tag entries: %w", err) } @@ -50,7 +51,7 @@ func (r *tagRepository) Add(libraryID int, tags ...model.Tag) error { // UpdateCounts updates the library_tag table with per-library statistics. // Only genres are being updated for now. -func (r *tagRepository) UpdateCounts() error { +func (r *tagRepository) UpdateCounts(ctx context.Context) error { template := ` INSERT INTO library_tag (tag_id, library_id, %[1]s_count) SELECT jt.value as tag_id, %[1]s.library_id, count(distinct %[1]s.id) as %[1]s_count @@ -65,8 +66,8 @@ DO UPDATE SET %[1]s_count = excluded.%[1]s_count; for _, table := range []string{"album", "media_file"} { start := time.Now() query := Expr(fmt.Sprintf(template, table)) - c, err := r.executeSQL(query) - log.Debug(r.ctx, "Updated library tag counts", "table", table, "elapsed", time.Since(start), "updated", c) + c, err := r.executeSQL(ctx, query) + log.Debug(ctx, "Updated library tag counts", "table", table, "elapsed", time.Since(start), "updated", c) if err != nil { return fmt.Errorf("updating %s library tag counts: %w", table, err) } @@ -74,14 +75,14 @@ DO UPDATE SET %[1]s_count = excluded.%[1]s_count; return nil } -func (r *tagRepository) GetAll(name model.TagName, options ...model.QueryOptions) (model.TagList, error) { - sq := r.newSelect(options...).Where(Eq{"tag.tag_name": name}) +func (r *tagRepository) GetAll(ctx context.Context, name model.TagName, options ...model.QueryOptions) (model.TagList, error) { + sq := r.newSelect(ctx, options...).Where(Eq{"tag.tag_name": name}) res := model.TagList{} - err := r.queryAll(sq, &res) + err := r.queryAll(ctx, sq, &res) return res, err } -func (r *tagRepository) purgeUnused() error { +func (r *tagRepository) purgeUnused(ctx context.Context) error { del := Delete(r.tableName).Where(` id not in (select jt.value from album left join json_tree(album.tags, '$') as jt @@ -93,14 +94,14 @@ func (r *tagRepository) purgeUnused() error { where atom is not null and key = 'id') `) - c, err := r.executeSQL(del) + c, err := r.executeSQL(ctx, del) if err != nil { return fmt.Errorf("error purging %s unused tags: %w", r.tableName, err) } if c > 0 { - log.Debug(r.ctx, "Purged unused tags", "totalDeleted", c, "table", r.tableName) + log.Debug(ctx, "Purged unused tags", "totalDeleted", c, "table", r.tableName) } return err } -var _ model.ResourceRepository = &tagRepository{} +var _ rest.Repository[model.Tag] = &tagRepository{} diff --git a/persistence/tag_repository_test.go b/persistence/tag_repository_test.go index 9a019c30e..730d26357 100644 --- a/persistence/tag_repository_test.go +++ b/persistence/tag_repository_test.go @@ -18,15 +18,15 @@ import ( var _ = Describe("TagRepository", func() { var repo model.TagRepository - var restRepo model.ResourceRepository + var restRepo rest.Repository[model.Tag] var ctx context.Context BeforeEach(func() { DeferCleanup(configtest.SetupConfig()) ctx = request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid", UserName: "johndoe", IsAdmin: true}) - tagRepo := NewTagRepository(ctx, GetDBXBuilder()) + tagRepo := NewTagRepository(GetDBXBuilder()) repo = tagRepo - restRepo = tagRepo.(model.ResourceRepository) + restRepo = tagRepo // Clean the database before each test to ensure isolation db := GetDBXBuilder() @@ -48,7 +48,7 @@ var _ = Describe("TagRepository", func() { return model.Tag{ID: id.NewTagID(name, value), TagName: model.TagName(name), TagValue: value} } - err = repo.Add(1, + err = repo.Add(ctx, 1, // Genre tags newTag("genre", "rock"), newTag("genre", "pop"), @@ -84,17 +84,16 @@ var _ = Describe("TagRepository", func() { TagValue: "experimental", } - err := repo.Add(1, newTag) + err := repo.Add(ctx, 1, newTag) Expect(err).ToNot(HaveOccurred()) // Verify tag was added - result, err := restRepo.Read(newTag.ID) + resultTag, err := restRepo.Read(ctx, newTag.ID) Expect(err).ToNot(HaveOccurred()) - resultTag := result.(*model.Tag) Expect(resultTag.TagValue).To(Equal("experimental")) // Check count increased - count, err := restRepo.Count() + count, err := restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(21))) // 20 from dataset + 1 new }) @@ -107,15 +106,15 @@ var _ = Describe("TagRepository", func() { TagValue: "rock", } - count, err := restRepo.Count() + count, err := restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(20))) // Still 20 tags - err = repo.Add(1, duplicateTag) + err = repo.Add(ctx, 1, duplicateTag) Expect(err).ToNot(HaveOccurred()) // Should not error // Count should remain the same - count, err = restRepo.Count() + count, err = restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(20))) // Still 20 tags }) @@ -123,7 +122,7 @@ var _ = Describe("TagRepository", func() { Describe("UpdateCounts", func() { It("should update tag counts successfully", func() { - err := repo.UpdateCounts() + err := repo.UpdateCounts(ctx) Expect(err).ToNot(HaveOccurred()) }) @@ -133,7 +132,7 @@ var _ = Describe("TagRepository", func() { _, err := db.NewQuery("DELETE FROM tag").Execute() Expect(err).ToNot(HaveOccurred()) - err = repo.UpdateCounts() + err = repo.UpdateCounts(ctx) Expect(err).ToNot(HaveOccurred()) }) @@ -159,7 +158,7 @@ var _ = Describe("TagRepository", func() { Expect(err).ToNot(HaveOccurred()) // This should not fail with foreign key constraint error - err = repo.UpdateCounts() + err = repo.UpdateCounts(ctx) Expect(err).ToNot(HaveOccurred()) // Cleanup @@ -189,7 +188,7 @@ var _ = Describe("TagRepository", func() { Expect(err).ToNot(HaveOccurred()) // This should not fail with foreign key constraint error - err = repo.UpdateCounts() + err = repo.UpdateCounts(ctx) Expect(err).ToNot(HaveOccurred()) // Cleanup @@ -201,7 +200,7 @@ var _ = Describe("TagRepository", func() { Describe("Count", func() { It("should return correct count of tags", func() { - count, err := restRepo.Count() + count, err := restRepo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(int64(20))) // From the test dataset }) @@ -210,25 +209,23 @@ var _ = Describe("TagRepository", func() { Describe("Read", func() { It("should return existing tag", func() { rockID := id.NewTagID("genre", "rock") - result, err := restRepo.Read(rockID) + resultTag, err := restRepo.Read(ctx, rockID) Expect(err).ToNot(HaveOccurred()) - resultTag := result.(*model.Tag) Expect(resultTag.ID).To(Equal(rockID)) Expect(resultTag.TagName).To(Equal(model.TagName("genre"))) Expect(resultTag.TagValue).To(Equal("rock")) }) It("should return error for non-existent tag", func() { - _, err := restRepo.Read("non-existent-id") + _, err := restRepo.Read(ctx, "non-existent-id") Expect(err).To(HaveOccurred()) }) }) Describe("ReadAll", func() { It("should return all tags from dataset", func() { - result, err := restRepo.ReadAll() + tags, err := restRepo.ReadAll(ctx) Expect(err).ToNot(HaveOccurred()) - tags := result.(model.TagList) Expect(tags).To(HaveLen(20)) }) @@ -236,9 +233,8 @@ var _ = Describe("TagRepository", func() { options := rest.QueryOptions{ Filters: map[string]any{"name": "%rock%"}, // Tags containing 'rock' } - result, err := restRepo.ReadAll(options) + tags, err := restRepo.ReadAll(ctx, options) Expect(err).ToNot(HaveOccurred()) - tags := result.(model.TagList) Expect(tags).To(HaveLen(2)) // "rock" and "Alternative Rock" // Verify all returned tags contain 'rock' in their value @@ -251,9 +247,8 @@ var _ = Describe("TagRepository", func() { options := rest.QueryOptions{ Filters: map[string]any{"name": "%e%"}, // Tags containing 'e' } - result, err := restRepo.ReadAll(options) + tags, err := restRepo.ReadAll(ctx, options) Expect(err).ToNot(HaveOccurred()) - tags := result.(model.TagList) Expect(tags).To(HaveLen(8)) // electronic, house, trance, energetic, Blues, decade x2, Alternative Rock // Verify all returned tags contain 'e' in their value @@ -268,9 +263,8 @@ var _ = Describe("TagRepository", func() { Sort: "name", Order: "asc", } - result, err := restRepo.ReadAll(options) + tags, err := restRepo.ReadAll(ctx, options) Expect(err).ToNot(HaveOccurred()) - tags := result.(model.TagList) Expect(tags).To(HaveLen(7)) Expect(slices.IsSortedFunc(tags, func(a, b model.Tag) int { @@ -284,9 +278,8 @@ var _ = Describe("TagRepository", func() { Sort: "name", Order: "desc", } - result, err := restRepo.ReadAll(options) + tags, err := restRepo.ReadAll(ctx, options) Expect(err).ToNot(HaveOccurred()) - tags := result.(model.TagList) Expect(tags).To(HaveLen(7)) Expect(slices.IsSortedFunc(tags, func(a, b model.Tag) int { @@ -294,18 +287,4 @@ var _ = Describe("TagRepository", func() { })) }) }) - - Describe("EntityName", func() { - It("should return correct entity name", func() { - name := restRepo.EntityName() - Expect(name).To(Equal("tag")) - }) - }) - - Describe("NewInstance", func() { - It("should return new tag instance", func() { - instance := restRepo.NewInstance() - Expect(instance).To(BeAssignableToTypeOf(model.Tag{})) - }) - }) }) diff --git a/persistence/transcoding_repository.go b/persistence/transcoding_repository.go index fdf67806d..9db56e143 100644 --- a/persistence/transcoding_repository.go +++ b/persistence/transcoding_repository.go @@ -13,63 +13,62 @@ type transcodingRepository struct { sqlRepository } -func NewTranscodingRepository(ctx context.Context, db dbx.Builder) model.TranscodingRepository { +func NewTranscodingRepository(db dbx.Builder) model.TranscodingRepository { r := &transcodingRepository{} - r.ctx = ctx r.db = db r.registerModel(&model.Transcoding{}, nil) return r } -func (r *transcodingRepository) Get(id string) (*model.Transcoding, error) { - sel := r.newSelect().Columns("*").Where(Eq{"id": id}) +func (r *transcodingRepository) Get(ctx context.Context, id string) (*model.Transcoding, error) { + sel := r.newSelect(ctx).Columns("*").Where(Eq{"id": id}) var res model.Transcoding - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *transcodingRepository) CountAll(qo ...model.QueryOptions) (int64, error) { - return r.count(Select(), qo...) +func (r *transcodingRepository) CountAll(ctx context.Context, qo ...model.QueryOptions) (int64, error) { + return r.count(ctx, Select(), qo...) } -func (r *transcodingRepository) FindByFormat(format string) (*model.Transcoding, error) { - sel := r.newSelect().Columns("*").Where(Eq{"target_format": format}) +func (r *transcodingRepository) FindByFormat(ctx context.Context, format string) (*model.Transcoding, error) { + sel := r.newSelect(ctx).Columns("*").Where(Eq{"target_format": format}) var res model.Transcoding - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) return &res, err } -func (r *transcodingRepository) Put(t *model.Transcoding) error { - if !loggedUser(r.ctx).IsAdmin { +func (r *transcodingRepository) Put(ctx context.Context, t *model.Transcoding) error { + if !loggedUser(ctx).IsAdmin { return rest.ErrPermissionDenied } - _, err := r.put(t.ID, t) + _, err := r.put(ctx, t.ID, t) return err } -func (r *transcodingRepository) Count(options ...rest.QueryOptions) (int64, error) { - return r.count(Select(), r.parseRestOptions(r.ctx, options...)) +func (r *transcodingRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.count(ctx, Select(), r.parseRestOptions(ctx, options...)) } -func (r *transcodingRepository) Read(id string) (any, error) { - res, err := r.Get(id) +func (r *transcodingRepository) Read(ctx context.Context, id string) (*model.Transcoding, error) { + res, err := r.Get(ctx, id) if err != nil { return nil, err } - if !loggedUser(r.ctx).IsAdmin { + if !loggedUser(ctx).IsAdmin { res.Command = "" } return res, nil } -func (r *transcodingRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - sel := r.newSelect(r.parseRestOptions(r.ctx, options...)).Columns("*") +func (r *transcodingRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.Transcoding, error) { + sel := r.newSelect(ctx, r.parseRestOptions(ctx, options...)).Columns("*") res := model.Transcodings{} - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } - if !loggedUser(r.ctx).IsAdmin { + if !loggedUser(ctx).IsAdmin { for i := range res { res[i].Command = "" } @@ -77,39 +76,35 @@ func (r *transcodingRepository) ReadAll(options ...rest.QueryOptions) (any, erro return res, nil } -func (r *transcodingRepository) EntityName() string { - return "transcoding" -} - -func (r *transcodingRepository) NewInstance() any { - return &model.Transcoding{} -} - -func (r *transcodingRepository) Save(entity any) (string, error) { - if !loggedUser(r.ctx).IsAdmin { +func (r *transcodingRepository) Save(ctx context.Context, t *model.Transcoding) (string, error) { + if !loggedUser(ctx).IsAdmin { return "", rest.ErrPermissionDenied } - t := entity.(*model.Transcoding) - return r.put(t.ID, t) + return r.put(ctx, t.ID, t) } -func (r *transcodingRepository) Update(id string, entity any, cols ...string) error { - if !loggedUser(r.ctx).IsAdmin { +func (r *transcodingRepository) Update(ctx context.Context, id string, entity model.Transcoding, cols ...string) error { + if !loggedUser(ctx).IsAdmin { return rest.ErrPermissionDenied } - t := entity.(*model.Transcoding) + t := &entity t.ID = id - _, err := r.put(id, t) + _, err := r.put(ctx, id, t) return err } -func (r *transcodingRepository) Delete(id string) error { - if !loggedUser(r.ctx).IsAdmin { +func (r *transcodingRepository) Delete(ctx context.Context, ids ...string) error { + if !loggedUser(ctx).IsAdmin { return rest.ErrPermissionDenied } - return r.deleteByID(id) + for _, id := range ids { + if err := r.deleteByID(ctx, id); err != nil { + return err + } + } + return nil } var _ model.TranscodingRepository = (*transcodingRepository)(nil) -var _ rest.Repository = (*transcodingRepository)(nil) -var _ rest.Persistable = (*transcodingRepository)(nil) +var _ rest.Repository[model.Transcoding] = (*transcodingRepository)(nil) +var _ rest.Persistable[model.Transcoding] = (*transcodingRepository)(nil) diff --git a/persistence/transcoding_repository_test.go b/persistence/transcoding_repository_test.go index 3be46bae3..372b5801f 100644 --- a/persistence/transcoding_repository_test.go +++ b/persistence/transcoding_repository_test.go @@ -1,6 +1,8 @@ package persistence import ( + "context" + "github.com/deluan/rest" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" @@ -11,83 +13,78 @@ import ( var _ = Describe("TranscodingRepository", func() { var repo model.TranscodingRepository - var adminRepo model.TranscodingRepository + var ctx, adminCtx context.Context BeforeEach(func() { - ctx := log.NewContext(GinkgoT().Context()) - ctx = request.WithUser(ctx, regularUser) - repo = NewTranscodingRepository(ctx, GetDBXBuilder()) - - adminCtx := log.NewContext(GinkgoT().Context()) - adminCtx = request.WithUser(adminCtx, adminUser) - adminRepo = NewTranscodingRepository(adminCtx, GetDBXBuilder()) + ctx = request.WithUser(log.NewContext(GinkgoT().Context()), regularUser) + adminCtx = request.WithUser(ctx, adminUser) + repo = NewTranscodingRepository(GetDBXBuilder()) }) AfterEach(func() { // Clean up any transcoding created during the tests - tc, err := adminRepo.FindByFormat("test_format") + tc, err := repo.FindByFormat(adminCtx, "test_format") if err == nil { - err = adminRepo.(*transcodingRepository).Delete(tc.ID) + err = repo.Delete(adminCtx, tc.ID) Expect(err).ToNot(HaveOccurred()) } }) Describe("Admin User", func() { It("creates a new transcoding", func() { - base, err := adminRepo.CountAll() + base, err := repo.CountAll(adminCtx) Expect(err).ToNot(HaveOccurred()) - err = adminRepo.Put(&model.Transcoding{ID: "new", Name: "new", TargetFormat: "test_format", DefaultBitRate: 320, Command: "ffmpeg"}) + err = repo.Put(adminCtx, &model.Transcoding{ID: "new", Name: "new", TargetFormat: "test_format", DefaultBitRate: 320, Command: "ffmpeg"}) Expect(err).ToNot(HaveOccurred()) - count, err := adminRepo.CountAll() + count, err := repo.CountAll(adminCtx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(Equal(base + 1)) }) It("updates an existing transcoding", func() { tr := &model.Transcoding{ID: "upd", Name: "old", TargetFormat: "test_format", DefaultBitRate: 100, Command: "ffmpeg"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) tr.Name = "updated" - err := adminRepo.Put(tr) + err := repo.Put(adminCtx, tr) Expect(err).ToNot(HaveOccurred()) - res, err := adminRepo.FindByFormat("test_format") + res, err := repo.FindByFormat(adminCtx, "test_format") Expect(err).ToNot(HaveOccurred()) Expect(res.Name).To(Equal("updated")) }) It("deletes a transcoding", func() { - err := adminRepo.Put(&model.Transcoding{ID: "to-delete", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 256, Command: "ffmpeg"}) + err := repo.Put(adminCtx, &model.Transcoding{ID: "to-delete", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 256, Command: "ffmpeg"}) Expect(err).ToNot(HaveOccurred()) - err = adminRepo.(*transcodingRepository).Delete("to-delete") + err = repo.Delete(adminCtx, "to-delete") Expect(err).ToNot(HaveOccurred()) - _, err = adminRepo.Get("to-delete") + _, err = repo.Get(adminCtx, "to-delete") Expect(err).To(MatchError(model.ErrNotFound)) }) It("returns not found when deleting a missing transcoding", func() { - err := adminRepo.(*transcodingRepository).Delete("does-not-exist") + err := repo.(*transcodingRepository).Delete(adminCtx, "does-not-exist") Expect(err).To(MatchError(model.ErrNotFound)) }) It("reads the Command field via the REST Read method", func() { tr := &model.Transcoding{ID: "adminread", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg -secret"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) - res, err := adminRepo.(*transcodingRepository).Read("adminread") + res, err := repo.Read(adminCtx, "adminread") Expect(err).ToNot(HaveOccurred()) - Expect(res.(*model.Transcoding).Command).To(Equal("ffmpeg -secret")) + Expect(res.Command).To(Equal("ffmpeg -secret")) }) }) Describe("Regular User", func() { It("reads a transcoding but with the Command field redacted", func() { tr := &model.Transcoding{ID: "readreg", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg -secret"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) - res, err := repo.(*transcodingRepository).Read("readreg") + t, err := repo.Read(ctx, "readreg") Expect(err).ToNot(HaveOccurred()) - t := res.(*model.Transcoding) Expect(t.Name).To(Equal("temp")) Expect(t.TargetFormat).To(Equal("test_format")) Expect(t.Command).To(BeEmpty()) @@ -95,11 +92,10 @@ var _ = Describe("TranscodingRepository", func() { It("lists transcodings but with the Command field redacted", func() { tr := &model.Transcoding{ID: "listreg", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg -secret"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) - res, err := repo.(*transcodingRepository).ReadAll() + list, err := repo.ReadAll(ctx) Expect(err).ToNot(HaveOccurred()) - list := res.(model.Transcodings) Expect(list).ToNot(BeEmpty()) for _, t := range list { Expect(t.Command).To(BeEmpty()) @@ -107,16 +103,16 @@ var _ = Describe("TranscodingRepository", func() { }) It("counts transcodings", func() { - count, err := repo.(*transcodingRepository).Count() + count, err := repo.Count(ctx) Expect(err).ToNot(HaveOccurred()) Expect(count).To(BeNumerically(">=", 0)) }) It("can still resolve a transcoding for streaming via Get (Command not redacted)", func() { tr := &model.Transcoding{ID: "streamreg", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg -secret"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) - res, err := repo.Get("streamreg") + res, err := repo.Get(ctx, "streamreg") Expect(err).ToNot(HaveOccurred()) Expect(res.ID).To(Equal("streamreg")) Expect(res.Command).To(Equal("ffmpeg -secret")) @@ -124,38 +120,34 @@ var _ = Describe("TranscodingRepository", func() { It("can still resolve a transcoding for streaming via FindByFormat (Command not redacted)", func() { tr := &model.Transcoding{ID: "fmtreg", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg -secret"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) - res, err := repo.FindByFormat("test_format") + res, err := repo.FindByFormat(ctx, "test_format") Expect(err).ToNot(HaveOccurred()) Expect(res.ID).To(Equal("fmtreg")) Expect(res.Command).To(Equal("ffmpeg -secret")) }) It("fails to create", func() { - err := repo.Put(&model.Transcoding{ID: "bad", Name: "bad", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg"}) + err := repo.Put(ctx, &model.Transcoding{ID: "bad", Name: "bad", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg"}) Expect(err).To(Equal(rest.ErrPermissionDenied)) }) It("fails to update", func() { tr := &model.Transcoding{ID: "updreg", Name: "old", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) tr.Name = "bad" - err := repo.Put(tr) + err := repo.Put(ctx, tr) Expect(err).To(Equal(rest.ErrPermissionDenied)) - - //_ = adminRepo.(*transcodingRepository).Delete("updreg") }) It("fails to delete", func() { tr := &model.Transcoding{ID: "delreg", Name: "temp", TargetFormat: "test_format", DefaultBitRate: 64, Command: "ffmpeg"} - Expect(adminRepo.Put(tr)).To(Succeed()) + Expect(repo.Put(adminCtx, tr)).To(Succeed()) - err := repo.(*transcodingRepository).Delete("delreg") + err := repo.Delete(ctx, "delreg") Expect(err).To(Equal(rest.ErrPermissionDenied)) - - //_ = adminRepo.(*transcodingRepository).Delete("delreg") }) }) }) diff --git a/persistence/user_props_repository.go b/persistence/user_props_repository.go index 9307385a2..59d7d332f 100644 --- a/persistence/user_props_repository.go +++ b/persistence/user_props_repository.go @@ -13,17 +13,16 @@ type userPropsRepository struct { sqlRepository } -func NewUserPropsRepository(ctx context.Context, db dbx.Builder) model.UserPropsRepository { +func NewUserPropsRepository(db dbx.Builder) model.UserPropsRepository { r := &userPropsRepository{} - r.ctx = ctx r.db = db r.tableName = "user_props" return r } -func (r userPropsRepository) Put(userId, key string, value string) error { +func (r userPropsRepository) Put(ctx context.Context, userId, key string, value string) error { update := Update(r.tableName).Set("value", value).Where(And{Eq{"user_id": userId}, Eq{"key": key}}) - count, err := r.executeSQL(update) + count, err := r.executeSQL(ctx, update) if err != nil { return err } @@ -31,24 +30,24 @@ func (r userPropsRepository) Put(userId, key string, value string) error { return nil } insert := Insert(r.tableName).Columns("user_id", "key", "value").Values(userId, key, value) - _, err = r.executeSQL(insert) + _, err = r.executeSQL(ctx, insert) return err } -func (r userPropsRepository) Get(userId, key string) (string, error) { +func (r userPropsRepository) Get(ctx context.Context, userId, key string) (string, error) { sel := Select("value").From(r.tableName).Where(And{Eq{"user_id": userId}, Eq{"key": key}}) resp := struct { Value string }{} - err := r.queryOne(sel, &resp) + err := r.queryOne(ctx, sel, &resp) if err != nil { return "", err } return resp.Value, nil } -func (r userPropsRepository) DefaultGet(userId, key string, defaultValue string) (string, error) { - value, err := r.Get(userId, key) +func (r userPropsRepository) DefaultGet(ctx context.Context, userId, key string, defaultValue string) (string, error) { + value, err := r.Get(ctx, userId, key) if errors.Is(err, model.ErrNotFound) { return defaultValue, nil } @@ -58,6 +57,6 @@ func (r userPropsRepository) DefaultGet(userId, key string, defaultValue string) return value, nil } -func (r userPropsRepository) Delete(userId, key string) error { - return r.delete(And{Eq{"user_id": userId}, Eq{"key": key}}) +func (r userPropsRepository) Delete(ctx context.Context, userId, key string) error { + return r.delete(ctx, And{Eq{"user_id": userId}, Eq{"key": key}}) } diff --git a/persistence/user_repository.go b/persistence/user_repository.go index 9399bcb53..20b4e5125 100644 --- a/persistence/user_repository.go +++ b/persistence/user_repository.go @@ -53,25 +53,24 @@ var ( encKey []byte ) -func NewUserRepository(ctx context.Context, db dbx.Builder) model.UserRepository { +func NewUserRepository(db dbx.Builder) model.UserRepository { r := &userRepository{} - r.ctx = ctx r.db = db r.tableName = "user" r.registerModel(&model.User{}, map[string]filterFunc{ "id": idFilter(r.tableName), - "password": invalidFilter(ctx), + "password": invalidFilter, "name": startsWithFilter(r.tableName + ".name"), }) once.Do(func() { - _ = r.initPasswordEncryptionKey() + _ = r.initPasswordEncryptionKey(context.Background()) }) return r } // selectUserWithLibraries returns a SelectBuilder that includes library information -func (r *userRepository) selectUserWithLibraries(options ...model.QueryOptions) SelectBuilder { - return r.newSelect(options...). +func (r *userRepository) selectUserWithLibraries(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + return r.newSelect(ctx, options...). Columns(`user.*`, `COALESCE(json_group_array(json_object( 'id', library.id, @@ -89,37 +88,37 @@ func (r *userRepository) selectUserWithLibraries(options ...model.QueryOptions) GroupBy("user.id") } -func (r *userRepository) CountAll(qo ...model.QueryOptions) (int64, error) { - return r.count(Select(), qo...) +func (r *userRepository) CountAll(ctx context.Context, qo ...model.QueryOptions) (int64, error) { + return r.count(ctx, Select(), qo...) } -func (r *userRepository) Get(id string) (*model.User, error) { - sel := r.selectUserWithLibraries().Where(Eq{"user.id": id}) +func (r *userRepository) Get(ctx context.Context, id string) (*model.User, error) { + sel := r.selectUserWithLibraries(ctx).Where(Eq{"user.id": id}) var res dbUser - err := r.queryOne(sel, &res) + err := r.queryOne(ctx, sel, &res) if err != nil { return nil, err } return res.User, nil } -func (r *userRepository) GetAll(options ...model.QueryOptions) (model.Users, error) { - sel := r.selectUserWithLibraries(options...) +func (r *userRepository) GetAll(ctx context.Context, options ...model.QueryOptions) (model.Users, error) { + sel := r.selectUserWithLibraries(ctx, options...) var res dbUsers - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) if err != nil { return nil, err } return res.toModels(), nil } -func (r *userRepository) Put(u *model.User) error { +func (r *userRepository) Put(ctx context.Context, u *model.User) error { if u.ID == "" { u.ID = id.NewRandom() } u.UpdatedAt = time.Now() if u.NewPassword != "" { - _ = r.encryptPassword(u) + _ = r.encryptPassword(ctx, u) } values, err := toSQLArgs(*u) if err != nil { @@ -134,7 +133,7 @@ func (r *userRepository) Put(u *model.User) error { var epoch int if u.NewPassword != "" { var res struct{ TokenEpoch int } - err = r.queryOne(update.Set("token_epoch", Expr("token_epoch + 1")). + err = r.queryOne(ctx, update.Set("token_epoch", Expr("token_epoch + 1")). Suffix("RETURNING token_epoch"), &res) switch { case errors.Is(err, model.ErrNotFound): @@ -145,7 +144,7 @@ func (r *userRepository) Put(u *model.User) error { epoch = res.TokenEpoch } } else { - count, err := r.executeSQL(update) + count, err := r.executeSQL(ctx, update) if err != nil { return err } @@ -154,7 +153,7 @@ func (r *userRepository) Put(u *model.User) error { if isNewUser { values["created_at"] = time.Now() insert := Insert(r.tableName).SetMap(values) - _, err = r.executeSQL(insert) + _, err = r.executeSQL(ctx, insert) if err != nil { return err } @@ -166,7 +165,7 @@ func (r *userRepository) Put(u *model.User) error { "INSERT OR IGNORE INTO user_library (user_id, library_id) SELECT ?, id FROM library", u.ID, ) - if _, err := r.executeSQL(sql); err != nil { + if _, err := r.executeSQL(ctx, sql); err != nil { return fmt.Errorf("failed to assign all libraries to admin user: %w", err) } } else if isNewUser { // Only for new regular users @@ -175,117 +174,108 @@ func (r *userRepository) Put(u *model.User) error { "INSERT OR IGNORE INTO user_library (user_id, library_id) SELECT ?, id FROM library WHERE default_new_users = true", u.ID, ) - if _, err := r.executeSQL(sql); err != nil { + if _, err := r.executeSQL(ctx, sql); err != nil { return fmt.Errorf("failed to assign default libraries to new user: %w", err) } } // Only the caller's own token can be refreshed in-flight; an admin resetting another // user must keep their own epoch. - if u.NewPassword != "" && !isNewUser && loggedUser(r.ctx).ID == u.ID { - request.SetTokenEpoch(r.ctx, epoch) + if u.NewPassword != "" && !isNewUser && loggedUser(ctx).ID == u.ID { + request.SetTokenEpoch(ctx, epoch) } return nil } -func (r *userRepository) FindFirstAdmin() (*model.User, error) { - sel := r.selectUserWithLibraries(model.QueryOptions{Sort: "updated_at", Max: 1}).Where(Eq{"user.is_admin": true}) +func (r *userRepository) FindFirstAdmin(ctx context.Context) (*model.User, error) { + sel := r.selectUserWithLibraries(ctx, model.QueryOptions{Sort: "updated_at", Max: 1}).Where(Eq{"user.is_admin": true}) var usr dbUser - err := r.queryOne(sel, &usr) + err := r.queryOne(ctx, sel, &usr) if err != nil { return nil, err } return usr.User, nil } -func (r *userRepository) FindByUsername(username string) (*model.User, error) { - sel := r.selectUserWithLibraries().Where(Expr("user.user_name = ? COLLATE NOCASE", username)) +func (r *userRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) { + sel := r.selectUserWithLibraries(ctx).Where(Expr("user.user_name = ? COLLATE NOCASE", username)) var usr dbUser - err := r.queryOne(sel, &usr) + err := r.queryOne(ctx, sel, &usr) if err != nil { return nil, err } return usr.User, nil } -func (r *userRepository) FindByUsernameWithPassword(username string) (*model.User, error) { - usr, err := r.FindByUsername(username) +func (r *userRepository) FindByUsernameWithPassword(ctx context.Context, username string) (*model.User, error) { + usr, err := r.FindByUsername(ctx, username) if err != nil { return nil, err } - _ = r.decryptPassword(usr) + _ = r.decryptPassword(ctx, usr) return usr, nil } -func (r *userRepository) UpdateLastLoginAt(id string) error { +func (r *userRepository) UpdateLastLoginAt(ctx context.Context, id string) error { upd := Update(r.tableName).Where(Eq{"id": id}).Set("last_login_at", time.Now()) - _, err := r.executeSQL(upd) + _, err := r.executeSQL(ctx, upd) return err } -func (r *userRepository) UpdateLastAccessAt(id string) error { +func (r *userRepository) UpdateLastAccessAt(ctx context.Context, id string) error { now := time.Now() upd := Update(r.tableName).Where(Eq{"id": id}).Set("last_access_at", now) - _, err := r.executeSQL(upd) + _, err := r.executeSQL(ctx, upd) return err } -func (r *userRepository) Count(options ...rest.QueryOptions) (int64, error) { - usr := loggedUser(r.ctx) +func (r *userRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + usr := loggedUser(ctx) if !usr.IsAdmin { return 0, rest.ErrPermissionDenied } - return r.CountAll(r.parseRestOptions(r.ctx, options...)) + return r.CountAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *userRepository) Read(id string) (any, error) { - usr := loggedUser(r.ctx) +func (r *userRepository) Read(ctx context.Context, id string) (*model.User, error) { + usr := loggedUser(ctx) if !usr.IsAdmin && usr.ID != id { return nil, rest.ErrPermissionDenied } - return r.Get(id) + return r.Get(ctx, id) } -func (r *userRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - usr := loggedUser(r.ctx) +func (r *userRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.User, error) { + usr := loggedUser(ctx) if !usr.IsAdmin { return nil, rest.ErrPermissionDenied } - return r.GetAll(r.parseRestOptions(r.ctx, options...)) + return r.GetAll(ctx, r.parseRestOptions(ctx, options...)) } -func (r *userRepository) EntityName() string { - return "user" -} - -func (r *userRepository) NewInstance() any { - return &model.User{} -} - -func (r *userRepository) Save(entity any) (string, error) { - usr := loggedUser(r.ctx) +func (r *userRepository) Save(ctx context.Context, u *model.User) (string, error) { + usr := loggedUser(ctx) if !usr.IsAdmin { return "", rest.ErrPermissionDenied } - u := entity.(*model.User) - if err := validateUsernameUnique(r, u); err != nil { + if err := validateUsernameUnique(ctx, r, u); err != nil { return "", err } if err := validateScrobbleFilter(u); err != nil { return "", err } - err := r.Put(u) + err := r.Put(ctx, u) if err != nil { return "", err } return u.ID, err } -func (r *userRepository) Update(id string, entity any, _ ...string) error { - u := entity.(*model.User) +func (r *userRepository) Update(ctx context.Context, id string, entity model.User, _ ...string) error { + u := &entity u.ID = id - usr := loggedUser(r.ctx) + usr := loggedUser(ctx) if !usr.IsAdmin && usr.ID != u.ID { return rest.ErrPermissionDenied } @@ -298,19 +288,19 @@ func (r *userRepository) Update(id string, entity any, _ ...string) error { } // Decrypt the user's existing password before validating. This is required otherwise the existing password entered by the user will never match. - if err := r.decryptPassword(usr); err != nil { + if err := r.decryptPassword(ctx, usr); err != nil { return err } if err := validatePasswordChange(u, usr); err != nil { return err } - if err := validateUsernameUnique(r, u); err != nil { + if err := validateUsernameUnique(ctx, r, u); err != nil { return err } if err := validateScrobbleFilter(u); err != nil { return err } - return r.Put(u) + return r.Put(ctx, u) } func validatePasswordChange(newUser *model.User, logged *model.User) error { @@ -339,8 +329,8 @@ func validatePasswordChange(newUser *model.User, logged *model.User) error { return nil } -func validateUsernameUnique(r model.UserRepository, u *model.User) error { - usr, err := r.FindByUsername(u.UserName) +func validateUsernameUnique(ctx context.Context, r model.UserRepository, u *model.User) error { + usr, err := r.FindByUsername(ctx, u.UserName) if errors.Is(err, model.ErrNotFound) { return nil } @@ -380,18 +370,18 @@ func invalidScrobbleFilter() error { }} } -func (r *userRepository) Delete(id string) error { - usr := loggedUser(r.ctx) +func (r *userRepository) Delete(ctx context.Context, ids ...string) error { + usr := loggedUser(ctx) if !usr.IsAdmin { return rest.ErrPermissionDenied } - if err := r.deleteByID(id); err != nil { - return err - } - - // Clean up orphaned plugin references for the deleted user - if err := cleanupPluginUserReferences(r.db, id); err != nil { - log.Error(r.ctx, "Failed to cleanup plugin user references", "userID", id, err) + for _, id := range ids { + if err := r.deleteByID(ctx, id); err != nil { + return err + } + if err := cleanupPluginUserReferences(r.db, id); err != nil { + log.Error(ctx, "Failed to cleanup plugin user references", "userID", id, err) + } } return nil } @@ -401,7 +391,7 @@ func keyTo32Bytes(input string) []byte { return data[0:] } -func (r *userRepository) initPasswordEncryptionKey() error { +func (r *userRepository) initPasswordEncryptionKey(ctx context.Context) error { encKey = keyTo32Bytes(consts.DefaultEncryptionKey) if conf.Server.PasswordEncryptionKey == "" { return nil @@ -410,8 +400,8 @@ func (r *userRepository) initPasswordEncryptionKey() error { key := keyTo32Bytes(conf.Server.PasswordEncryptionKey) keySum := fmt.Sprintf("%x", sha256.Sum256(key)) - props := NewPropertyRepository(r.ctx, r.db) - savedKeySum, err := props.Get(consts.PasswordsEncryptedKey) + props := NewPropertyRepository(r.db) + savedKeySum, err := props.Get(ctx, consts.PasswordsEncryptedKey) // If passwords are already encrypted if err == nil { @@ -425,24 +415,24 @@ func (r *userRepository) initPasswordEncryptionKey() error { // if not, try to re-encrypt all current passwords with new encryption key, // assuming they were encrypted with the DefaultEncryptionKey - sql := r.newSelect().Columns("id", "user_name", "password") + sql := r.newSelect(ctx).Columns("id", "user_name", "password") users := model.Users{} - err = r.queryAll(sql, &users) + err = r.queryAll(ctx, sql, &users) if err != nil { log.Error("Could not encrypt all passwords", err) return err } log.Warn("New PasswordEncryptionKey set. Encrypting all passwords", "numUsers", len(users)) - if err = r.decryptAllPasswords(users); err != nil { + if err = r.decryptAllPasswords(ctx, users); err != nil { return err } encKey = key for i := range users { u := users[i] u.NewPassword = u.Password - if err := r.encryptPassword(&u); err == nil { + if err := r.encryptPassword(ctx, &u); err == nil { upd := Update(r.tableName).Set("password", u.NewPassword).Where(Eq{"id": u.ID}) - _, err = r.executeSQL(upd) + _, err = r.executeSQL(ctx, upd) if err != nil { log.Error("Password NOT encrypted! This may cause problems!", "user", u.UserName, "id", u.ID, err) } else { @@ -451,7 +441,7 @@ func (r *userRepository) initPasswordEncryptionKey() error { } } - err = props.Put(consts.PasswordsEncryptedKey, keySum) + err = props.Put(ctx, consts.PasswordsEncryptedKey, keySum) if err != nil { log.Error("Could not flag passwords as encrypted. It will cause login errors", err) return err @@ -460,10 +450,10 @@ func (r *userRepository) initPasswordEncryptionKey() error { } // encrypts u.NewPassword -func (r *userRepository) encryptPassword(u *model.User) error { - encPassword, err := utils.Encrypt(r.ctx, encKey, u.NewPassword) +func (r *userRepository) encryptPassword(ctx context.Context, u *model.User) error { + encPassword, err := utils.Encrypt(ctx, encKey, u.NewPassword) if err != nil { - log.Error(r.ctx, "Error encrypting user's password", "user", u.UserName, err) + log.Error(ctx, "Error encrypting user's password", "user", u.UserName, err) return err } u.NewPassword = encPassword @@ -471,19 +461,19 @@ func (r *userRepository) encryptPassword(u *model.User) error { } // decrypts u.Password -func (r *userRepository) decryptPassword(u *model.User) error { - plaintext, err := utils.Decrypt(r.ctx, encKey, u.Password) +func (r *userRepository) decryptPassword(ctx context.Context, u *model.User) error { + plaintext, err := utils.Decrypt(ctx, encKey, u.Password) if err != nil { - log.Error(r.ctx, "Error decrypting user's password", "user", u.UserName, err) + log.Error(ctx, "Error decrypting user's password", "user", u.UserName, err) return err } u.Password = plaintext return nil } -func (r *userRepository) decryptAllPasswords(users model.Users) error { +func (r *userRepository) decryptAllPasswords(ctx context.Context, users model.Users) error { for i := range users { - if err := r.decryptPassword(&users[i]); err != nil { + if err := r.decryptPassword(ctx, &users[i]); err != nil { return err } } @@ -492,7 +482,7 @@ func (r *userRepository) decryptAllPasswords(users model.Users) error { // Library association methods -func (r *userRepository) GetUserLibraries(userID string) (model.Libraries, error) { +func (r *userRepository) GetUserLibraries(ctx context.Context, userID string) (model.Libraries, error) { sel := Select("l.*"). From("library l"). Join("user_library ul ON l.id = ul.library_id"). @@ -500,14 +490,14 @@ func (r *userRepository) GetUserLibraries(userID string) (model.Libraries, error OrderBy("l.name") var res model.Libraries - err := r.queryAll(sel, &res) + err := r.queryAll(ctx, sel, &res) return res, err } -func (r *userRepository) SetUserLibraries(userID string, libraryIDs []int) error { +func (r *userRepository) SetUserLibraries(ctx context.Context, userID string, libraryIDs []int) error { // Remove existing associations delSql := Delete("user_library").Where(Eq{"user_id": userID}) - if _, err := r.executeSQL(delSql); err != nil { + if _, err := r.executeSQL(ctx, delSql); err != nil { return err } @@ -517,12 +507,12 @@ func (r *userRepository) SetUserLibraries(userID string, libraryIDs []int) error for _, libID := range libraryIDs { insert = insert.Values(userID, libID) } - _, err := r.executeSQL(insert) + _, err := r.executeSQL(ctx, insert) return err } return nil } var _ model.UserRepository = (*userRepository)(nil) -var _ rest.Repository = (*userRepository)(nil) -var _ rest.Persistable = (*userRepository)(nil) +var _ rest.Repository[model.User] = (*userRepository)(nil) +var _ rest.Persistable[model.User] = (*userRepository)(nil) diff --git a/persistence/user_repository_test.go b/persistence/user_repository_test.go index 9a51b17dd..0e776fc3a 100644 --- a/persistence/user_repository_test.go +++ b/persistence/user_repository_test.go @@ -21,9 +21,11 @@ import ( var _ = Describe("UserRepository", func() { var repo model.UserRepository + var ctx context.Context BeforeEach(func() { - repo = NewUserRepository(log.NewContext(GinkgoT().Context()), GetDBXBuilder()) + ctx = log.NewContext(GinkgoT().Context()) + repo = NewUserRepository(GetDBXBuilder()) }) Describe("Put/Get/FindByUsername", func() { @@ -36,20 +38,20 @@ var _ = Describe("UserRepository", func() { IsAdmin: true, } It("saves the user to the DB", func() { - Expect(repo.Put(&usr)).To(BeNil()) + Expect(repo.Put(ctx, &usr)).To(BeNil()) }) It("returns the newly created user", func() { - actual, err := repo.Get("123") + actual, err := repo.Get(ctx, "123") Expect(err).ToNot(HaveOccurred()) Expect(actual.Name).To(Equal("Admin")) }) It("find the user by case-insensitive username", func() { - actual, err := repo.FindByUsername("aDmIn") + actual, err := repo.FindByUsername(ctx, "aDmIn") Expect(err).ToNot(HaveOccurred()) Expect(actual.Name).To(Equal("Admin")) }) It("find the user by username and decrypts the password", func() { - actual, err := repo.FindByUsernameWithPassword("aDmIn") + actual, err := repo.FindByUsernameWithPassword(ctx, "aDmIn") Expect(err).ToNot(HaveOccurred()) Expect(actual.Name).To(Equal("Admin")) Expect(actual.Password).To(Equal("wordpass")) @@ -57,27 +59,27 @@ var _ = Describe("UserRepository", func() { It("updates the name and keep the same password", func() { usr.Name = "Jane Doe" usr.NewPassword = "" - Expect(repo.Put(&usr)).To(BeNil()) + Expect(repo.Put(ctx, &usr)).To(BeNil()) - actual, err := repo.FindByUsernameWithPassword("admin") + actual, err := repo.FindByUsernameWithPassword(ctx, "admin") Expect(err).ToNot(HaveOccurred()) Expect(actual.Name).To(Equal("Jane Doe")) Expect(actual.Password).To(Equal("wordpass")) }) It("updates password if specified", func() { usr.NewPassword = "newpass" - Expect(repo.Put(&usr)).To(BeNil()) + Expect(repo.Put(ctx, &usr)).To(BeNil()) - actual, err := repo.FindByUsernameWithPassword("admin") + actual, err := repo.FindByUsernameWithPassword(ctx, "admin") Expect(err).ToNot(HaveOccurred()) Expect(actual.Password).To(Equal("newpass")) }) It("persists and reads back the scrobble filter", func() { usr := model.User{ID: "u-filter", UserName: "u-filter", Name: "Filter User", ScrobbleFilter: `{"all":[{"contains":{"title":"????"}}]}`} - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - saved, err := repo.Get("u-filter") + saved, err := repo.Get(ctx, "u-filter") Expect(err).ToNot(HaveOccurred()) Expect(saved.ScrobbleFilter).To(Equal(`{"all":[{"contains":{"title":"????"}}]}`)) }) @@ -88,7 +90,7 @@ var _ = Describe("UserRepository", func() { "values ('u-rawsql', 'u-rawsql', 'Raw', '', '', datetime('now'), datetime('now'))").Execute() Expect(err).ToNot(HaveOccurred()) - saved, err := repo.Get("u-rawsql") + saved, err := repo.Get(ctx, "u-rawsql") Expect(err).ToNot(HaveOccurred()) Expect(saved.ScrobbleFilter).To(Equal("")) }) @@ -233,36 +235,34 @@ var _ = Describe("UserRepository", func() { Describe("Delete", func() { It("returns not found for a missing user", func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) - adminRepo := NewUserRepository(adminCtx, GetDBXBuilder()).(*userRepository) - Expect(adminRepo.Delete("does-not-exist")).To(MatchError(model.ErrNotFound)) + adminRepo := NewUserRepository(GetDBXBuilder()).(*userRepository) + Expect(adminRepo.Delete(adminCtx, "does-not-exist")).To(MatchError(model.ErrNotFound)) }) }) Describe("ReadAll name filter", func() { - var adminRepo model.ResourceRepository + var adminRepo model.UserRepository + var adminCtx context.Context BeforeEach(func() { - adminCtx := request.WithUser(GinkgoT().Context(), model.User{ID: "admin-id", UserName: "admin", IsAdmin: true}) - adminRepo = NewUserRepository(adminCtx, GetDBXBuilder()).(model.ResourceRepository) + adminCtx = request.WithUser(ctx, model.User{ID: "admin-id", UserName: "admin", IsAdmin: true}) + adminRepo = NewUserRepository(GetDBXBuilder()) for _, u := range []model.User{ {ID: "filter-alice", UserName: "alice_filter", Name: "Alice Filter", NewPassword: "x"}, {ID: "filter-bob", UserName: "bob_filter", Name: "Bob Filter", NewPassword: "x"}, } { - Expect(adminRepo.(model.UserRepository).Put(&u)).To(Succeed()) + Expect(adminRepo.Put(adminCtx, &u)).To(Succeed()) } }) AfterEach(func() { - ur := adminRepo.(model.UserRepository) - _ = ur.Delete("filter-alice") - _ = ur.Delete("filter-bob") + _ = adminRepo.Delete(adminCtx, "filter-alice", "filter-bob") }) It("matches users whose name starts with the given prefix", func() { - res, err := adminRepo.ReadAll(rest.QueryOptions{Filters: map[string]any{"name": "Alice"}}) + users, err := adminRepo.ReadAll(adminCtx, rest.QueryOptions{Filters: map[string]any{"name": "Alice"}}) Expect(err).ToNot(HaveOccurred()) - users := res.(model.Users) var names []string for _, u := range users { @@ -273,9 +273,8 @@ var _ = Describe("UserRepository", func() { }) It("does not match names by mid-string substring (startsWith, not contains)", func() { - res, err := adminRepo.ReadAll(rest.QueryOptions{Filters: map[string]any{"name": "Filter"}}) + users, err := adminRepo.ReadAll(adminCtx, rest.QueryOptions{Filters: map[string]any{"name": "Filter"}}) Expect(err).ToNot(HaveOccurred()) - users := res.(model.Users) for _, u := range users { Expect(u.ID).ToNot(Or(Equal("filter-alice"), Equal("filter-bob")), @@ -290,17 +289,17 @@ var _ = Describe("UserRepository", func() { BeforeEach(func() { existingUser = &model.User{ID: "1", UserName: "johndoe"} repo = tests.CreateMockUserRepo() - err := repo.Put(existingUser) + err := repo.Put(ctx, existingUser) Expect(err).ToNot(HaveOccurred()) }) It("allows unique usernames", func() { var newUser = &model.User{ID: "2", UserName: "unique_username"} - err := validateUsernameUnique(repo, newUser) + err := validateUsernameUnique(ctx, repo, newUser) Expect(err).ToNot(HaveOccurred()) }) It("returns ValidationError if username already exists", func() { var newUser = &model.User{ID: "2", UserName: "johndoe"} - err := validateUsernameUnique(repo, newUser) + err := validateUsernameUnique(ctx, repo, newUser) var verr *rest.ValidationError isValidationError := errors.As(err, &verr) @@ -311,7 +310,7 @@ var _ = Describe("UserRepository", func() { repo.Error = errors.New("fake error") var newUser = &model.User{ID: "2", UserName: "newuser"} - err := validateUsernameUnique(repo, newUser) + err := validateUsernameUnique(ctx, repo, newUser) Expect(err).To(MatchError("fake error")) }) }) @@ -330,39 +329,39 @@ var _ = Describe("UserRepository", func() { NewPassword: "password", IsAdmin: false, } - Expect(repo.Put(&testUser)).To(BeNil()) + Expect(repo.Put(ctx, &testUser)).To(BeNil()) userID = testUser.ID library1 = model.Library{ID: 0, Name: "Library 500", Path: "/path/500"} library2 = model.Library{ID: 0, Name: "Library 501", Path: "/path/501"} // Create test libraries - libRepo := NewLibraryRepository(log.NewContext(context.TODO()), GetDBXBuilder()) - Expect(libRepo.Put(&library1)).To(BeNil()) - Expect(libRepo.Put(&library2)).To(BeNil()) + libRepo := NewLibraryRepository(GetDBXBuilder()) + Expect(libRepo.Put(ctx, &library1)).To(BeNil()) + Expect(libRepo.Put(ctx, &library2)).To(BeNil()) }) AfterEach(func() { // Clean up user-library associations to ensure test isolation - _ = repo.SetUserLibraries(userID, []int{}) + _ = repo.SetUserLibraries(ctx, userID, []int{}) // Clean up test libraries to ensure isolation between test groups - libRepo := NewLibraryRepository(log.NewContext(context.TODO()), GetDBXBuilder()) - _ = libRepo.(*libraryRepository).delete(squirrel.Eq{"id": []int{library1.ID, library2.ID}}) + libRepo := NewLibraryRepository(GetDBXBuilder()) + _ = libRepo.(*libraryRepository).delete(ctx, squirrel.Eq{"id": []int{library1.ID, library2.ID}}) }) Describe("GetUserLibraries", func() { It("returns empty list when user has no library associations", func() { - libraries, err := repo.GetUserLibraries("non-existent-user") + libraries, err := repo.GetUserLibraries(ctx, "non-existent-user") Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(0)) }) It("returns user's associated libraries", func() { - err := repo.SetUserLibraries(userID, []int{library1.ID, library2.ID}) + err := repo.SetUserLibraries(ctx, userID, []int{library1.ID, library2.ID}) Expect(err).ToNot(HaveOccurred()) - libraries, err := repo.GetUserLibraries(userID) + libraries, err := repo.GetUserLibraries(ctx, userID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(2)) @@ -374,24 +373,24 @@ var _ = Describe("UserRepository", func() { Describe("SetUserLibraries", func() { It("sets user's library associations", func() { libraryIDs := []int{library1.ID, library2.ID} - err := repo.SetUserLibraries(userID, libraryIDs) + err := repo.SetUserLibraries(ctx, userID, libraryIDs) Expect(err).ToNot(HaveOccurred()) - libraries, err := repo.GetUserLibraries(userID) + libraries, err := repo.GetUserLibraries(ctx, userID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(2)) }) It("replaces existing associations", func() { // Set initial associations - err := repo.SetUserLibraries(userID, []int{library1.ID, library2.ID}) + err := repo.SetUserLibraries(ctx, userID, []int{library1.ID, library2.ID}) Expect(err).ToNot(HaveOccurred()) // Replace with just one library - err = repo.SetUserLibraries(userID, []int{library1.ID}) + err = repo.SetUserLibraries(ctx, userID, []int{library1.ID}) Expect(err).ToNot(HaveOccurred()) - libraries, err := repo.GetUserLibraries(userID) + libraries, err := repo.GetUserLibraries(ctx, userID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(1)) Expect(libraries[0].ID).To(Equal(library1.ID)) @@ -399,14 +398,14 @@ var _ = Describe("UserRepository", func() { It("removes all associations when passed empty slice", func() { // Set initial associations - err := repo.SetUserLibraries(userID, []int{library1.ID, library2.ID}) + err := repo.SetUserLibraries(ctx, userID, []int{library1.ID, library2.ID}) Expect(err).ToNot(HaveOccurred()) // Remove all - err = repo.SetUserLibraries(userID, []int{}) + err = repo.SetUserLibraries(ctx, userID, []int{}) Expect(err).ToNot(HaveOccurred()) - libraries, err := repo.GetUserLibraries(userID) + libraries, err := repo.GetUserLibraries(ctx, userID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(0)) }) @@ -422,10 +421,10 @@ var _ = Describe("UserRepository", func() { ) BeforeEach(func() { - libRepo = NewLibraryRepository(log.NewContext(context.TODO()), GetDBXBuilder()) + libRepo = NewLibraryRepository(GetDBXBuilder()) // Count initial libraries - existingLibs, err := libRepo.GetAll() + existingLibs, err := libRepo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) initialLibCount = len(existingLibs) @@ -433,16 +432,16 @@ var _ = Describe("UserRepository", func() { library2 = model.Library{ID: 0, Name: "Admin Test Library 2", Path: "/admin/test/path2"} // Create test libraries - Expect(libRepo.Put(&library1)).To(BeNil()) - Expect(libRepo.Put(&library2)).To(BeNil()) + Expect(libRepo.Put(ctx, &library1)).To(BeNil()) + Expect(libRepo.Put(ctx, &library2)).To(BeNil()) }) AfterEach(func() { // Clean up test libraries and their associations - _ = libRepo.(*libraryRepository).delete(squirrel.Eq{"id": []int{library1.ID, library2.ID}}) + _ = libRepo.(*libraryRepository).delete(ctx, squirrel.Eq{"id": []int{library1.ID, library2.ID}}) // Clean up user-library associations for these test libraries - _, _ = repo.(*userRepository).executeSQL(squirrel.Delete("user_library").Where(squirrel.Eq{"library_id": []int{library1.ID, library2.ID}})) + _, _ = repo.(*userRepository).executeSQL(ctx, squirrel.Delete("user_library").Where(squirrel.Eq{"library_id": []int{library1.ID, library2.ID}})) }) It("automatically assigns all libraries to admin users when created", func() { @@ -455,11 +454,11 @@ var _ = Describe("UserRepository", func() { IsAdmin: true, } - err := repo.Put(&adminUser) + err := repo.Put(ctx, &adminUser) Expect(err).ToNot(HaveOccurred()) // Admin should automatically have access to all libraries (including existing ones) - libraries, err := repo.GetUserLibraries(adminUser.ID) + libraries, err := repo.GetUserLibraries(ctx, adminUser.ID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(initialLibCount + 2)) // Initial libraries + our 2 test libraries @@ -481,20 +480,20 @@ var _ = Describe("UserRepository", func() { IsAdmin: false, } - err := repo.Put(®ularUser) + err := repo.Put(ctx, ®ularUser) Expect(err).ToNot(HaveOccurred()) // Give them access to just one library - err = repo.SetUserLibraries(regularUser.ID, []int{library1.ID}) + err = repo.SetUserLibraries(ctx, regularUser.ID, []int{library1.ID}) Expect(err).ToNot(HaveOccurred()) // Promote to admin regularUser.IsAdmin = true - err = repo.Put(®ularUser) + err = repo.Put(ctx, ®ularUser) Expect(err).ToNot(HaveOccurred()) // Should now have access to all libraries (including existing ones) - libraries, err := repo.GetUserLibraries(regularUser.ID) + libraries, err := repo.GetUserLibraries(ctx, regularUser.ID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(initialLibCount + 2)) // Initial libraries + our 2 test libraries @@ -516,11 +515,11 @@ var _ = Describe("UserRepository", func() { IsAdmin: false, } - err := repo.Put(®ularUser) + err := repo.Put(ctx, ®ularUser) Expect(err).ToNot(HaveOccurred()) // Regular user should be assigned to default libraries (library ID 1 from migration) - libraries, err := repo.GetUserLibraries(regularUser.ID) + libraries, err := repo.GetUserLibraries(ctx, regularUser.ID) Expect(err).ToNot(HaveOccurred()) Expect(libraries).To(HaveLen(1)) Expect(libraries[0].ID).To(Equal(1)) @@ -537,13 +536,13 @@ var _ = Describe("UserRepository", func() { ) BeforeEach(func() { - libRepo = NewLibraryRepository(log.NewContext(context.TODO()), GetDBXBuilder()) + libRepo = NewLibraryRepository(GetDBXBuilder()) library1 = model.Library{ID: 0, Name: "Field Test Library 1", Path: "/field/test/path1"} library2 = model.Library{ID: 0, Name: "Field Test Library 2", Path: "/field/test/path2"} // Create test libraries - Expect(libRepo.Put(&library1)).To(BeNil()) - Expect(libRepo.Put(&library2)).To(BeNil()) + Expect(libRepo.Put(ctx, &library1)).To(BeNil()) + Expect(libRepo.Put(ctx, &library2)).To(BeNil()) // Create test user testUser = model.User{ @@ -554,23 +553,23 @@ var _ = Describe("UserRepository", func() { NewPassword: "password", IsAdmin: false, } - Expect(repo.Put(&testUser)).To(BeNil()) + Expect(repo.Put(ctx, &testUser)).To(BeNil()) // Assign libraries to user - Expect(repo.SetUserLibraries(testUser.ID, []int{library1.ID, library2.ID})).To(BeNil()) + Expect(repo.SetUserLibraries(ctx, testUser.ID, []int{library1.ID, library2.ID})).To(BeNil()) }) AfterEach(func() { // Clean up test libraries and their associations - _ = libRepo.(*libraryRepository).delete(squirrel.Eq{"id": []int{library1.ID, library2.ID}}) - _ = repo.(*userRepository).delete(squirrel.Eq{"id": testUser.ID}) + _ = libRepo.(*libraryRepository).delete(ctx, squirrel.Eq{"id": []int{library1.ID, library2.ID}}) + _ = repo.(*userRepository).delete(ctx, squirrel.Eq{"id": testUser.ID}) // Clean up user-library associations for these test libraries - _, _ = repo.(*userRepository).executeSQL(squirrel.Delete("user_library").Where(squirrel.Eq{"library_id": []int{library1.ID, library2.ID}})) + _, _ = repo.(*userRepository).executeSQL(ctx, squirrel.Delete("user_library").Where(squirrel.Eq{"library_id": []int{library1.ID, library2.ID}})) }) It("populates Libraries field when getting a single user", func() { - user, err := repo.Get(testUser.ID) + user, err := repo.Get(ctx, testUser.ID) Expect(err).ToNot(HaveOccurred()) Expect(user.Libraries).To(HaveLen(2)) @@ -591,7 +590,7 @@ var _ = Describe("UserRepository", func() { }) It("populates Libraries field when getting all users", func() { - users, err := repo.(*userRepository).GetAll() + users, err := repo.(*userRepository).GetAll(ctx) Expect(err).ToNot(HaveOccurred()) // Find our test user in the results @@ -607,7 +606,7 @@ var _ = Describe("UserRepository", func() { }) It("populates Libraries field when finding user by username", func() { - user, err := repo.FindByUsername(testUser.UserName) + user, err := repo.FindByUsername(ctx, testUser.UserName) Expect(err).ToNot(HaveOccurred()) Expect(user.Libraries).To(HaveLen(2)) @@ -625,10 +624,10 @@ var _ = Describe("UserRepository", func() { NewPassword: "password", IsAdmin: false, } - Expect(repo.Put(&userWithoutLibs)).To(BeNil()) - defer func() { _ = repo.(*userRepository).delete(squirrel.Eq{"id": userWithoutLibs.ID}) }() + Expect(repo.Put(ctx, &userWithoutLibs)).To(BeNil()) + defer func() { _ = repo.(*userRepository).delete(ctx, squirrel.Eq{"id": userWithoutLibs.ID}) }() - user, err := repo.Get(userWithoutLibs.ID) + user, err := repo.Get(ctx, userWithoutLibs.ID) Expect(err).ToNot(HaveOccurred()) Expect(user.Libraries).ToNot(BeNil()) // Regular users should be assigned to default libraries (library ID 1 from migration) @@ -686,8 +685,8 @@ var _ = Describe("UserRepository", func() { Describe("filters", func() { It("qualifies id filter with table name", func() { r := repo.(*userRepository) - qo := r.parseRestOptions(r.ctx, rest.QueryOptions{Filters: map[string]any{"id": "123"}}) - sel := r.selectUserWithLibraries(qo) + qo := r.parseRestOptions(ctx, rest.QueryOptions{Filters: map[string]any{"id": "123"}}) + sel := r.selectUserWithLibraries(ctx, qo) query, _, err := r.toSQL(sel) Expect(err).NotTo(HaveOccurred()) Expect(query).To(ContainSubstring("user.id = {:p0}")) @@ -705,29 +704,28 @@ var _ = Describe("UserRepository", func() { } BeforeEach(func() { - ctx := log.NewContext(context.TODO()) ctx = request.WithUser(ctx, model.User{ID: "userid", IsAdmin: true}) - repo = NewUserRepository(ctx, GetDBXBuilder()) + repo = NewUserRepository(GetDBXBuilder()) usr = newUser() - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) }) It("starts at zero for a new user", func() { - got, err := repo.Get(usr.ID) + got, err := repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(0)) }) It("increments once per password change", func() { usr.NewPassword = "second" - Expect(repo.Put(&usr)).To(Succeed()) - got, err := repo.Get(usr.ID) + Expect(repo.Put(ctx, &usr)).To(Succeed()) + got, err := repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(1)) usr.NewPassword = "third" - Expect(repo.Put(&usr)).To(Succeed()) - got, err = repo.Get(usr.ID) + Expect(repo.Put(ctx, &usr)).To(Succeed()) + got, err = repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(2)) }) @@ -735,9 +733,9 @@ var _ = Describe("UserRepository", func() { It("leaves the epoch alone when the password is untouched", func() { usr.NewPassword = "" usr.Name = "Renamed" - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - got, err := repo.Get(usr.ID) + got, err := repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(0)) Expect(got.Name).To(Equal("Renamed")) @@ -754,11 +752,11 @@ var _ = Describe("UserRepository", func() { ctx := log.NewContext(context.TODO()) ctx = request.WithUser(ctx, model.User{ID: usr.ID}) ctx = request.WithTokenEpochHolder(ctx) - own := NewUserRepository(ctx, GetDBXBuilder()) + own := NewUserRepository(GetDBXBuilder()) u := usr u.NewPassword = "concurrent" - if err := own.Put(&u); err != nil { + if err := own.Put(ctx, &u); err != nil { return // the shared in-memory test DB can raise SQLITE_LOCKED } epoch, ok := request.TokenEpochFrom(ctx) @@ -778,73 +776,73 @@ var _ = Describe("UserRepository", func() { }) Describe("Put and the token epoch", func() { - newRepo := func(actingUserID string) model.UserRepository { + newRepo := func(actingUserID string) (context.Context, model.UserRepository) { ctx := log.NewContext(context.TODO()) ctx = request.WithUser(ctx, model.User{ID: actingUserID, IsAdmin: true}) ctx = request.WithTokenEpochHolder(ctx) - return NewUserRepository(ctx, GetDBXBuilder()) + return ctx, NewUserRepository(GetDBXBuilder()) } It("does not bump when creating a user", func() { - repo := newRepo("admin") + ctx, repo := newRepo("admin") usr := model.User{ID: id.NewRandom(), UserName: "fresh", NewPassword: "pw1"} - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - got, err := repo.Get(usr.ID) + got, err := repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(0)) }) It("bumps when the password changes", func() { - repo := newRepo("admin") + ctx, repo := newRepo("admin") usr := model.User{ID: id.NewRandom(), UserName: "changer", NewPassword: "pw1"} - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) usr.NewPassword = "pw2" - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - got, err := repo.Get(usr.ID) + got, err := repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(1)) }) It("does not bump on an edit that leaves the password alone", func() { - repo := newRepo("admin") + ctx, repo := newRepo("admin") usr := model.User{ID: id.NewRandom(), UserName: "renamer", NewPassword: "pw1"} - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) usr.NewPassword = "" usr.Name = "New Display Name" - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - got, err := repo.Get(usr.ID) + got, err := repo.Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(got.TokenEpoch).To(Equal(0)) }) It("signals the new epoch when a user changes their own password", func() { userID := id.NewRandom() - repo := newRepo(userID) + ctx, repo := newRepo(userID) usr := model.User{ID: userID, UserName: "self", NewPassword: "pw1"} - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) usr.NewPassword = "pw2" - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - epoch, ok := request.TokenEpochFrom(repo.(*userRepository).ctx) + epoch, ok := request.TokenEpochFrom(ctx) Expect(ok).To(BeTrue()) Expect(epoch).To(Equal(1)) }) It("does not signal when an admin changes someone else's password", func() { - repo := newRepo("some-admin") + ctx, repo := newRepo("some-admin") usr := model.User{ID: id.NewRandom(), UserName: "other", NewPassword: "pw1"} - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) usr.NewPassword = "pw2" - Expect(repo.Put(&usr)).To(Succeed()) + Expect(repo.Put(ctx, &usr)).To(Succeed()) - _, ok := request.TokenEpochFrom(repo.(*userRepository).ctx) + _, ok := request.TokenEpochFrom(ctx) Expect(ok).To(BeFalse()) }) }) diff --git a/plugins/host_library.go b/plugins/host_library.go index 3d9f61b4f..25c14a813 100644 --- a/plugins/host_library.go +++ b/plugins/host_library.go @@ -37,7 +37,7 @@ func (s *libraryServiceImpl) GetLibrary(ctx context.Context, id int32) (*host.Li return nil, fmt.Errorf("library not accessible: library ID %d is not in the allowed list", id) } - lib, err := s.ds.Library(ctx).Get(int(id)) + lib, err := s.ds.Library().Get(ctx, int(id)) if err != nil { return nil, fmt.Errorf("library not found: %w", err) } @@ -55,7 +55,7 @@ func (s *libraryServiceImpl) isLibraryAccessible(id int) bool { } func (s *libraryServiceImpl) GetAllLibraries(ctx context.Context) ([]host.Library, error) { - libs, err := s.ds.Library(ctx).GetAll() + libs, err := s.ds.Library().GetAll(ctx) if err != nil { return nil, fmt.Errorf("failed to get libraries: %w", err) } diff --git a/plugins/host_library_test.go b/plugins/host_library_test.go index 00a953b24..edd4b546f 100644 --- a/plugins/host_library_test.go +++ b/plugins/host_library_test.go @@ -47,7 +47,7 @@ var _ = Describe("LibraryService", Ordered, func() { } lib.LastScanAt = lib.LastScanAt.Add(0) // Ensure time is set - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(model.Libraries{*lib}) result, err := service.GetLibrary(ctx, 1) @@ -77,7 +77,7 @@ var _ = Describe("LibraryService", Ordered, func() { TotalDuration: 1800.0, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(model.Libraries{*lib}) result, err := service.GetLibrary(ctx, 2) @@ -91,7 +91,7 @@ var _ = Describe("LibraryService", Ordered, func() { It("should return error for non-existent library", func() { service = newLibraryService(ds, &LibraryPermission{Reason: new("test")}, nil, true).(*libraryServiceImpl) - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(model.Libraries{}) _, err := service.GetLibrary(ctx, 999) @@ -109,7 +109,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 2, Name: "Jazz", Path: "/music/jazz", TotalSongs: 50}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) results, err := service.GetAllLibraries(ctx) @@ -131,7 +131,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 2, Name: "Jazz", Path: "/music/jazz", TotalSongs: 50}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) results, err := service.GetAllLibraries(ctx) @@ -154,7 +154,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 3, Name: "Classical", Path: "/music/classical", TotalSongs: 75}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) results, err := service.GetAllLibraries(ctx) @@ -172,7 +172,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 2, Name: "Jazz", Path: "/music/jazz", TotalSongs: 50}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) // Requesting library 1 which is not in the allowed list @@ -189,7 +189,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 2, Name: "Jazz", Path: "/music/jazz", TotalSongs: 50}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) result, err := service.GetLibrary(ctx, 2) @@ -206,7 +206,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 2, Name: "Jazz", Path: "/music/jazz", TotalSongs: 50}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) results, err := service.GetAllLibraries(ctx) @@ -222,7 +222,7 @@ var _ = Describe("LibraryService", Ordered, func() { {ID: 2, Name: "Jazz", Path: "/music/jazz", TotalSongs: 50}, } - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(libs) results, err := service.GetAllLibraries(ctx) @@ -291,7 +291,7 @@ var _ = Describe("LibraryService", Ordered, func() { Expect(manager.ds).ToNot(BeNil()) ctx := context.Background() - libs, err := manager.ds.Library(adminContext(ctx)).GetAll() + libs, err := manager.ds.Library().GetAll(adminContext(ctx)) Expect(err).ToNot(HaveOccurred()) Expect(libs).To(HaveLen(1)) Expect(libs[0].Path).To(Equal("/tmp/test-music")) diff --git a/plugins/host_matcher_test.go b/plugins/host_matcher_test.go index f0967628c..7d3e47fae 100644 --- a/plugins/host_matcher_test.go +++ b/plugins/host_matcher_test.go @@ -193,7 +193,7 @@ var _ = Describe("MatcherService", Ordered, func() { mediaFileRepo.SetData(model.MediaFiles{mf}) userRepo = tests.CreateMockUserRepo() - Expect(userRepo.Put(&model.User{ID: "u-alice", UserName: "alice"})).To(Succeed()) + Expect(userRepo.Put(GinkgoT().Context(), &model.User{ID: "u-alice", UserName: "alice"})).To(Succeed()) ds = &tests.MockDataStore{MockedMediaFile: mediaFileRepo, MockedUser: userRepo} }) @@ -269,7 +269,7 @@ var _ = Describe("MatcherService", Ordered, func() { _, err := svc.MatchSongs(callerCtx, input, host.MatchOptions{}) Expect(err).ToNot(HaveOccurred()) - usr, ok := request.UserFrom(capturing.lastMediaFileCtx) + usr, ok := request.UserFrom(capturing.lastMediaFileCtx()) Expect(ok).To(BeTrue()) Expect(usr.IsAdmin).To(BeTrue()) Expect(usr.ID).ToNot(Equal("u-caller")) @@ -283,7 +283,7 @@ var _ = Describe("MatcherService", Ordered, func() { _, err := svc.MatchSongs(callerCtx, input, host.MatchOptions{Username: "alice"}) Expect(err).ToNot(HaveOccurred()) - usr, ok := request.UserFrom(capturing.lastMediaFileCtx) + usr, ok := request.UserFrom(capturing.lastMediaFileCtx()) Expect(ok).To(BeTrue()) Expect(usr.ID).To(Equal("u-alice")) }) @@ -395,7 +395,7 @@ var _ = Describe("MatcherService Integration", Ordered, func() { mediaFileRepo.SetData(model.MediaFiles{hit}) userRepo := tests.CreateMockUserRepo() - Expect(userRepo.Put(&model.User{ID: "u-alice", UserName: "alice"})).To(Succeed()) + Expect(userRepo.Put(GinkgoT().Context(), &model.User{ID: "u-alice", UserName: "alice"})).To(Succeed()) dataStore := &tests.MockDataStore{ MockedPlugin: mockPluginRepo, @@ -475,14 +475,38 @@ var _ = Describe("MatcherService Integration", Ordered, func() { }) }) -// ctxCapturingDataStore records the context passed to MediaFile so tests can assert -// which user the matcher resolved before querying the library. +// ctxCapturingDataStore records the context the media file queries run with, so tests +// can assert which user the matcher resolved before querying the library. type ctxCapturingDataStore struct { *tests.MockDataStore - lastMediaFileCtx context.Context + repo *ctxCapturingMediaFileRepo } -func (d *ctxCapturingDataStore) MediaFile(ctx context.Context) model.MediaFileRepository { - d.lastMediaFileCtx = ctx - return d.MockDataStore.MediaFile(ctx) +func (d *ctxCapturingDataStore) MediaFile() model.MediaFileRepository { + if d.repo == nil { + d.repo = &ctxCapturingMediaFileRepo{MediaFileRepository: d.MockDataStore.MediaFile()} + } + return d.repo +} + +func (d *ctxCapturingDataStore) lastMediaFileCtx() context.Context { + if d.repo == nil { + return nil + } + return d.repo.lastCtx +} + +type ctxCapturingMediaFileRepo struct { + model.MediaFileRepository + lastCtx context.Context +} + +func (r *ctxCapturingMediaFileRepo) GetAll(ctx context.Context, options ...model.QueryOptions) (model.MediaFiles, error) { + r.lastCtx = ctx + return r.MediaFileRepository.GetAll(ctx, options...) +} + +func (r *ctxCapturingMediaFileRepo) GetAllByTags(ctx context.Context, tag model.TagName, values []string, options ...model.QueryOptions) (model.MediaFiles, error) { + r.lastCtx = ctx + return r.MediaFileRepository.GetAllByTags(ctx, tag, values, options...) } diff --git a/plugins/host_scrobbleretriever.go b/plugins/host_scrobbleretriever.go index 7417d7c50..780b0274d 100644 --- a/plugins/host_scrobbleretriever.go +++ b/plugins/host_scrobbleretriever.go @@ -41,7 +41,7 @@ func (s *scrobbleRetrieverServiceImpl) getFirstLastScrobble(ctx context.Context, return nil, err } - scrobbles, err := s.ds.Scrobble(ctx).GetAll(model.QueryOptions{Sort: "submission_time", Order: order, Max: 1}) + scrobbles, err := s.ds.Scrobble().GetAll(ctx, model.QueryOptions{Sort: "submission_time", Order: order, Max: 1}) if err != nil { return nil, err } @@ -80,7 +80,7 @@ func (s *scrobbleRetrieverServiceImpl) GetScrobbles(ctx context.Context, usernam // Fetch one more item than requested. The last item is the next timestamp to fetch lookahead := options.MaxItems + 1 - scrobbles, err := s.ds.Scrobble(ctx).GetAll(model.QueryOptions{ + scrobbles, err := s.ds.Scrobble().GetAll(ctx, model.QueryOptions{ Max: lookahead, Filters: scrobbleRangeFilters(options.FromTimestamp, options.ToTimestamp), // The id tiebreak makes the order of equal timestamps stable, which is what @@ -142,7 +142,7 @@ func (s *scrobbleRetrieverServiceImpl) GetScrobbleCount(ctx context.Context, use return 0, err } - return s.ds.Scrobble(ctx).CountAll(model.QueryOptions{ + return s.ds.Scrobble().CountAll(ctx, model.QueryOptions{ Filters: scrobbleRangeFilters(options.FromTimestamp, options.ToTimestamp), }) } diff --git a/plugins/host_scrobbleretriever_test.go b/plugins/host_scrobbleretriever_test.go index aa92c9eb5..6721c1487 100644 --- a/plugins/host_scrobbleretriever_test.go +++ b/plugins/host_scrobbleretriever_test.go @@ -75,36 +75,36 @@ var _ = Describe("Scrobble Retriever Host Function", Ordered, func() { conf.Server.Plugins.Folder = conf.NewDir(tmpDir) conf.Server.Plugins.AutoReload = false - userRepo := dataStore.User(ctx) + userRepo := dataStore.User() // Add test users - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(ctx, &model.User{ ID: "user1", UserName: "testuser", IsAdmin: false, }) - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(ctx, &model.User{ ID: "admin1", UserName: "adminuser", IsAdmin: true, }) - err = dataStore.MediaFile(ctx).Put(&model.MediaFile{ID: "1", LibraryID: 1}) + err = dataStore.MediaFile().Put(ctx, &model.MediaFile{ID: "1", LibraryID: 1}) Expect(err).To(BeNil()) - err = dataStore.MediaFile(ctx).Put(&model.MediaFile{ID: "2", LibraryID: 1}) + err = dataStore.MediaFile().Put(ctx, &model.MediaFile{ID: "2", LibraryID: 1}) Expect(err).To(BeNil()) - err = dataStore.MediaFile(ctx).Put(&model.MediaFile{ID: "3", LibraryID: 1}) + err = dataStore.MediaFile().Put(ctx, &model.MediaFile{ID: "3", LibraryID: 1}) Expect(err).To(BeNil()) scrobbleCtx := request.WithUser(GinkgoT().Context(), model.User{ID: "admin1", UserName: "adminuser"}) - scrobbleRepo := dataStore.Scrobble(scrobbleCtx) - err = scrobbleRepo.RecordScrobble("1", time.Unix(0, 0)) + scrobbleRepo := dataStore.Scrobble() + err = scrobbleRepo.RecordScrobble(scrobbleCtx, "1", time.Unix(0, 0)) Expect(err).To(BeNil()) - err = scrobbleRepo.RecordScrobble("2", time.Unix(1, 0)) + err = scrobbleRepo.RecordScrobble(scrobbleCtx, "2", time.Unix(1, 0)) Expect(err).To(BeNil()) - err = scrobbleRepo.RecordScrobble("3", time.Unix(2, 0)) + err = scrobbleRepo.RecordScrobble(scrobbleCtx, "3", time.Unix(2, 0)) Expect(err).To(BeNil()) - err = scrobbleRepo.RecordScrobble("1", time.Unix(2, 0)) + err = scrobbleRepo.RecordScrobble(scrobbleCtx, "1", time.Unix(2, 0)) Expect(err).To(BeNil()) // Create and configure manager @@ -125,7 +125,7 @@ var _ = Describe("Scrobble Retriever Host Function", Ordered, func() { dataStore.MockedPlugin = tests.CreateMockPluginRepo() - mockPluginRepo := dataStore.Plugin(GinkgoT().Context()).(*tests.MockPluginRepo) + mockPluginRepo := dataStore.Plugin().(*tests.MockPluginRepo) mockPluginRepo.Permitted = true enabledPlugin := model.Plugin{ ID: "test-scrobble-retriever", @@ -300,10 +300,10 @@ var _ = Describe("Scrobble Retriever Host Function", Ordered, func() { BeforeAll(func() { scrobbleCtx := request.WithUser(GinkgoT().Context(), model.User{ID: "admin1", UserName: "adminuser"}) - scrobbleRepo := dataStore.Scrobble(scrobbleCtx) + scrobbleRepo := dataStore.Scrobble() for i := range 5 { - err := scrobbleRepo.RecordScrobble("3", time.Unix(100, 0)) + err := scrobbleRepo.RecordScrobble(scrobbleCtx, "3", time.Unix(100, 0)) Expect(err).To(BeNil()) scrobble := host.ScrobbleRef{ID: 5 + int64(i), MediaFileID: "3", SubmissionTime: 100} diff --git a/plugins/host_storage_test.go b/plugins/host_storage_test.go index 9fc58c396..9d8df23a0 100644 --- a/plugins/host_storage_test.go +++ b/plugins/host_storage_test.go @@ -85,7 +85,7 @@ var _ = Describe("Storage Host Function", Ordered, func() { } manager.SetSubsonicRouter(router) - mockPluginRepo := dataStore.Plugin(GinkgoT().Context()).(*tests.MockPluginRepo) + mockPluginRepo := dataStore.Plugin().(*tests.MockPluginRepo) mockPluginRepo.Permitted = true // Setup config diff --git a/plugins/host_subsonicapi.go b/plugins/host_subsonicapi.go index dba58d795..a8ff12140 100644 --- a/plugins/host_subsonicapi.go +++ b/plugins/host_subsonicapi.go @@ -138,7 +138,7 @@ func (s *subsonicAPIServiceImpl) checkPermissions(ctx context.Context, username } // Look up the user by username to get their ID - usr, err := s.ds.User(ctx).FindByUsername(username) + usr, err := s.ds.User().FindByUsername(ctx, username) if err != nil { if errors.Is(err, model.ErrNotFound) { return fmt.Errorf("username %s not found", username) diff --git a/plugins/host_subsonicapi_test.go b/plugins/host_subsonicapi_test.go index 4b941bc43..0d9c75ee7 100644 --- a/plugins/host_subsonicapi_test.go +++ b/plugins/host_subsonicapi_test.go @@ -1,6 +1,7 @@ package plugins import ( + "context" "crypto/sha256" "encoding/hex" "encoding/json" @@ -51,12 +52,12 @@ var _ = Describe("SubsonicAPI Host Function", Ordered, func() { dataStore = &tests.MockDataStore{MockedUser: userRepo} // Add test users - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(GinkgoT().Context(), &model.User{ ID: "user1", UserName: "testuser", IsAdmin: false, }) - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(GinkgoT().Context(), &model.User{ ID: "admin1", UserName: "adminuser", IsAdmin: true, @@ -77,7 +78,7 @@ var _ = Describe("SubsonicAPI Host Function", Ordered, func() { hash := sha256.Sum256(wasmData) hashHex := hex.EncodeToString(hash[:]) - mockPluginRepo := dataStore.Plugin(GinkgoT().Context()).(*tests.MockPluginRepo) + mockPluginRepo := dataStore.Plugin().(*tests.MockPluginRepo) mockPluginRepo.Permitted = true enabledPlugin := model.Plugin{ ID: "test-subsonicapi-plugin", @@ -234,27 +235,29 @@ var _ = Describe("SubsonicAPI Host Function", Ordered, func() { var _ = Describe("SubsonicAPIService", func() { var ( + ctx context.Context router *fakeSubsonicRouter userRepo *tests.MockedUserRepo dataStore *tests.MockDataStore ) BeforeEach(func() { + ctx = GinkgoT().Context() router = &fakeSubsonicRouter{} userRepo = tests.CreateMockUserRepo() dataStore = &tests.MockDataStore{MockedUser: userRepo} - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(ctx, &model.User{ ID: "user1", UserName: "testuser", IsAdmin: false, }) - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(ctx, &model.User{ ID: "admin1", UserName: "adminuser", IsAdmin: true, }) - _ = userRepo.Put(&model.User{ + _ = userRepo.Put(ctx, &model.User{ ID: "user2", UserName: "alloweduser", IsAdmin: false, @@ -267,7 +270,6 @@ var _ = Describe("SubsonicAPIService", func() { // allowedUserIDs contains "user2", but testuser is "user1" service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{"user2"}, false)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/ping?u=testuser") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("not authorized")) @@ -277,7 +279,6 @@ var _ = Describe("SubsonicAPIService", func() { // allowedUserIDs contains "user2" which is "alloweduser" service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{"user2"}, false)) - ctx := GinkgoT().Context() response, err := service.Call(ctx, "/ping?u=alloweduser") Expect(err).ToNot(HaveOccurred()) Expect(response).To(ContainSubstring("ok")) @@ -287,7 +288,6 @@ var _ = Describe("SubsonicAPIService", func() { // allowedUserIDs only contains "user1" (testuser), not "admin1" service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{"user1"}, false)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/ping?u=adminuser") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("not authorized")) @@ -297,7 +297,6 @@ var _ = Describe("SubsonicAPIService", func() { // allowedUserIDs contains "admin1" service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{"admin1"}, false)) - ctx := GinkgoT().Context() response, err := service.Call(ctx, "/ping?u=adminuser") Expect(err).ToNot(HaveOccurred()) Expect(response).To(ContainSubstring("ok")) @@ -308,7 +307,6 @@ var _ = Describe("SubsonicAPIService", func() { It("allows all users regardless of allowed list", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() response, err := service.Call(ctx, "/ping?u=testuser") Expect(err).ToNot(HaveOccurred()) Expect(response).To(ContainSubstring("ok")) @@ -317,7 +315,6 @@ var _ = Describe("SubsonicAPIService", func() { It("allows admin users when allUsers is true", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() response, err := service.Call(ctx, "/ping?u=adminuser") Expect(err).ToNot(HaveOccurred()) Expect(response).To(ContainSubstring("ok")) @@ -328,7 +325,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error when no users are configured", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, false)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/ping?u=testuser") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("no users configured")) @@ -337,7 +333,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error for empty user list", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{}, false)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/ping?u=testuser") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("no users configured")) @@ -349,7 +344,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error for missing username parameter", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/ping") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("missing required parameter")) @@ -358,7 +352,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error for invalid URL", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "://invalid") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("invalid URL")) @@ -367,7 +360,6 @@ var _ = Describe("SubsonicAPIService", func() { It("extracts endpoint from path correctly", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{"user1"}, false)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/rest/ping.view?u=testuser") Expect(err).ToNot(HaveOccurred()) @@ -380,7 +372,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns binary data and content-type", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() contentType, data, err := service.CallRaw(ctx, "/getCoverArt?u=testuser&id=al-1") Expect(err).ToNot(HaveOccurred()) Expect(contentType).To(Equal("image/png")) @@ -390,7 +381,6 @@ var _ = Describe("SubsonicAPIService", func() { It("does not set f=json parameter", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, _, err := service.CallRaw(ctx, "/getCoverArt?u=testuser&id=al-1") Expect(err).ToNot(HaveOccurred()) @@ -402,7 +392,6 @@ var _ = Describe("SubsonicAPIService", func() { It("enforces permission checks", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess([]string{"user2"}, false)) - ctx := GinkgoT().Context() _, _, err := service.CallRaw(ctx, "/getCoverArt?u=testuser&id=al-1") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("not authorized")) @@ -411,7 +400,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error when username is missing", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, _, err := service.CallRaw(ctx, "/getCoverArt") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("missing required parameter")) @@ -420,7 +408,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error when router is nil", func() { service := newSubsonicAPIService("test-plugin", nil, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, _, err := service.CallRaw(ctx, "/getCoverArt?u=testuser") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("router not available")) @@ -429,7 +416,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error for invalid URL", func() { service := newSubsonicAPIService("test-plugin", router, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, _, err := service.CallRaw(ctx, "://invalid") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("invalid URL")) @@ -440,7 +426,6 @@ var _ = Describe("SubsonicAPIService", func() { It("returns error when router is nil", func() { service := newSubsonicAPIService("test-plugin", nil, dataStore, newUserAccess(nil, true)) - ctx := GinkgoT().Context() _, err := service.Call(ctx, "/ping?u=testuser") Expect(err).To(HaveOccurred()) Expect(err.Error()).To(ContainSubstring("router not available")) diff --git a/plugins/host_taskqueue.go b/plugins/host_taskqueue.go index 2f74c0aa4..b15cfaf56 100644 --- a/plugins/host_taskqueue.go +++ b/plugins/host_taskqueue.go @@ -71,7 +71,7 @@ type taskQueueServiceImpl struct { manager *Manager maxConcurrency int32 db *sql.DB - ctx context.Context + ctx context.Context //nolint:containedctx // service lifecycle ctx for the worker goroutines cancel context.CancelFunc wg sync.WaitGroup mu sync.Mutex diff --git a/plugins/host_users.go b/plugins/host_users.go index a56c8f866..28f4dc9b0 100644 --- a/plugins/host_users.go +++ b/plugins/host_users.go @@ -23,7 +23,7 @@ func newUsersService(ds model.DataStore, allowedUsers []string, allUsers bool) h } func (s *usersServiceImpl) GetUsers(ctx context.Context) ([]host.User, error) { - users, err := s.ds.User(ctx).GetAll() + users, err := s.ds.User().GetAll(ctx) if err != nil { return nil, err } diff --git a/plugins/host_users_test.go b/plugins/host_users_test.go index 1721d3ee2..56edd9f67 100644 --- a/plugins/host_users_test.go +++ b/plugins/host_users_test.go @@ -35,21 +35,21 @@ var _ = Describe("UsersService", Ordered, func() { var mockUserRepo *tests.MockedUserRepo BeforeEach(func() { - mockUserRepo = ds.User(ctx).(*tests.MockedUserRepo) + mockUserRepo = ds.User().(*tests.MockedUserRepo) // Add test users - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user1", UserName: "alice", Name: "Alice Admin", IsAdmin: true, }) - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user2", UserName: "bob", Name: "Bob User", IsAdmin: false, }) - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user3", UserName: "charlie", Name: "Charlie User", @@ -144,21 +144,21 @@ var _ = Describe("UsersService", Ordered, func() { var mockUserRepo *tests.MockedUserRepo BeforeEach(func() { - mockUserRepo = ds.User(ctx).(*tests.MockedUserRepo) + mockUserRepo = ds.User().(*tests.MockedUserRepo) // Add test users - alice is admin, bob and charlie are not - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user1", UserName: "alice", Name: "Alice Admin", IsAdmin: true, }) - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user2", UserName: "bob", Name: "Bob User", IsAdmin: false, }) - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user3", UserName: "charlie", Name: "Charlie User", @@ -458,20 +458,20 @@ func setupTestUsersPlugin() (*testUsersSetup, error) { } // createTestUsers creates standard test users in the mock repo -func createTestUsers(mockUserRepo *tests.MockedUserRepo) { - _ = mockUserRepo.Put(&model.User{ +func createTestUsers(ctx context.Context, mockUserRepo *tests.MockedUserRepo) { + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user1", UserName: "alice", Name: "Alice Admin", IsAdmin: true, }) - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user2", UserName: "bob", Name: "Bob User", IsAdmin: false, }) - _ = mockUserRepo.Put(&model.User{ + _ = mockUserRepo.Put(ctx, &model.User{ ID: "user3", UserName: "charlie", Name: "Charlie User", @@ -560,7 +560,7 @@ func setupUsersIntegrationManagerWithEnabled(enabled, allUsers bool, allowedUser }}) mockUserRepo := tests.CreateMockUserRepo() - createTestUsers(mockUserRepo) + createTestUsers(GinkgoT().Context(), mockUserRepo) dataStore := &tests.MockDataStore{ MockedPlugin: mockPluginRepo, diff --git a/plugins/host_websocket.go b/plugins/host_websocket.go index d9d82665c..a58deb129 100644 --- a/plugins/host_websocket.go +++ b/plugins/host_websocket.go @@ -56,7 +56,7 @@ type wsConnection struct { // webSocketServiceImpl implements host.WebSocketService. // It provides plugins with WebSocket communication capabilities. type webSocketServiceImpl struct { - baseCtx context.Context // bounds the read loops, which outlive the Connect() call + baseCtx context.Context //nolint:containedctx // bounds the read loops, which outlive the Connect() call pluginName string manager *Manager requiredHosts []string diff --git a/plugins/manager.go b/plugins/manager.go index bab67e987..90247a75c 100644 --- a/plugins/manager.go +++ b/plugins/manager.go @@ -49,7 +49,7 @@ type PluginMetricsRecorder interface { type Manager struct { mu sync.RWMutex plugins map[string]*plugin - ctx context.Context + ctx context.Context //nolint:containedctx // manager lifecycle ctx, cancelled by Stop cancel context.CancelFunc cache wazero.CompilationCache stopped atomic.Bool // Set to true when Stop() is called @@ -134,7 +134,7 @@ func (m *Manager) Start(ctx context.Context) error { // Clear previous error states so plugins can be retried on restart adminCtx := adminContext(ctx) - if err := m.ds.Plugin(adminCtx).ClearErrors(); err != nil { + if err := m.ds.Plugin().ClearErrors(adminCtx); err != nil { log.Error(ctx, "Error clearing plugin errors", err) } @@ -323,9 +323,9 @@ func (m *Manager) EnablePlugin(ctx context.Context, id string) error { } adminCtx := adminContext(ctx) - repo := m.ds.Plugin(adminCtx) + repo := m.ds.Plugin() - plugin, err := repo.Get(id) + plugin, err := repo.Get(adminCtx, id) if err != nil { return fmt.Errorf("getting plugin from DB: %w", err) } @@ -344,7 +344,7 @@ func (m *Manager) EnablePlugin(ctx context.Context, id string) error { // Store error and return plugin.LastError = err.Error() plugin.UpdatedAt = time.Now() - _ = repo.Put(plugin) + _ = repo.Put(adminCtx, plugin) return fmt.Errorf("loading plugin: %w", err) } @@ -352,7 +352,7 @@ func (m *Manager) EnablePlugin(ctx context.Context, id string) error { plugin.Enabled = true plugin.LastError = "" plugin.UpdatedAt = time.Now() - if err := repo.Put(plugin); err != nil { + if err := repo.Put(adminCtx, plugin); err != nil { // Unload since we couldn't update DB _ = m.unloadPlugin(id) return fmt.Errorf("updating plugin in DB: %w", err) @@ -371,9 +371,9 @@ func (m *Manager) DisablePlugin(ctx context.Context, id string) error { } adminCtx := adminContext(ctx) - repo := m.ds.Plugin(adminCtx) + repo := m.ds.Plugin() - plugin, err := repo.Get(id) + plugin, err := repo.Get(adminCtx, id) if err != nil { return fmt.Errorf("getting plugin from DB: %w", err) } @@ -390,7 +390,7 @@ func (m *Manager) DisablePlugin(ctx context.Context, id string) error { // Update DB plugin.Enabled = false plugin.UpdatedAt = time.Now() - if err := repo.Put(plugin); err != nil { + if err := repo.Put(adminCtx, plugin); err != nil { return fmt.Errorf("updating plugin in DB: %w", err) } @@ -408,9 +408,9 @@ func (m *Manager) ValidatePluginConfig(ctx context.Context, id, configJSON strin } adminCtx := adminContext(ctx) - repo := m.ds.Plugin(adminCtx) + repo := m.ds.Plugin() - plugin, err := repo.Get(id) + plugin, err := repo.Get(adminCtx, id) if err != nil { return fmt.Errorf("getting plugin from DB: %w", err) } @@ -476,9 +476,9 @@ func (m *Manager) updatePluginSettings(ctx context.Context, id string, updateFn } adminCtx := adminContext(ctx) - repo := m.ds.Plugin(adminCtx) + repo := m.ds.Plugin() - plugin, err := repo.Get(id) + plugin, err := repo.Get(adminCtx, id) if err != nil { return fmt.Errorf("getting plugin from DB: %w", err) } @@ -512,7 +512,7 @@ func (m *Manager) updatePluginSettings(ctx context.Context, id string, updateFn log.Debug(ctx, "Plugin was not loaded", "plugin", id) } plugin.Enabled = false - if err := repo.Put(plugin); err != nil { + if err := repo.Put(adminCtx, plugin); err != nil { return fmt.Errorf("updating plugin in DB: %w", err) } log.Info(ctx, "Disabled plugin due to "+disableReason, "plugin", id) @@ -520,7 +520,7 @@ func (m *Manager) updatePluginSettings(ctx context.Context, id string, updateFn return nil } - if err := repo.Put(plugin); err != nil { + if err := repo.Put(adminCtx, plugin); err != nil { return fmt.Errorf("updating plugin in DB: %w", err) } @@ -532,7 +532,7 @@ func (m *Manager) updatePluginSettings(ctx context.Context, id string, updateFn if err := m.loadPluginWithConfig(plugin); err != nil { plugin.LastError = err.Error() plugin.Enabled = false - _ = repo.Put(plugin) + _ = repo.Put(adminCtx, plugin) return fmt.Errorf("reloading plugin: %w", err) } } @@ -586,10 +586,10 @@ func (m *Manager) UnloadDisabledPlugins(ctx context.Context) { } adminCtx := adminContext(ctx) - repo := m.ds.Plugin(adminCtx) + repo := m.ds.Plugin() // Get all disabled plugins from the database - plugins, err := repo.GetAll(model.QueryOptions{ + plugins, err := repo.GetAll(adminCtx, model.QueryOptions{ Filters: squirrel.Eq{"enabled": false}, }) if err != nil { diff --git a/plugins/manager_loader.go b/plugins/manager_loader.go index 024b14b72..cca87e5b0 100644 --- a/plugins/manager_loader.go +++ b/plugins/manager_loader.go @@ -219,9 +219,9 @@ func (m *Manager) loadEnabledPlugins(ctx context.Context) error { } adminCtx := adminContext(ctx) - repo := m.ds.Plugin(adminCtx) + repo := m.ds.Plugin() - plugins, err := repo.GetAll() + plugins, err := repo.GetAll(adminCtx) if err != nil { return fmt.Errorf("reading plugins from DB: %w", err) } @@ -257,7 +257,7 @@ func (m *Manager) loadEnabledPlugins(ctx context.Context) error { plugin.LastError = err.Error() plugin.Enabled = false plugin.UpdatedAt = time.Now() - if putErr := repo.Put(&plugin); putErr != nil { + if putErr := repo.Put(adminCtx, &plugin); putErr != nil { log.Error(ctx, "Failed to update plugin error in DB", "plugin", plugin.ID, putErr) } } @@ -269,7 +269,7 @@ func (m *Manager) loadEnabledPlugins(ctx context.Context) error { if plugin.LastError != "" && m.transient == nil { plugin.LastError = "" plugin.UpdatedAt = time.Now() - if putErr := repo.Put(&plugin); putErr != nil { + if putErr := repo.Put(adminCtx, &plugin); putErr != nil { log.Error(ctx, "Failed to clear plugin error in DB", "plugin", plugin.ID, putErr) } } @@ -347,7 +347,7 @@ func (m *Manager) loadPluginWithConfig(p *model.Plugin) error { if pkg.Manifest.HasLibraryFilesystemPermission() { adminCtx := adminContext(ctx) - libraries, err := m.ds.Library(adminCtx).GetAll() + libraries, err := m.ds.Library().GetAll(adminCtx) if err != nil { return fmt.Errorf("failed to get libraries for filesystem access: %w", err) } diff --git a/plugins/manager_plugin.go b/plugins/manager_plugin.go index 13375a70f..015cd7f4c 100644 --- a/plugins/manager_plugin.go +++ b/plugins/manager_plugin.go @@ -126,7 +126,7 @@ func (a userAccess) resolve(ctx context.Context, ds model.DataStore, username st if !a.allUsers && len(a.userIDSet) == 0 { return nil, fmt.Errorf("plugin is not authorized to scope by user") } - usr, err := ds.User(ctx).FindByUsername(username) + usr, err := ds.User().FindByUsername(ctx, username) if err != nil { if errors.Is(err, model.ErrNotFound) { return nil, fmt.Errorf("user %q not found", username) diff --git a/plugins/manager_readonly_test.go b/plugins/manager_readonly_test.go index 019b14fbf..0c88459c0 100644 --- a/plugins/manager_readonly_test.go +++ b/plugins/manager_readonly_test.go @@ -1,6 +1,7 @@ package plugins import ( + "context" "os" "path/filepath" @@ -14,11 +15,16 @@ import ( var _ = Describe("Manager.LoadPlugins", func() { var ( + ctx context.Context mgr *Manager repo *tests.MockPluginRepo tmpDir string ) + BeforeEach(func() { + ctx = GinkgoT().Context() + }) + // newManager builds a manager over rows the caller can corrupt, with no Subsonic router: a CLI // has none, and Start would log.Fatal on that. newManager := func(rows model.Plugins) *Manager { @@ -54,7 +60,7 @@ var _ = Describe("Manager.LoadPlugins", func() { It("detects capabilities without a Subsonic router configured", func() { mgr = newManager(nil) - Expect(mgr.LoadPlugins(GinkgoT().Context(), []string{"test-metadata-agent", "broken"}, false)).To(Succeed()) + Expect(mgr.LoadPlugins(ctx, []string{"test-metadata-agent", "broken"}, false)).To(Succeed()) Expect(mgr.PluginNames(string(CapabilityMetadataAgent))).To(ContainElement("test-metadata-agent")) }) @@ -70,9 +76,9 @@ var _ = Describe("Manager.LoadPlugins", func() { It("leaves the stored row untouched", func() { mgr = newManager(brokenRows()) - Expect(mgr.LoadPlugins(GinkgoT().Context(), []string{"test-metadata-agent", "broken"}, false)).To(Succeed()) + Expect(mgr.LoadPlugins(ctx, []string{"test-metadata-agent", "broken"}, false)).To(Succeed()) - stored, err := repo.Get("broken") + stored, err := repo.Get(ctx, "broken") Expect(err).ToNot(HaveOccurred()) Expect(stored.Enabled).To(BeTrue(), "inspecting a plugin must never disable it") Expect(stored.LastError).To(BeEmpty()) @@ -83,9 +89,9 @@ var _ = Describe("Manager.LoadPlugins", func() { It("still disables it when not read-only", func() { mgr = newManager(brokenRows()) - Expect(mgr.loadEnabledPlugins(GinkgoT().Context())).To(Succeed()) + Expect(mgr.loadEnabledPlugins(ctx)).To(Succeed()) - stored, err := repo.Get("broken") + stored, err := repo.Get(ctx, "broken") Expect(err).ToNot(HaveOccurred()) Expect(stored.Enabled).To(BeFalse()) Expect(stored.LastError).ToNot(BeEmpty()) @@ -97,7 +103,7 @@ var _ = Describe("Manager.LoadPlugins", func() { It("does not load a plugin that is not in the agent list", func() { mgr = newManager(nil) - Expect(mgr.LoadPlugins(GinkgoT().Context(), []string{"some-other-agent"}, false)).To(Succeed()) + Expect(mgr.LoadPlugins(ctx, []string{"some-other-agent"}, false)).To(Succeed()) Expect(mgr.PluginNames(string(CapabilityMetadataAgent))).To(BeEmpty()) }) @@ -105,7 +111,7 @@ var _ = Describe("Manager.LoadPlugins", func() { It("does nothing when no agents are configured", func() { mgr = newManager(nil) - Expect(mgr.LoadPlugins(GinkgoT().Context(), nil, false)).To(Succeed()) + Expect(mgr.LoadPlugins(ctx, nil, false)).To(Succeed()) Expect(mgr.PluginNames(string(CapabilityMetadataAgent))).To(BeEmpty()) // Not even the wazero cache: with nothing to load there is nothing to compile. @@ -116,7 +122,7 @@ var _ = Describe("Manager.LoadPlugins", func() { mgr = newManager(nil) conf.Server.Plugins.Enabled = false - Expect(mgr.LoadPlugins(GinkgoT().Context(), []string{"test-metadata-agent", "broken"}, false)).To(Succeed()) + Expect(mgr.LoadPlugins(ctx, []string{"test-metadata-agent", "broken"}, false)).To(Succeed()) Expect(mgr.PluginNames(string(CapabilityMetadataAgent))).To(BeEmpty()) }) diff --git a/plugins/manager_sync.go b/plugins/manager_sync.go index 480fa1bb9..33b4f81d4 100644 --- a/plugins/manager_sync.go +++ b/plugins/manager_sync.go @@ -65,7 +65,7 @@ func (m *Manager) addPluginToDB(ctx context.Context, repo model.PluginRepository CreatedAt: now, UpdatedAt: now, } - if err := repo.Put(newPlugin); err != nil { + if err := repo.Put(ctx, newPlugin); err != nil { return fmt.Errorf("adding plugin to DB: %w", err) } log.Info(ctx, "Discovered new plugin", "plugin", name) @@ -88,7 +88,7 @@ func (m *Manager) updatePluginInDB(ctx context.Context, repo model.PluginReposit dbPlugin.Enabled = false dbPlugin.LastError = "" dbPlugin.UpdatedAt = time.Now() - if err := repo.Put(dbPlugin); err != nil { + if err := repo.Put(ctx, dbPlugin); err != nil { return fmt.Errorf("updating plugin in DB: %w", err) } log.Info(ctx, "Plugin file changed", "plugin", dbPlugin.ID, "wasEnabled", wasEnabled) @@ -105,7 +105,7 @@ func (m *Manager) removePluginFromDB(ctx context.Context, repo model.PluginRepos log.Debug(ctx, "Plugin not loaded during removal", "plugin", pluginID, err) } } - if err := repo.Delete(pluginID); err != nil { + if err := repo.Delete(ctx, pluginID); err != nil { return fmt.Errorf("deleting plugin from DB: %w", err) } // Discard any scrobbles still buffered for the removed plugin, so they are @@ -115,7 +115,7 @@ func (m *Manager) removePluginFromDB(ctx context.Context, repo model.PluginRepos // wipe the builtin Last.fm retry queue. if scrobbler.IsBuiltinScrobbler(pluginID) { log.Debug(ctx, "Keeping buffered scrobbles: name is owned by a builtin scrobbler", "plugin", pluginID) - } else if err := m.ds.ScrobbleBuffer(ctx).Discard(pluginID); err != nil { + } else if err := m.ds.ScrobbleBuffer().Discard(ctx, pluginID); err != nil { log.Error(ctx, "Error discarding buffered scrobbles for removed plugin", "plugin", pluginID, err) } log.Info(ctx, "Plugin removed", "plugin", pluginID) @@ -162,8 +162,8 @@ func (m *Manager) syncPlugins(ctx context.Context, folder string) error { log.Debug(ctx, "Plugin sync: scanned folder", "folder", folder, "entriesTotal", len(entries), "pluginsFound", len(filesOnDisk)) // Get all plugins from DB - repo := m.ds.Plugin(adminCtx) - dbPlugins, err := repo.GetAll() + repo := m.ds.Plugin() + dbPlugins, err := repo.GetAll(adminCtx) if err != nil { return fmt.Errorf("reading plugins from DB: %w", err) } @@ -192,7 +192,7 @@ func (m *Manager) syncPlugins(ctx context.Context, folder string) error { if dbPlugin.Path != path { dbPlugin.Path = path dbPlugin.UpdatedAt = now - if err := repo.Put(dbPlugin); err != nil { + if err := repo.Put(adminCtx, dbPlugin); err != nil { log.Error(ctx, "Failed to update plugin path in DB", "plugin", name, err) } } @@ -215,7 +215,7 @@ func (m *Manager) syncPlugins(ctx context.Context, folder string) error { } dbPlugin.Enabled = false } - if putErr := repo.Put(dbPlugin); putErr != nil { + if putErr := repo.Put(adminCtx, dbPlugin); putErr != nil { log.Error(ctx, "Failed to update plugin in DB", "plugin", name, err) } } @@ -225,12 +225,12 @@ func (m *Manager) syncPlugins(ctx context.Context, folder string) error { if !exists { // New plugin - add to DB as disabled - if err := m.addPluginToDB(ctx, repo, name, path, metadata); err != nil { + if err := m.addPluginToDB(adminCtx, repo, name, path, metadata); err != nil { log.Error(ctx, "Failed to add plugin to DB", "plugin", name, err) } } else { // Plugin changed - update DB - if err := m.updatePluginInDB(ctx, repo, dbPlugin, path, metadata); err != nil { + if err := m.updatePluginInDB(adminCtx, repo, dbPlugin, path, metadata); err != nil { log.Error(ctx, "Failed to update plugin in DB", "plugin", name, err) } } @@ -240,7 +240,7 @@ func (m *Manager) syncPlugins(ctx context.Context, folder string) error { // Remove plugins no longer on disk for _, dbPlugin := range pluginsInDB { - if err := m.removePluginFromDB(ctx, repo, dbPlugin); err != nil { + if err := m.removePluginFromDB(adminCtx, repo, dbPlugin); err != nil { log.Error(ctx, "Failed to delete plugin from DB", "plugin", dbPlugin.ID, err) } } diff --git a/plugins/manager_sync_test.go b/plugins/manager_sync_test.go index dd64dcd3f..e190abd34 100644 --- a/plugins/manager_sync_test.go +++ b/plugins/manager_sync_test.go @@ -13,11 +13,13 @@ import ( ) var _ = Describe("syncPlugins", func() { + var ctx context.Context var m *Manager var repo *tests.MockPluginRepo var folder string BeforeEach(func() { + ctx = GinkgoT().Context() folder = GinkgoT().TempDir() repo = tests.CreateMockPluginRepo() repo.SetData(model.Plugins{}) @@ -36,7 +38,7 @@ var _ = Describe("syncPlugins", func() { Expect(m.syncPlugins(context.Background(), folder)).To(Succeed()) - _, err := repo.Get("my-plugin") + _, err := repo.Get(ctx, "my-plugin") Expect(err).ToNot(HaveOccurred()) }) @@ -46,18 +48,23 @@ var _ = Describe("syncPlugins", func() { Expect(m.syncPlugins(context.Background(), folder)).To(Succeed()) - all, err := repo.GetAll() + all, err := repo.GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(all).To(BeEmpty()) }) }) var _ = Describe("removePluginFromDB", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) + It("discards buffered scrobbles for the removed plugin", func() { - ctx := context.Background() buffer := tests.CreateMockedScrobbleBufferRepo() - Expect(buffer.Enqueue("my-plugin", "user1", "track1", time.Now())).To(Succeed()) - Expect(buffer.Enqueue("other-plugin", "user1", "track2", time.Now())).To(Succeed()) + Expect(buffer.Enqueue(ctx, "my-plugin", "user1", "track1", time.Now())).To(Succeed()) + Expect(buffer.Enqueue(ctx, "other-plugin", "user1", "track2", time.Now())).To(Succeed()) repo := tests.CreateMockPluginRepo() plugin := model.Plugin{ID: "my-plugin", Enabled: false} @@ -70,22 +77,21 @@ var _ = Describe("removePluginFromDB", func() { } Expect(m.removePluginFromDB(ctx, repo, &plugin)).To(Succeed()) - _, err := repo.Get("my-plugin") + _, err := repo.Get(ctx, "my-plugin") Expect(err).To(MatchError(model.ErrNotFound)) - remaining, err := buffer.Length() + remaining, err := buffer.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(remaining).To(Equal(int64(1))) - entry, err := buffer.Next("other-plugin", "user1") + entry, err := buffer.Next(ctx, "other-plugin", "user1") Expect(err).ToNot(HaveOccurred()) Expect(entry).ToNot(BeNil(), "entries of other services must be kept") }) It("keeps buffered scrobbles of a builtin scrobbler sharing the removed plugin's name", func() { - ctx := context.Background() scrobbler.Register("builtin-svc", func(model.DataStore) scrobbler.Scrobbler { return nil }) buffer := tests.CreateMockedScrobbleBufferRepo() - Expect(buffer.Enqueue("builtin-svc", "user1", "track1", time.Now())).To(Succeed()) + Expect(buffer.Enqueue(ctx, "builtin-svc", "user1", "track1", time.Now())).To(Succeed()) repo := tests.CreateMockPluginRepo() plugin := model.Plugin{ID: "builtin-svc", Enabled: false} @@ -96,7 +102,7 @@ var _ = Describe("removePluginFromDB", func() { } Expect(m.removePluginFromDB(ctx, repo, &plugin)).To(Succeed()) - remaining, err := buffer.Length() + remaining, err := buffer.Length(ctx) Expect(err).ToNot(HaveOccurred()) Expect(remaining).To(Equal(int64(1)), "builtin scrobbler queue must not be wiped") }) diff --git a/plugins/manager_watcher.go b/plugins/manager_watcher.go index f7f658be9..ce9e25a9d 100644 --- a/plugins/manager_watcher.go +++ b/plugins/manager_watcher.go @@ -157,7 +157,7 @@ func (m *Manager) processPluginEvent(pluginName string) { log.Debug(m.ctx, "Plugin event action", "plugin", pluginName, "action", action, "path", ndpPath) ctx := adminContext(m.ctx) - repo := m.ds.Plugin(ctx) + repo := m.ds.Plugin() switch action { case actionUpdate: @@ -168,7 +168,7 @@ func (m *Manager) processPluginEvent(pluginName string) { return } - dbPlugin, err := repo.Get(pluginName) + dbPlugin, err := repo.Get(ctx, pluginName) if err != nil { // Plugin not in DB yet, need full manifest extraction to add it metadata, extractErr := m.extractManifest(ndpPath) @@ -176,7 +176,7 @@ func (m *Manager) processPluginEvent(pluginName string) { log.Error(m.ctx, "Failed to extract manifest from new plugin", "plugin", pluginName, extractErr) return } - if addErr := m.addPluginToDB(m.ctx, repo, pluginName, ndpPath, metadata); addErr != nil { + if addErr := m.addPluginToDB(ctx, repo, pluginName, ndpPath, metadata); addErr != nil { log.Error(m.ctx, "Failed to add plugin to DB", "plugin", pluginName, addErr) } return @@ -198,23 +198,23 @@ func (m *Manager) processPluginEvent(pluginName string) { _ = m.unloadPlugin(pluginName) dbPlugin.Enabled = false } - _ = repo.Put(dbPlugin) + _ = repo.Put(ctx, dbPlugin) return } - if err := m.updatePluginInDB(m.ctx, repo, dbPlugin, ndpPath, metadata); err != nil { + if err := m.updatePluginInDB(ctx, repo, dbPlugin, ndpPath, metadata); err != nil { log.Error(m.ctx, "Failed to update plugin in DB", "plugin", pluginName, err) } case actionRemove: // File removed - unload if enabled, delete from DB - dbPlugin, err := repo.Get(pluginName) + dbPlugin, err := repo.Get(ctx, pluginName) if err != nil { log.Debug(m.ctx, "Removed plugin not in DB", "plugin", pluginName) return } - if err := m.removePluginFromDB(m.ctx, repo, dbPlugin); err != nil { + if err := m.removePluginFromDB(ctx, repo, dbPlugin); err != nil { log.Error(m.ctx, "Failed to delete plugin from DB", "plugin", pluginName, err) } } diff --git a/plugins/manager_watcher_test.go b/plugins/manager_watcher_test.go index 99326bde1..17d9489dc 100644 --- a/plugins/manager_watcher_test.go +++ b/plugins/manager_watcher_test.go @@ -31,7 +31,7 @@ var _ = Describe("Plugin Watcher", func() { _ = manager.unloadPlugin("test-metadata-agent") _ = os.Remove(filepath.Join(tmpDir, "test-metadata-agent"+PackageExtension)) // Also remove from DB so tests start with a clean slate - _ = manager.ds.Plugin(ctx).Delete("test-metadata-agent") + _ = manager.ds.Plugin().Delete(ctx, "test-metadata-agent") }) // Helper to copy test plugin into the temp folder @@ -51,7 +51,7 @@ var _ = Describe("Plugin Watcher", func() { // Clean up: unload plugin if loaded, remove copied file, delete from DB _ = manager.unloadPlugin("test-metadata-agent") _ = os.Remove(filepath.Join(tmpDir, "test-metadata-agent"+PackageExtension)) - _ = manager.ds.Plugin(ctx).Delete("test-metadata-agent") + _ = manager.ds.Plugin().Delete(ctx, "test-metadata-agent") }) It("adds plugin to DB when file exists", func() { @@ -62,8 +62,8 @@ var _ = Describe("Plugin Watcher", func() { Expect(manager.PluginNames(string(CapabilityMetadataAgent))).ToNot(ContainElement("test-metadata-agent")) // Verify it was added to DB - repo := manager.ds.Plugin(ctx) - plugin, err := repo.Get("test-metadata-agent") + repo := manager.ds.Plugin() + plugin, err := repo.Get(ctx, "test-metadata-agent") Expect(err).ToNot(HaveOccurred()) Expect(plugin.ID).To(Equal("test-metadata-agent")) Expect(plugin.Enabled).To(BeFalse()) @@ -80,11 +80,11 @@ var _ = Describe("Plugin Watcher", func() { // Modify the stored SHA256 in DB to simulate a file change // (In reality, the file would have different content) - repo := manager.ds.Plugin(ctx) - plugin, err := repo.Get("test-metadata-agent") + repo := manager.ds.Plugin() + plugin, err := repo.Get(ctx, "test-metadata-agent") Expect(err).ToNot(HaveOccurred()) plugin.SHA256 = "different-hash-to-simulate-change" - err = repo.Put(plugin) + err = repo.Put(ctx, plugin) Expect(err).ToNot(HaveOccurred()) // Simulate modification - the plugin should be disabled and unloaded @@ -94,7 +94,7 @@ var _ = Describe("Plugin Watcher", func() { Expect(manager.PluginNames(string(CapabilityMetadataAgent))).ToNot(ContainElement("test-metadata-agent")) // But still in DB (just disabled) - plugin, err = repo.Get("test-metadata-agent") + plugin, err = repo.Get(ctx, "test-metadata-agent") Expect(err).ToNot(HaveOccurred()) Expect(plugin.Enabled).To(BeFalse()) }) @@ -115,8 +115,8 @@ var _ = Describe("Plugin Watcher", func() { Expect(manager.PluginNames(string(CapabilityMetadataAgent))).ToNot(ContainElement("test-metadata-agent")) // And removed from DB - repo := manager.ds.Plugin(ctx) - _, err = repo.Get("test-metadata-agent") + repo := manager.ds.Plugin() + _, err = repo.Get(ctx, "test-metadata-agent") Expect(err).To(HaveOccurred()) }) }) diff --git a/scanner/controller.go b/scanner/controller.go index bfb396c6d..5eed6c58d 100644 --- a/scanner/controller.go +++ b/scanner/controller.go @@ -94,7 +94,7 @@ type scanner interface { } type controller struct { - rootCtx context.Context + rootCtx context.Context //nolint:containedctx // scanner lifecycle ctx ds model.DataStore broker events.Broker metrics metrics.Metrics @@ -108,7 +108,7 @@ type controller struct { // getLastScanTime returns the most recent scan time across all libraries func (s *controller) getLastScanTime(ctx context.Context) (time.Time, error) { - libs, err := s.ds.Library(ctx).GetAll(model.QueryOptions{ + libs, err := s.ds.Library().GetAll(ctx, model.QueryOptions{ Sort: "last_scan_at", Order: "desc", Max: 1, @@ -126,9 +126,9 @@ func (s *controller) getLastScanTime(ctx context.Context) (time.Time, error) { // getScanInfo retrieves scan status from the database func (s *controller) getScanInfo(ctx context.Context) (scanType string, elapsed time.Duration, lastErr string) { - lastErr, _ = s.ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "") - scanType, _ = s.ds.Property(ctx).DefaultGet(consts.LastScanTypeKey, "") - startTimeStr, _ := s.ds.Property(ctx).DefaultGet(consts.LastScanStartTimeKey, "") + lastErr, _ = s.ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "") + scanType, _ = s.ds.Property().DefaultGet(ctx, consts.LastScanTypeKey, "") + startTimeStr, _ := s.ds.Property().DefaultGet(ctx, consts.LastScanStartTimeKey, "") if startTimeStr != "" { startTime, err := time.Parse(time.RFC3339, startTimeStr) @@ -185,7 +185,7 @@ func (s *controller) Status(ctx context.Context) (*model.ScannerStatus, error) { } func (s *controller) getCounters(ctx context.Context) (int64, int64, error) { - libs, err := s.ds.Library(ctx).GetAll() + libs, err := s.ds.Library().GetAll(ctx) if err != nil { return 0, 0, fmt.Errorf("library count: %w", err) } @@ -238,7 +238,7 @@ func (s *controller) ScanFolders(requestCtx context.Context, fullScan bool, targ } // Store scan error in database so it can be displayed in the UI if scanError != nil { - _ = s.ds.Property(ctx).Put(consts.LastScanErrorKey, scanError.Error()) + _ = s.ds.Property().Put(ctx, consts.LastScanErrorKey, scanError.Error()) } // Refresh the query-planner statistics after a successful full scan. This must run in the // server process: with the external scanner, an ANALYZE in the subprocess is invisible to the @@ -324,7 +324,7 @@ func (s *controller) includesUnscannedLibrary(ctx context.Context, targets []mod // anyIncludedLibrary reports whether any library included in the scan (all of them when targets is // empty) matches pred. func anyIncludedLibrary(ctx context.Context, ds model.DataStore, targets []model.ScanTarget, pred func(model.Library) bool) bool { - libraries, err := ds.Library(ctx).GetAll() + libraries, err := ds.Library().GetAll(ctx) if err != nil { return false } diff --git a/scanner/controller_test.go b/scanner/controller_test.go index 45d202904..bdcb99eda 100644 --- a/scanner/controller_test.go +++ b/scanner/controller_test.go @@ -35,7 +35,7 @@ var _ = Describe("Controller", func() { }) It("includes last scan error", func() { - Expect(ds.Property(ctx).Put(consts.LastScanErrorKey, "boom")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.LastScanErrorKey, "boom")).To(Succeed()) status, err := ctrl.Status(ctx) Expect(err).ToNot(HaveOccurred()) Expect(status.LastError).To(Equal("boom")) @@ -43,8 +43,8 @@ var _ = Describe("Controller", func() { It("includes scan type and error in status", func() { // Set up test data in property repo - Expect(ds.Property(ctx).Put(consts.LastScanErrorKey, "test error")).To(Succeed()) - Expect(ds.Property(ctx).Put(consts.LastScanTypeKey, "full")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.LastScanErrorKey, "test error")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.LastScanTypeKey, "full")).To(Succeed()) // Get status and verify basic info status, err := ctrl.Status(ctx) diff --git a/scanner/image_changes.go b/scanner/image_changes.go index a9c4de365..3ea7ce57c 100644 --- a/scanner/image_changes.go +++ b/scanner/image_changes.go @@ -54,7 +54,7 @@ func (c *imageChangeCollector) enqueue(ctx context.Context) { if len(items) == 0 { continue } - if err := c.ds.ArtworkQueue(ctx).Enqueue(items...); err != nil { + if err := c.ds.ArtworkQueue().Enqueue(ctx, items...); err != nil { log.Warn(ctx, "Scanner: could not enqueue artwork for image changes", "lib", lib.Name, err) continue } @@ -77,7 +77,7 @@ func (c *imageChangeCollector) queueItems(ctx context.Context, lib model.Library var items []model.ArtworkQueueItem - albumIDs, err := c.ds.MediaFile(ctx).GetAlbumIDsByFolder(lib, folderIDs...) + albumIDs, err := c.ds.MediaFile().GetAlbumIDsByFolder(ctx, lib, folderIDs...) if err != nil { return nil, err } @@ -90,7 +90,7 @@ func (c *imageChangeCollector) queueItems(ctx context.Context, lib model.Library } // The resolver climbs to the library root, so the subtree below the folder is the affected set. // A failure here must not discard the album items already collected. - artistIDs, err := c.ds.Album(ctx).GetSoleAlbumArtistIDsInSubtrees(lib, artistFolderPaths...) + artistIDs, err := c.ds.Album().GetSoleAlbumArtistIDsInSubtrees(ctx, lib, artistFolderPaths...) if err != nil { log.Warn(ctx, "Scanner: could not map image changes to artists", "lib", lib.Name, err) return items, nil diff --git a/scanner/phase_1_folders.go b/scanner/phase_1_folders.go index feefde032..7b6a6b097 100644 --- a/scanner/phase_1_folders.go +++ b/scanner/phase_1_folders.go @@ -61,7 +61,7 @@ type scanJob struct { func newScanJob(ctx context.Context, ds model.DataStore, lib model.Library, fullScan bool, targetFolders []string) (*scanJob, error) { // Get folder updates, optionally filtered to specific target folders - lastUpdates, err := ds.Folder(ctx).GetFolderUpdateInfo(lib, targetFolders...) + lastUpdates, err := ds.Folder().GetFolderUpdateInfo(ctx, lib, targetFolders...) if err != nil { return nil, fmt.Errorf("getting last updates: %w", err) } @@ -124,8 +124,8 @@ func (j *scanJob) createFolderEntry(path string) *folderEntry { type phaseFolders struct { jobs []*scanJob ds model.DataStore - ctx context.Context - walkCtx context.Context // cancelled when a folder fails to persist, so the walk stops early + ctx context.Context //nolint:containedctx // phase runs under a single scan ctx + walkCtx context.Context //nolint:containedctx // cancelled when a folder fails to persist, so the walk stops early stopWalk context.CancelCauseFunc state *scanState prevAlbumPIDConf string @@ -139,7 +139,7 @@ func (p *phaseFolders) description() string { func (p *phaseFolders) producer() ppl.Producer[*folderEntry] { return ppl.NewProducer(func(put func(entry *folderEntry)) error { var err error - p.prevAlbumPIDConf, err = p.ds.Property(p.ctx).DefaultGet(consts.PIDAlbumKey, "") + p.prevAlbumPIDConf, err = p.ds.Property().DefaultGet(p.ctx, consts.PIDAlbumKey, "") if err != nil { return fmt.Errorf("getting album PID conf: %w", err) } @@ -217,7 +217,7 @@ func (p *phaseFolders) processFolder(entry *folderEntry) (*folderEntry, error) { } // Load children mediafiles from DB - cursor, err := p.ds.MediaFile(p.ctx).GetCursor(model.QueryOptions{ + cursor, err := p.ds.MediaFile().GetCursor(p.ctx, model.QueryOptions{ Filters: squirrel.And{squirrel.Eq{"folder_id": entry.id}}, }) if err != nil { @@ -367,34 +367,34 @@ func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, en albumIDMap := maps.Clone(entry.albumIDMap) // Instantiate all repositories just once per folder - folderRepo := tx.Folder(ctx) - tagRepo := tx.Tag(ctx) - artistRepo := tx.Artist(ctx) - libraryRepo := tx.Library(ctx) - albumRepo := tx.Album(ctx) - mfRepo := tx.MediaFile(ctx) + folderRepo := tx.Folder() + tagRepo := tx.Tag() + artistRepo := tx.Artist() + libraryRepo := tx.Library() + albumRepo := tx.Album() + mfRepo := tx.MediaFile() // Save folder to DB folder := entry.toFolder() - err := folderRepo.Put(folder) + err := folderRepo.Put(ctx, folder) if err != nil { return fmt.Errorf("persisting folder: %w", err) } // Save all tags to DB - err = tagRepo.Add(entry.job.lib.ID, entry.tags...) + err = tagRepo.Add(ctx, entry.job.lib.ID, entry.tags...) if err != nil { return fmt.Errorf("persisting tags: %w", err) } // Save all new/modified artists to DB. Their information will be incomplete, but they will be refreshed later for i := range entry.artists { - err = artistRepo.Put(&entry.artists[i], "name", + err = artistRepo.Put(ctx, &entry.artists[i], "name", "mbz_artist_id", "sort_artist_name", "order_artist_name", "full_text", "search_normalized", "updated_at") if err != nil { return fmt.Errorf("persisting artist %q: %w", entry.artists[i].Name, err) } - err = libraryRepo.AddArtist(entry.job.lib.ID, entry.artists[i].ID) + err = libraryRepo.AddArtist(ctx, entry.job.lib.ID, entry.artists[i].ID) if err != nil { return fmt.Errorf("adding artist %q to library: %w", entry.artists[i].Name, err) } @@ -416,7 +416,7 @@ func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, en // Save all tracks to DB for i := range entry.tracks { - err = mfRepo.Put(&entry.tracks[i]) + err = mfRepo.Put(ctx, &entry.tracks[i]) if err != nil { return fmt.Errorf("persisting track %q: %w", entry.tracks[i].Path, err) } @@ -425,14 +425,14 @@ func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, en // A re-imported track returns to unresolved so new embedded art is picked up lazily. if len(entry.tracks) > 0 { trackIDs := slice.Map(entry.tracks, func(t model.MediaFile) string { return t.ID }) - if err := tx.Artwork(ctx).DeleteForItems(model.KindMediaFileArtwork, trackIDs); err != nil { + if err := tx.Artwork().DeleteForItems(ctx, model.KindMediaFileArtwork, trackIDs); err != nil { log.Warn(ctx, "Scanner: could not invalidate media_file artwork", err) } } // Mark all missing tracks as not available if len(entry.missingTracks) > 0 { - err = mfRepo.MarkMissing(true, entry.missingTracks...) + err = mfRepo.MarkMissing(ctx, true, entry.missingTracks...) if err != nil { return fmt.Errorf("marking missing tracks: %w", err) } @@ -442,7 +442,7 @@ func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, en return mf.AlbumID, struct{}{} }) albumsToUpdate := slices.Collect(maps.Keys(groupedMissingTracks)) - err = albumRepo.Touch(albumsToUpdate...) + err = albumRepo.Touch(ctx, albumsToUpdate...) if err != nil { return fmt.Errorf("touching albums %v: %w", albumsToUpdate, err) } @@ -451,12 +451,12 @@ func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, en // Enqueue artwork resolution for changed albums/artists. Never fails the scan. // A full scan re-imports every track, so a re-import is no evidence the art changed. if len(queueItems) > 0 { - queue := tx.ArtworkQueue(ctx) + queue := tx.ArtworkQueue() enqueue := queue.Enqueue if p.state.fullScan { enqueue = queue.EnqueueIfMissing } - if err := enqueue(queueItems...); err != nil { + if err := enqueue(ctx, queueItems...); err != nil { log.Warn(ctx, "Scanner: could not enqueue artwork resolution", err) } } @@ -467,7 +467,7 @@ func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, en func (p *phaseFolders) persistAlbum(repo model.AlbumRepository, a *model.Album, idMap map[string]string) error { prevID := idMap[a.ID] log.Trace(p.ctx, "Persisting album", "album", a.Name, "albumArtist", a.AlbumArtist, "id", a.ID, "prevID", cmp.Or(prevID, "nil")) - if err := repo.Put(a); err != nil { + if err := repo.Put(p.ctx, a); err != nil { return fmt.Errorf("persisting album %s: %w", a.ID, err) } if prevID == "" { @@ -476,13 +476,13 @@ func (p *phaseFolders) persistAlbum(repo model.AlbumRepository, a *model.Album, // Reassign annotation from previous album to new album log.Trace(p.ctx, "Reassigning album annotations", "from", prevID, "to", a.ID, "album", a.Name) - if err := repo.ReassignAnnotation(prevID, a.ID); err != nil { + if err := repo.ReassignAnnotation(p.ctx, prevID, a.ID); err != nil { log.Warn(p.ctx, "Scanner: Could not reassign annotations", "from", prevID, "to", a.ID, "album", a.Name, err) p.state.sendWarning(fmt.Sprintf("Could not reassign annotations from %s to %s ('%s'): %v", prevID, a.ID, a.Name, err)) } // Keep created_at field from previous instance of the album - if err := repo.CopyAttributes(prevID, a.ID, "created_at"); err != nil { + if err := repo.CopyAttributes(p.ctx, prevID, a.ID, "created_at"); err != nil { // Silently ignore when the previous album is not found if !errors.Is(err, model.ErrNotFound) { log.Warn(p.ctx, "Scanner: Could not copy fields", "from", prevID, "to", a.ID, "album", a.Name, err) @@ -520,14 +520,14 @@ func (p *phaseFolders) finalize(err error) error { continue } folderIDs := slices.Collect(maps.Keys(job.lastUpdates)) - if err := tx.Folder(ctx).MarkMissing(true, folderIDs...); err != nil { + if err := tx.Folder().MarkMissing(ctx, true, folderIDs...); err != nil { return fmt.Errorf("marking missing folders in %s: %w", job.lib.Name, err) } - if err := tx.MediaFile(ctx).MarkMissingByFolder(true, folderIDs...); err != nil { + if err := tx.MediaFile().MarkMissingByFolder(ctx, true, folderIDs...); err != nil { return fmt.Errorf("marking tracks in missing folders in %s: %w", job.lib.Name, err) } // Touch all albums that have missing folders, so they get refreshed in later phases - if _, err := tx.Album(ctx).TouchByMissingFolder(); err != nil { + if _, err := tx.Album().TouchByMissingFolder(ctx); err != nil { return fmt.Errorf("touching albums with missing folders in %s: %w", job.lib.Name, err) } } diff --git a/scanner/phase_2_missing_tracks.go b/scanner/phase_2_missing_tracks.go index 6ccc9a46c..945257e47 100644 --- a/scanner/phase_2_missing_tracks.go +++ b/scanner/phase_2_missing_tracks.go @@ -33,7 +33,7 @@ type missingTracks struct { // 4. Updates the database with the new locations of the matched files and removes the old entries. // 5. Logs the results and finalizes the phase by reporting the total number of matched files. type phaseMissingTracks struct { - ctx context.Context + ctx context.Context //nolint:containedctx // phase runs under a single scan ctx ds model.DataStore totalMatched atomic.Uint32 state *scanState @@ -71,7 +71,7 @@ func (p *phaseMissingTracks) produce(put func(tracks *missingTracks)) error { } for _, lib := range p.state.libraries { log.Debug(p.ctx, "Scanner: Checking missing tracks", "libraryId", lib.ID, "libraryName", lib.Name) - cursor, err := p.ds.MediaFile(p.ctx).GetMissingAndMatching(lib.ID) + cursor, err := p.ds.MediaFile().GetMissingAndMatching(p.ctx, lib.ID) if err != nil { return fmt.Errorf("loading missing tracks for library %s: %w", lib.Name, err) } @@ -232,7 +232,7 @@ func (p *phaseMissingTracks) processCrossLibraryMoves(in *missingTracks) (*missi func (p *phaseMissingTracks) findCrossLibraryMatch(missing model.MediaFile) (model.MediaFile, error) { // First tier: Search by MusicBrainz Track ID if available if missing.MbzReleaseTrackID != "" { - matches, err := p.ds.MediaFile(p.ctx).FindRecentFilesByMBZTrackID(missing, missing.CreatedAt) + matches, err := p.ds.MediaFile().FindRecentFilesByMBZTrackID(p.ctx, missing, missing.CreatedAt) if err != nil { log.Error(p.ctx, "Scanner: Error searching for recent files by MBZ Track ID", "mbzTrackID", missing.MbzReleaseTrackID, err) } else { @@ -251,7 +251,7 @@ func (p *phaseMissingTracks) findCrossLibraryMatch(missing model.MediaFile) (mod } // Second tier: Search by intrinsic properties (title, size, suffix, etc.) - matches, err := p.ds.MediaFile(p.ctx).FindRecentFilesByProperties(missing, missing.CreatedAt) + matches, err := p.ds.MediaFile().FindRecentFilesByProperties(p.ctx, missing, missing.CreatedAt) if err != nil { log.Error(p.ctx, "Scanner: Error searching for recent files by properties", "missing", missing.Path, err) return model.MediaFile{}, err @@ -298,25 +298,25 @@ func (p *phaseMissingTracks) moveMatched(target, missing model.MediaFile) error // Update the target media file with the missing file's ID. This effectively "moves" the track // to the new location while keeping its annotations and references intact. moved.ID = missing.ID - if err := tx.MediaFile(ctx).Put(&moved); err != nil { + if err := tx.MediaFile().Put(ctx, &moved); err != nil { return fmt.Errorf("update matched track: %w", err) } // Discard the new mediafile row (the one that was moved to) - if err := tx.MediaFile(ctx).Delete(target.ID); err != nil { + if err := tx.MediaFile().Delete(ctx, target.ID); err != nil { return fmt.Errorf("delete discarded track: %w", err) } if reassignAlbum { // Reassign direct album annotations (starred, rating) log.Debug(ctx, "Scanner: Reassigning album annotations", "from", oldAlbumID, "to", newAlbumID) - if err := tx.Album(ctx).ReassignAnnotation(oldAlbumID, newAlbumID); err != nil { + if err := tx.Album().ReassignAnnotation(ctx, oldAlbumID, newAlbumID); err != nil { log.Warn(ctx, "Scanner: Could not reassign album annotations", "from", oldAlbumID, "to", newAlbumID, err) } // Keep created_at field from previous instance of the album, so moved albums // don't appear in "Recently Added" - if err := tx.Album(ctx).CopyAttributes(oldAlbumID, newAlbumID, "created_at"); err != nil { + if err := tx.Album().CopyAttributes(ctx, oldAlbumID, newAlbumID, "created_at"); err != nil { if !errors.Is(err, model.ErrNotFound) { log.Warn(ctx, "Scanner: Could not copy album created_at", "from", oldAlbumID, "to", newAlbumID, err) } @@ -360,7 +360,7 @@ func (p *phaseMissingTracks) purgeMissing() error { var deletedCount int64 err := p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { var err error - deletedCount, err = tx.MediaFile(ctx).DeleteAllMissing() + deletedCount, err = tx.MediaFile().DeleteAllMissing(ctx) return err }, "scanner: purge missing") if err != nil { diff --git a/scanner/phase_2_missing_tracks_test.go b/scanner/phase_2_missing_tracks_test.go index b7aa52f90..f61aa4244 100644 --- a/scanner/phase_2_missing_tracks_test.go +++ b/scanner/phase_2_missing_tracks_test.go @@ -131,8 +131,8 @@ var _ = Describe("phaseMissingTracks", func() { missingTrack := model.MediaFile{ID: "1", PID: "A", Path: "dir1/path1.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} matchedTrack := model.MediaFile{ID: "2", PID: "A", Path: "dir2/path2.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -144,7 +144,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(phase.totalMatched.Load()).To(Equal(uint32(1))) Expect(state.changesDetected.Load()).To(BeTrue()) - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) }) @@ -156,8 +156,8 @@ var _ = Describe("phaseMissingTracks", func() { probe = &probeTxDS{MockDataStore: ds.(*tests.MockDataStore)} probe.MockedAlbum = tests.CreateMockAlbumRepo() phase = createPhaseMissingTracks(ctx, state, probe) - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) }) It("claims the target album before the transaction, so a concurrent move skips it", func() { @@ -190,8 +190,8 @@ var _ = Describe("phaseMissingTracks", func() { It("keeps the moved track", func() { missingTrack := model.MediaFile{ID: "1", PID: "A", Path: "dir1/path1.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} matchedTrack := model.MediaFile{ID: "2", PID: "A", Path: "dir2/path2.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) _, err := phase.processMissingTracks(&missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -199,7 +199,7 @@ var _ = Describe("phaseMissingTracks", func() { }) Expect(err).ToNot(HaveOccurred()) - movedTrack, err := ds.MediaFile(ctx).Get("1") + movedTrack, err := ds.MediaFile().Get(ctx, "1") Expect(err).ToNot(HaveOccurred()) Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) }) @@ -217,8 +217,8 @@ var _ = Describe("phaseMissingTracks", func() { } missingTrack := model.MediaFile{ID: "1", PID: "A", AlbumID: "old-album", Path: "dir1/path1.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} matchedTrack := model.MediaFile{ID: "2", PID: "A", AlbumID: "new-album", Path: "dir2/path2.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) _, err := phase.processMissingTracks(&missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -233,8 +233,8 @@ var _ = Describe("phaseMissingTracks", func() { missingTrack := model.MediaFile{ID: "1", PID: "A", Path: "path1.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} matchedTrack := model.MediaFile{ID: "2", PID: "A", Path: "path1.flac", Tags: model.Tags{"title": []string{"title1"}}, Size: 200} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -246,7 +246,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(phase.totalMatched.Load()).To(Equal(uint32(1))) Expect(state.changesDetected.Load()).To(BeTrue()) - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) Expect(movedTrack.Size).To(Equal(matchedTrack.Size)) }) @@ -255,8 +255,8 @@ var _ = Describe("phaseMissingTracks", func() { missingTrack := model.MediaFile{ID: "1", PID: "A", Path: "dir1/path1.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} matchedTrack := model.MediaFile{ID: "2", PID: "A", Path: "dir2/path2.flac", Tags: model.Tags{"title": []string{"different title"}}, Size: 200} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -268,7 +268,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(phase.totalMatched.Load()).To(Equal(uint32(1))) Expect(state.changesDetected.Load()).To(BeTrue()) - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) Expect(movedTrack.Size).To(Equal(matchedTrack.Size)) }) @@ -278,9 +278,9 @@ var _ = Describe("phaseMissingTracks", func() { matchedEquivalent := model.MediaFile{ID: "2", PID: "A", Path: "dir1/file1.flac", Tags: model.Tags{"title": []string{"title1"}}, Size: 200} matchedExact := model.MediaFile{ID: "3", PID: "A", Path: "dir2/file2.mp3", Tags: model.Tags{"title": []string{"title1"}}, Size: 100} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedEquivalent) - _ = ds.MediaFile(ctx).Put(&matchedExact) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedEquivalent) + _ = ds.MediaFile().Put(ctx, &matchedExact) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -293,7 +293,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(phase.totalMatched.Load()).To(Equal(uint32(1))) Expect(state.changesDetected.Load()).To(BeTrue()) - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal(matchedExact.Path)) Expect(movedTrack.Size).To(Equal(matchedExact.Size)) }) @@ -303,9 +303,9 @@ var _ = Describe("phaseMissingTracks", func() { matched1 := model.MediaFile{ID: "2", PID: "A", Path: "dir1/file2.flac", Title: "another title", Size: 200} matched2 := model.MediaFile{ID: "3", PID: "A", Path: "dir2/file3.mp3", Title: "different title", Size: 100} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matched1) - _ = ds.MediaFile(ctx).Put(&matched2) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matched1) + _ = ds.MediaFile().Put(ctx, &matched2) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -318,7 +318,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeFalse()) // The missing track should still be the same - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal(missingTrack.Path)) Expect(movedTrack.Title).To(Equal(missingTrack.Title)) Expect(movedTrack.Size).To(Equal(missingTrack.Size)) @@ -333,9 +333,9 @@ var _ = Describe("phaseMissingTracks", func() { missingTrack2 := model.MediaFile{ID: "2", PID: "A", Path: "old_dir2/song.mp3", Title: "title1", Size: 100} matchedTrack := model.MediaFile{ID: "3", PID: "A", Path: "new_dir/song.mp3", Title: "title1", Size: 200} - _ = ds.MediaFile(ctx).Put(&missingTrack1) - _ = ds.MediaFile(ctx).Put(&missingTrack2) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack1) + _ = ds.MediaFile().Put(ctx, &missingTrack2) + _ = ds.MediaFile().Put(ctx, &matchedTrack) in := &missingTracks{ missing: []model.MediaFile{missingTrack1, missingTrack2}, @@ -349,11 +349,11 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeTrue()) // The matched track should have been consumed by the first missing track - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) // The second missing track should remain unchanged - unmatchedTrack, _ := ds.MediaFile(ctx).Get("2") + unmatchedTrack, _ := ds.MediaFile().Get(ctx, "2") Expect(unmatchedTrack.Path).To(Equal(missingTrack2.Path)) }) @@ -361,8 +361,8 @@ var _ = Describe("phaseMissingTracks", func() { missingTrack := model.MediaFile{ID: "1", PID: "A", Path: "path1.mp3", Tags: model.Tags{"title": []string{"title1"}}} matchedTrack := model.MediaFile{ID: "2", PID: "A", Path: "path1.mp3", Tags: model.Tags{"title": []string{"title1"}}} - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -370,7 +370,7 @@ var _ = Describe("phaseMissingTracks", func() { } // Simulate an error when moving the matched track by deleting the track from the DB - _ = ds.MediaFile(ctx).Delete("2") + _ = ds.MediaFile().Delete(ctx, "2") _, err := phase.processMissingTracks(in) Expect(err).To(HaveOccurred()) @@ -514,8 +514,8 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: scanStartTime.Add(-10 * time.Minute), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&movedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &movedTrack) in := &missingTracks{ lib: model.Library{ID: 1, Name: "Library 1"}, @@ -529,7 +529,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeTrue()) // Verify the move was performed - updatedTrack, _ := ds.MediaFile(ctx).Get("missing1") + updatedTrack, _ := ds.MediaFile().Get(ctx, "missing1") Expect(updatedTrack.Path).To(Equal("/lib2/track.mp3")) Expect(updatedTrack.LibraryID).To(Equal(2)) }) @@ -566,8 +566,8 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: scanStartTime.Add(-10 * time.Minute), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&movedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &movedTrack) in := &missingTracks{ lib: model.Library{ID: 1, Name: "Library 1"}, @@ -581,7 +581,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeTrue()) // Verify the move was performed - updatedTrack, _ := ds.MediaFile(ctx).Get("missing2") + updatedTrack, _ := ds.MediaFile().Get(ctx, "missing2") Expect(updatedTrack.Path).To(Equal("/lib2/track2.flac")) Expect(updatedTrack.LibraryID).To(Equal(2)) }) @@ -612,8 +612,8 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: scanStartTime.Add(-10 * time.Minute), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&sameLibTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &sameLibTrack) in := &missingTracks{ lib: model.Library{ID: 1, Name: "Library 1"}, @@ -670,9 +670,9 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: scanStartTime.Add(-5 * time.Minute), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&mbzTrack) - _ = ds.MediaFile(ctx).Put(&intrinsicTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &mbzTrack) + _ = ds.MediaFile().Put(ctx, &intrinsicTrack) in := &missingTracks{ lib: model.Library{ID: 1, Name: "Library 1"}, @@ -686,7 +686,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeTrue()) // Verify the MBZ track was chosen (not the intrinsic one) - updatedTrack, _ := ds.MediaFile(ctx).Get("missing4") + updatedTrack, _ := ds.MediaFile().Get(ctx, "missing4") Expect(updatedTrack.Path).To(Equal("/lib2/track4.mp3")) Expect(updatedTrack.LibraryID).To(Equal(2)) }) @@ -718,8 +718,8 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: scanStartTime.Add(-10 * time.Minute), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&equivalentTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &equivalentTrack) in := &missingTracks{ lib: model.Library{ID: 1, Name: "Library 1"}, @@ -733,7 +733,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeTrue()) // Verify the equivalent match was accepted - updatedTrack, _ := ds.MediaFile(ctx).Get("missing5") + updatedTrack, _ := ds.MediaFile().Get(ctx, "missing5") Expect(updatedTrack.Path).To(Equal("/lib2/different/track5.mp3")) Expect(updatedTrack.LibraryID).To(Equal(2)) }) @@ -788,9 +788,9 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: scanStartTime.Add(-5 * time.Minute), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&match1) - _ = ds.MediaFile(ctx).Put(&match2) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &match1) + _ = ds.MediaFile().Put(ctx, &match2) in := &missingTracks{ lib: model.Library{ID: 1, Name: "Library 1"}, @@ -804,7 +804,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(state.changesDetected.Load()).To(BeFalse()) // Verify no move was performed - unchangedTrack, _ := ds.MediaFile(ctx).Get("missing6") + unchangedTrack, _ := ds.MediaFile().Get(ctx, "missing6") Expect(unchangedTrack.Path).To(Equal("/lib1/track6.mp3")) Expect(unchangedTrack.LibraryID).To(Equal(1)) }) @@ -844,7 +844,7 @@ var _ = Describe("phaseMissingTracks", func() { var albumRepo *tests.MockAlbumRepo BeforeEach(func() { - albumRepo = ds.Album(ctx).(*tests.MockAlbumRepo) + albumRepo = ds.Album().(*tests.MockAlbumRepo) albumRepo.ReassignAnnotationCalls = make(map[string]string) albumRepo.CopyAttributesCalls = make(map[string]string) }) @@ -868,8 +868,8 @@ var _ = Describe("phaseMissingTracks", func() { Size: 100, } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) in := &missingTracks{ missing: []model.MediaFile{missingTrack}, @@ -879,7 +879,7 @@ var _ = Describe("phaseMissingTracks", func() { _, err := phase.processMissingTracks(in) Expect(err).ToNot(HaveOccurred()) - movedTrack, _ := ds.MediaFile(ctx).Get("1") + movedTrack, _ := ds.MediaFile().Get(ctx, "1") Expect(movedTrack.Path).To(Equal("new/song.mp3")) Expect(movedTrack.CreatedAt).To(Equal(originalTime)) }) @@ -905,21 +905,21 @@ var _ = Describe("phaseMissingTracks", func() { {ID: "new-album", LibraryID: 2, CreatedAt: time.Now()}, }) - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) err := phase.moveMatched(matchedTrack, missingTrack) Expect(err).ToNot(HaveOccurred()) // Track's created_at should be preserved from the missing file - movedTrack, _ := ds.MediaFile(ctx).Get("missing-ca") + movedTrack, _ := ds.MediaFile().Get(ctx, "missing-ca") Expect(movedTrack.CreatedAt).To(Equal(originalTime)) // Album's created_at should be copied from old to new Expect(albumRepo.CopyAttributesCalls).To(HaveKeyWithValue("old-album", "new-album")) // Verify the new album's CreatedAt was actually updated - newAlbum, err := albumRepo.Get("new-album") + newAlbum, err := albumRepo.Get(ctx, "new-album") Expect(err).ToNot(HaveOccurred()) Expect(newAlbum.CreatedAt).To(Equal(originalTime)) }) @@ -939,14 +939,14 @@ var _ = Describe("phaseMissingTracks", func() { CreatedAt: time.Now(), } - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) err := phase.moveMatched(matchedTrack, missingTrack) Expect(err).ToNot(HaveOccurred()) // Track's created_at should still be preserved - movedTrack, _ := ds.MediaFile(ctx).Get("missing-same") + movedTrack, _ := ds.MediaFile().Get(ctx, "missing-same") Expect(movedTrack.CreatedAt).To(Equal(originalTime)) // CopyAttributes should NOT have been called (same album) @@ -964,7 +964,7 @@ var _ = Describe("phaseMissingTracks", func() { ) BeforeEach(func() { - albumRepo = ds.Album(ctx).(*tests.MockAlbumRepo) + albumRepo = ds.Album().(*tests.MockAlbumRepo) albumRepo.ReassignAnnotationCalls = make(map[string]string) oldAlbumID = "old-album-id" @@ -999,8 +999,8 @@ var _ = Describe("phaseMissingTracks", func() { } // Store both tracks in the database - _ = ds.MediaFile(ctx).Put(&missingTrack) - _ = ds.MediaFile(ctx).Put(&matchedTrack) + _ = ds.MediaFile().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) }) When("album ID changes during cross-library move", func() { @@ -1033,7 +1033,7 @@ var _ = Describe("phaseMissingTracks", func() { Expect(err).ToNot(HaveOccurred()) // Verify that the track was still moved (ID should be updated) - movedTrack, err := ds.MediaFile(ctx).Get(missingTrack.ID) + movedTrack, err := ds.MediaFile().Get(ctx, missingTrack.ID) Expect(err).ToNot(HaveOccurred()) Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) }) diff --git a/scanner/phase_3_refresh_albums.go b/scanner/phase_3_refresh_albums.go index 964ab7408..58ad9cb80 100644 --- a/scanner/phase_3_refresh_albums.go +++ b/scanner/phase_3_refresh_albums.go @@ -26,7 +26,7 @@ import ( // 5. As a last step, it refreshes the artist statistics to reflect the changes type phaseRefreshAlbums struct { ds model.DataStore - ctx context.Context + ctx context.Context //nolint:containedctx // phase runs under a single scan ctx refreshed atomic.Uint32 skipped atomic.Uint32 state *scanState @@ -47,7 +47,7 @@ func (p *phaseRefreshAlbums) producer() ppl.Producer[*model.Album] { func (p *phaseRefreshAlbums) produce(put func(album *model.Album)) error { count := 0 for _, lib := range p.state.libraries { - cursor, err := p.ds.Album(p.ctx).GetTouchedAlbums(lib.ID) + cursor, err := p.ds.Album().GetTouchedAlbums(p.ctx, lib.ID) if err != nil { return fmt.Errorf("loading touched albums: %w", err) } @@ -76,7 +76,7 @@ func (p *phaseRefreshAlbums) stages() []ppl.Stage[*model.Album] { } func (p *phaseRefreshAlbums) filterUnmodified(album *model.Album) (*model.Album, error) { - mfs, err := p.ds.MediaFile(p.ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album_id": album.ID}}) + mfs, err := p.ds.MediaFile().GetAll(p.ctx, model.QueryOptions{Filters: squirrel.Eq{"album_id": album.ID}}) if err != nil { log.Error(p.ctx, "Error loading media files for album", "album_id", album.ID, err) return nil, err @@ -104,7 +104,7 @@ func (p *phaseRefreshAlbums) refreshAlbum(album *model.Album) (*model.Album, err } start := time.Now() err := p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { - return tx.Album(ctx).Put(album) + return tx.Album().Put(ctx, album) }, "scanner: refresh album") log.Debug(p.ctx, "Scanner: refreshing album", "album_id", album.ID, "name", album.Name, "songCount", album.SongCount, "elapsed", time.Since(start), err) if err != nil { @@ -135,7 +135,7 @@ func (p *phaseRefreshAlbums) finalize(err error) error { var cnt int64 err = p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { var txErr error - cnt, txErr = tx.Album(ctx).RefreshPlayCounts() + cnt, txErr = tx.Album().RefreshPlayCounts(ctx) return txErr }, "scanner: refresh album play counts") if err != nil { @@ -147,7 +147,7 @@ func (p *phaseRefreshAlbums) finalize(err error) error { start = time.Now() err = p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { var txErr error - cnt, txErr = tx.Artist(ctx).RefreshPlayCounts() + cnt, txErr = tx.Artist().RefreshPlayCounts(ctx) return txErr }, "scanner: refresh artist play counts") if err != nil { diff --git a/scanner/phase_3_refresh_albums_test.go b/scanner/phase_3_refresh_albums_test.go index 1f0baf428..2743da070 100644 --- a/scanner/phase_3_refresh_albums_test.go +++ b/scanner/phase_3_refresh_albums_test.go @@ -104,7 +104,7 @@ var _ = Describe("phaseRefreshAlbums", func() { Describe("refreshAlbum", func() { It("refreshes the album in the database", func() { - Expect(albumRepo.CountAll()).To(Equal(int64(0))) + Expect(albumRepo.CountAll(ctx)).To(Equal(int64(0))) album := &model.Album{ID: "album1", Name: "Album 1"} result, err := phase.refreshAlbum(album) @@ -112,7 +112,7 @@ var _ = Describe("phaseRefreshAlbums", func() { Expect(result).ToNot(BeNil()) Expect(result.ID).To(Equal("album1")) - savedAlbum, err := albumRepo.Get("album1") + savedAlbum, err := albumRepo.Get(ctx, "album1") Expect(err).ToNot(HaveOccurred()) Expect(savedAlbum).ToNot(BeNil()) diff --git a/scanner/phase_4_playlists.go b/scanner/phase_4_playlists.go index 4e11fa81d..bb67c1ba3 100644 --- a/scanner/phase_4_playlists.go +++ b/scanner/phase_4_playlists.go @@ -19,7 +19,7 @@ import ( ) type phasePlaylists struct { - ctx context.Context + ctx context.Context //nolint:containedctx // phase runs under a single scan ctx scanState *scanState ds model.DataStore pls playlists.Playlists @@ -53,7 +53,7 @@ func (p *phasePlaylists) produce(put func(entry *model.Folder)) error { // Resolve the admin at phase time (the producer runs late in the scan), so an // admin created while the scan was in progress is picked up. Assigned once, // before any put() below, so the channel send synchronizes it with the stages. - admin, err := p.ds.User(p.ctx).FindFirstAdmin() + admin, err := p.ds.User().FindFirstAdmin(p.ctx) if err != nil && !errors.Is(err, model.ErrNotFound) { return fmt.Errorf("finding admin user: %w", err) } @@ -71,9 +71,9 @@ func (p *phasePlaylists) produce(put func(entry *model.Folder)) error { p.pendingImport = pending var cursor model.FolderCursor if p.pendingImport { - cursor, err = p.ds.Folder(p.ctx).GetAllWithPlaylists() + cursor, err = p.ds.Folder().GetAllWithPlaylists(p.ctx) } else { - cursor, err = p.ds.Folder(p.ctx).GetTouchedWithPlaylists() + cursor, err = p.ds.Folder().GetTouchedWithPlaylists(p.ctx) } if err != nil { return fmt.Errorf("loading folders with playlists: %w", err) @@ -102,7 +102,7 @@ func (p *phasePlaylists) produce(put func(entry *model.Folder)) error { // the scan does not complete as successful without recording the recovery). func (p *phasePlaylists) deferImport() error { err := p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { - return tx.Property(ctx).Put(consts.PlaylistsImportPendingFlagKey, "1") + return tx.Property().Put(ctx, consts.PlaylistsImportPendingFlagKey, "1") }, "scanner: defer playlist import") if err != nil { return fmt.Errorf("recording pending playlist import: %w", err) @@ -113,7 +113,7 @@ func (p *phasePlaylists) deferImport() error { } func (p *phasePlaylists) importPending() (bool, error) { - v, err := p.ds.Property(p.ctx).DefaultGet(consts.PlaylistsImportPendingFlagKey, "0") + v, err := p.ds.Property().DefaultGet(p.ctx, consts.PlaylistsImportPendingFlagKey, "0") return v == "1", err } @@ -150,7 +150,7 @@ func (p *phasePlaylists) processPlaylistsInFolder(folder *model.Folder) (*model. } item := model.ArtworkQueueItem{ItemKind: model.KindPlaylistArtwork.Prefix(), ItemID: pls.ID, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityScan} - if err := p.ds.ArtworkQueue(p.ctx).Enqueue(item); err != nil { + if err := p.ds.ArtworkQueue().Enqueue(p.ctx, item); err != nil { log.Warn(p.ctx, "Scanner: could not enqueue playlist artwork", "id", pls.ID, err) } p.refreshed.Add(1) @@ -167,7 +167,7 @@ func (p *phasePlaylists) finalize(err error) error { p.scanState.changesDetected.Store(true) } if p.pendingImport && err == nil { - if derr := p.ds.Property(p.ctx).Delete(consts.PlaylistsImportPendingFlagKey); derr != nil { + if derr := p.ds.Property().Delete(p.ctx, consts.PlaylistsImportPendingFlagKey); derr != nil { log.Warn(p.ctx, "Scanner: Could not clear pending playlist-import flag", derr) } } diff --git a/scanner/phase_4_playlists_test.go b/scanner/phase_4_playlists_test.go index 93ec1a36d..303af338f 100644 --- a/scanner/phase_4_playlists_test.go +++ b/scanner/phase_4_playlists_test.go @@ -38,7 +38,7 @@ var _ = Describe("phasePlaylists", func() { folderRepo = &mockFolderRepository{} userRepo = tests.CreateMockUserRepo() // An admin user exists by default, so playlist import proceeds. - Expect(userRepo.Put(&model.User{ID: "123", UserName: "admin", IsAdmin: true})).To(Succeed()) + Expect(userRepo.Put(ctx, &model.User{ID: "123", UserName: "admin", IsAdmin: true})).To(Succeed()) propRepo = &tests.MockedPropertyRepo{} ds = &tests.MockDataStore{ MockedFolder: folderRepo, @@ -102,7 +102,7 @@ var _ = Describe("phasePlaylists", func() { Expect(err).ToNot(HaveOccurred()) Expect(called).To(BeFalse()) - v, _ := propRepo.Get(consts.PlaylistsImportPendingFlagKey) + v, _ := propRepo.Get(ctx, consts.PlaylistsImportPendingFlagKey) Expect(v).To(Equal("1")) }) @@ -113,7 +113,7 @@ var _ = Describe("phasePlaylists", func() { Expect(err).To(MatchError(ContainSubstring("finding admin user"))) // Must NOT have set the pending flag on a real error. - _, getErr := propRepo.Get(consts.PlaylistsImportPendingFlagKey) + _, getErr := propRepo.Get(ctx, consts.PlaylistsImportPendingFlagKey) Expect(getErr).To(HaveOccurred()) }) @@ -127,7 +127,7 @@ var _ = Describe("phasePlaylists", func() { }) It("imports all playlist folders when the pending flag is set", func() { - Expect(propRepo.Put(consts.PlaylistsImportPendingFlagKey, "1")).To(Succeed()) + Expect(propRepo.Put(ctx, consts.PlaylistsImportPendingFlagKey, "1")).To(Succeed()) folderRepo.SetAllData(map[*model.Folder]error{ {Path: "/path/to/folder1"}: nil, {Path: "/path/to/folder2"}: nil, @@ -146,22 +146,22 @@ var _ = Describe("phasePlaylists", func() { Describe("finalize", func() { It("clears the pending flag after a successful pending import", func() { - Expect(propRepo.Put(consts.PlaylistsImportPendingFlagKey, "1")).To(Succeed()) + Expect(propRepo.Put(ctx, consts.PlaylistsImportPendingFlagKey, "1")).To(Succeed()) phase.pendingImport = true Expect(phase.finalize(nil)).To(Succeed()) - _, err := propRepo.Get(consts.PlaylistsImportPendingFlagKey) + _, err := propRepo.Get(ctx, consts.PlaylistsImportPendingFlagKey) Expect(err).To(HaveOccurred()) // deleted }) It("keeps the pending flag when the import failed", func() { - Expect(propRepo.Put(consts.PlaylistsImportPendingFlagKey, "1")).To(Succeed()) + Expect(propRepo.Put(ctx, consts.PlaylistsImportPendingFlagKey, "1")).To(Succeed()) phase.pendingImport = true Expect(phase.finalize(errors.New("boom"))).To(HaveOccurred()) - v, _ := propRepo.Get(consts.PlaylistsImportPendingFlagKey) + v, _ := propRepo.Get(ctx, consts.PlaylistsImportPendingFlagKey) Expect(v).To(Equal("1")) }) }) @@ -204,7 +204,7 @@ var _ = Describe("phasePlaylists", func() { _, err := phase.processPlaylistsInFolder(folder) Expect(err).ToNot(HaveOccurred()) - queued, err := ds.ArtworkQueue(ctx).DequeueBatch(10) + queued, err := ds.ArtworkQueue().DequeueBatch(ctx, 10) Expect(err).ToNot(HaveOccurred()) Expect(queued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "pl"), @@ -264,11 +264,11 @@ func cursorFromData(data map[*model.Folder]error) model.FolderCursor { } } -func (f *mockFolderRepository) GetTouchedWithPlaylists() (model.FolderCursor, error) { +func (f *mockFolderRepository) GetTouchedWithPlaylists(context.Context) (model.FolderCursor, error) { return cursorFromData(f.data), nil } -func (f *mockFolderRepository) GetAllWithPlaylists() (model.FolderCursor, error) { +func (f *mockFolderRepository) GetAllWithPlaylists(context.Context) (model.FolderCursor, error) { return cursorFromData(f.allData), nil } diff --git a/scanner/scanner.go b/scanner/scanner.go index a8a192771..cd2fe3c8d 100644 --- a/scanner/scanner.go +++ b/scanner/scanner.go @@ -88,7 +88,7 @@ func (s *scannerImpl) scanFolders(ctx context.Context, fullScan bool, targets [] } // Get libraries and optionally filter by targets - allLibs, err := s.ds.Library(ctx).GetAll() + allLibs, err := s.ds.Library().GetAll(ctx) if err != nil { state.sendWarning(fmt.Sprintf("getting libraries: %s", err)) return @@ -131,8 +131,8 @@ func (s *scannerImpl) scanFolders(ctx context.Context, fullScan bool, targets [] if state.isSelectiveScan() { scanType += "-selective" } - _ = s.ds.Property(ctx).Put(consts.LastScanTypeKey, scanType) - _ = s.ds.Property(ctx).Put(consts.LastScanStartTimeKey, startTime.Format(time.RFC3339)) + _ = s.ds.Property().Put(ctx, consts.LastScanTypeKey, scanType) + _ = s.ds.Property().Put(ctx, consts.LastScanStartTimeKey, startTime.Format(time.RFC3339)) // if there was a full scan in progress, force a full scan if !state.fullScan { @@ -141,9 +141,9 @@ func (s *scannerImpl) scanFolders(ctx context.Context, fullScan bool, targets [] log.Info(ctx, "Scanner: Interrupted full scan detected", "lib", lib.Name) state.fullScan = true if state.isSelectiveScan() { - _ = s.ds.Property(ctx).Put(consts.LastScanTypeKey, "full-selective") + _ = s.ds.Property().Put(ctx, consts.LastScanTypeKey, "full-selective") } else { - _ = s.ds.Property(ctx).Put(consts.LastScanTypeKey, "full") + _ = s.ds.Property().Put(ctx, consts.LastScanTypeKey, "full") } break } @@ -190,12 +190,12 @@ func (s *scannerImpl) scanFolders(ctx context.Context, fullScan bool, targets [] ) if err != nil { log.Error(ctx, "Scanner: Finished with error", "duration", time.Since(startTime), err) - _ = s.ds.Property(ctx).Put(consts.LastScanErrorKey, err.Error()) + _ = s.ds.Property().Put(ctx, consts.LastScanErrorKey, err.Error()) state.sendError(err) return } - _ = s.ds.Property(ctx).Put(consts.LastScanErrorKey, "") + _ = s.ds.Property().Put(ctx, consts.LastScanErrorKey, "") if state.changesDetected.Load() { state.sendProgress(&ProgressInfo{ChangesDetected: true}) @@ -218,7 +218,7 @@ func (s *scannerImpl) prepareLibrariesForScan(ctx context.Context, state *scanSt if lib.LastScanStartedAt.IsZero() { // This is a new scan - mark it as started err := s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { - return tx.Library(ctx).ScanBegin(lib.ID, state.fullScan) + return tx.Library().ScanBegin(ctx, lib.ID, state.fullScan) }, "scanner: begin library scan") if err != nil { log.Error(ctx, "Scanner: Error marking scan start", "lib", lib.Name, err) @@ -227,7 +227,7 @@ func (s *scannerImpl) prepareLibrariesForScan(ctx context.Context, state *scanSt } // Reload library to get updated state (timestamps, etc.) - reloadedLib, err := s.ds.Library(ctx).Get(lib.ID) + reloadedLib, err := s.ds.Library().Get(ctx, lib.ID) if err != nil { log.Error(ctx, "Scanner: Error reloading library", "lib", lib.Name, err) state.sendWarning(err.Error()) @@ -291,7 +291,7 @@ func (s *scannerImpl) runEnqueueMissingArtwork(ctx context.Context, state *scanS var n int64 err := s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { var err error - n, err = tx.ArtworkQueue(ctx).EnqueueAllMissing(kind, model.ArtworkPriorityScan) + n, err = tx.ArtworkQueue().EnqueueAllMissing(ctx, kind, model.ArtworkPriorityScan) return err }, "scanner: enqueue missing artwork") if err != nil { @@ -312,7 +312,7 @@ func (s *scannerImpl) runRefreshStats(ctx context.Context, state *scanState) fun return nil } start := time.Now() - stats, err := s.ds.Artist(ctx).RefreshStats(state.fullScan) + stats, err := s.ds.Artist().RefreshStats(ctx, state.fullScan) if err != nil { log.Error(ctx, "Scanner: Error refreshing artists stats", err) return fmt.Errorf("refreshing artists stats: %w", err) @@ -321,7 +321,7 @@ func (s *scannerImpl) runRefreshStats(ctx context.Context, state *scanState) fun start = time.Now() err = s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { - return tx.Tag(ctx).UpdateCounts() + return tx.Tag().UpdateCounts(ctx) }, "scanner: update tag counts") if err != nil { log.Error(ctx, "Scanner: Error updating tag counts", err) @@ -337,18 +337,18 @@ func (s *scannerImpl) runUpdateLibraries(ctx context.Context, state *scanState) start := time.Now() return s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { for _, lib := range state.libraries { - if err := tx.Library(ctx).ScanEnd(lib.ID); err != nil { + if err := tx.Library().ScanEnd(ctx, lib.ID); err != nil { return fmt.Errorf("updating last scan completed for %s: %w", lib.Name, err) } - if err := tx.Property(ctx).Put(consts.PIDTrackKey, conf.Server.PID.Track); err != nil { + if err := tx.Property().Put(ctx, consts.PIDTrackKey, conf.Server.PID.Track); err != nil { return fmt.Errorf("updating track PID conf: %w", err) } - if err := tx.Property(ctx).Put(consts.PIDAlbumKey, conf.Server.PID.Album); err != nil { + if err := tx.Property().Put(ctx, consts.PIDAlbumKey, conf.Server.PID.Album); err != nil { return fmt.Errorf("updating album PID conf: %w", err) } if state.changesDetected.Load() { log.Debug(ctx, "Scanner: Refreshing library stats", "lib", lib.Name) - if err := tx.Library(ctx).RefreshStats(lib.ID); err != nil { + if err := tx.Library().RefreshStats(ctx, lib.ID); err != nil { return fmt.Errorf("refreshing library stats for %s: %w", lib.Name, err) } } else { diff --git a/scanner/scanner_benchmark_test.go b/scanner/scanner_benchmark_test.go index 65410d500..e6797df36 100644 --- a/scanner/scanner_benchmark_test.go +++ b/scanner/scanner_benchmark_test.go @@ -82,7 +82,7 @@ func BenchmarkScan(b *testing.B) { }) lib := model.Library{ID: 1, Name: "Fake Library", Path: "fake:///music"} - err := ds.Library(context.Background()).Put(&lib) + err := ds.Library().Put(b.Context(), &lib) if err != nil { b.Fatal(err) } diff --git a/scanner/scanner_multilibrary_test.go b/scanner/scanner_multilibrary_test.go index c0d5d4ece..546baf756 100644 --- a/scanner/scanner_multilibrary_test.go +++ b/scanner/scanner_multilibrary_test.go @@ -75,7 +75,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { IsAdmin: true, NewPassword: "password", } - Expect(ds.User(ctx).Put(&adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) s = scanner.New(ctx, ds, events.NoopBroker(), playlists.NewPlaylists(ds, artwork.NewUploader(ds)), metrics.NewNoopInstance()) @@ -83,8 +83,8 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { // Create two test libraries (let DB auto-assign IDs) lib1 = model.Library{Name: "Rock Collection", Path: "rock:///music"} lib2 = model.Library{Name: "Jazz Collection", Path: "jazz:///music"} - Expect(ds.Library(ctx).Put(&lib1)).To(Succeed()) - Expect(ds.Library(ctx).Put(&lib2)).To(Succeed()) + Expect(ds.Library().Put(ctx, &lib1)).To(Succeed()) + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) }) runScanner := func(ctx context.Context, fullScan bool) error { @@ -122,7 +122,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Check Rock library media files - rockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, Sort: "title", }) @@ -138,7 +138,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { } // Check Jazz library media files - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, Sort: "title", }) @@ -158,7 +158,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Check Rock library albums - rockAlbums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + rockAlbums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, Sort: "name", }) @@ -172,7 +172,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(rockAlbums[1].SongCount).To(Equal(2)) // Check Jazz library albums - jazzAlbums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + jazzAlbums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, Sort: "name", }) @@ -190,7 +190,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Check Rock library folders - rockFolders, err := ds.Folder(ctx).GetAll(model.QueryOptions{ + rockFolders, err := ds.Folder().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -201,7 +201,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { } // Check Jazz library folders - jazzFolders, err := ds.Folder(ctx).GetAll(model.QueryOptions{ + jazzFolders, err := ds.Folder().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -218,7 +218,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { // Check library-artist associations // Get all artists and check library associations - allArtists, err := ds.Artist(ctx).GetAll() + allArtists, err := ds.Artist().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) rockArtistNames := []string{} @@ -262,7 +262,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Check Rock library stats - rockLib, err := ds.Library(ctx).Get(lib1.ID) + rockLib, err := ds.Library().Get(ctx, lib1.ID) Expect(err).ToNot(HaveOccurred()) Expect(rockLib.TotalSongs).To(Equal(4)) Expect(rockLib.TotalAlbums).To(Equal(2)) @@ -271,7 +271,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(rockLib.TotalFolders).To(Equal(2)) // Abbey Road, IV (only folders with audio files) // Check Jazz library stats - jazzLib, err := ds.Library(ctx).Get(lib2.ID) + jazzLib, err := ds.Library().Get(ctx, lib2.ID) Expect(err).ToNot(HaveOccurred()) Expect(jazzLib.TotalSongs).To(Equal(4)) Expect(jazzLib.TotalAlbums).To(Equal(2)) @@ -285,25 +285,25 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify rock library stats - rockLib, err := ds.Library(ctx).Get(lib1.ID) + rockLib, err := ds.Library().Get(ctx, lib1.ID) Expect(err).ToNot(HaveOccurred()) Expect(rockLib.TotalSongs).To(Equal(4)) Expect(rockLib.TotalAlbums).To(Equal(2)) // Verify jazz library stats - jazzLib, err := ds.Library(ctx).Get(lib2.ID) + jazzLib, err := ds.Library().Get(ctx, lib2.ID) Expect(err).ToNot(HaveOccurred()) Expect(jazzLib.TotalSongs).To(Equal(4)) Expect(jazzLib.TotalAlbums).To(Equal(2)) // Verify that libraries don't interfere with each other - rockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) Expect(rockFiles).To(HaveLen(4)) - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -316,7 +316,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify that rock library only contains rock content - rockAlbums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + rockAlbums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -325,7 +325,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(rockAlbumNames).ToNot(ContainElements("Kind of Blue", "Giant Steps")) // Verify that jazz library only contains jazz content - jazzAlbums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + jazzAlbums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -365,7 +365,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { var rockCount, jazzCount int64 // Get Jeff Beck artist ID - jeffArtists, err := ds.Artist(ctx).GetAll(model.QueryOptions{ + jeffArtists, err := ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "Jeff Beck"}, }) Expect(err).ToNot(HaveOccurred()) @@ -389,14 +389,14 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(jazzCount).To(Equal(int64(1))) // Verify Jeff Beck albums are in correct libraries - rockAlbums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + rockAlbums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID, "album_artist": "Jeff Beck"}, }) Expect(err).ToNot(HaveOccurred()) Expect(rockAlbums).To(HaveLen(1)) Expect(rockAlbums[0].Name).To(Equal("Truth")) - jazzAlbums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + jazzAlbums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID, "album_artist": "Jeff Beck"}, }) Expect(err).ToNot(HaveOccurred()) @@ -426,13 +426,13 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify initial state - rockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) Expect(rockFiles).To(HaveLen(1)) - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -442,13 +442,13 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) // Verify counts remain the same - rockFiles, err = ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err = ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) Expect(rockFiles).To(HaveLen(1)) - jazzFiles, err = ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err = ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -485,7 +485,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) // Check that only the rock library file is marked as missing - missingRockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + missingRockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Eq{"library_id": lib1.ID}, squirrel.Eq{"missing": true}, @@ -496,7 +496,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(missingRockFiles[0].Title).To(Equal("Shoot to Thrill")) // Check that jazz library files are not affected - missingJazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + missingJazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Eq{"library_id": lib2.ID}, squirrel.Eq{"missing": true}, @@ -506,7 +506,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(missingJazzFiles).To(HaveLen(0)) // Verify non-missing files - presentRockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + presentRockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Eq{"library_id": lib1.ID}, squirrel.Eq{"missing": false}, @@ -548,7 +548,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(warnings).ToNot(BeEmpty(), "Should have warnings for filesystem errors") // Jazz library should have been scanned successfully - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -557,7 +557,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(jazzFiles[1].Title).To(BeElementOf("So What", "Freddie Freeloader")) // Rock library may have partial content (depending on scanner implementation) - rockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -565,12 +565,12 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { _ = rockFiles // Verify jazz library stats are correct - jazzLib, err := ds.Library(ctx).Get(lib2.ID) + jazzLib, err := ds.Library().Get(ctx, lib2.ID) Expect(err).ToNot(HaveOccurred()) Expect(jazzLib.TotalSongs).To(Equal(2)) // Error should be empty (warnings don't count as scan errors) - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "unset") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "unset") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(BeEmpty()) }) @@ -586,20 +586,20 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(warnings).ToNot(BeEmpty(), "Should have warnings for multiple filesystem errors") // Jazz library should be completely unaffected - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) Expect(jazzFiles).To(HaveLen(2)) // Jazz library statistics should be accurate - jazzLib, err := ds.Library(ctx).Get(lib2.ID) + jazzLib, err := ds.Library().Get(ctx, lib2.ID) Expect(err).ToNot(HaveOccurred()) Expect(jazzLib.TotalSongs).To(Equal(2)) Expect(jazzLib.TotalAlbums).To(Equal(1)) // Error should be empty (warnings don't count as scan errors) - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "unset") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "unset") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(BeEmpty()) }) @@ -623,7 +623,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { It("should propagate database errors and stop scanning", func() { // Install mock repo that injects DB error mfRepo := &mockMediaFileRepo{ - MediaFileRepository: ds.RealDS.MediaFile(ctx), + MediaFileRepository: ds.RealDS.MediaFile(), GetMissingAndMatchingError: errors.New("database connection failed"), } ds.MockedMediaFile = mfRepo @@ -632,7 +632,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, false)).To(MatchError(ContainSubstring("database connection failed"))) // Error should be recorded in scanner properties - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(ContainSubstring("database connection failed")) }) @@ -640,7 +640,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { It("should preserve error information in scanner properties", func() { // Install mock repo that injects DB error mfRepo := &mockMediaFileRepo{ - MediaFileRepository: ds.RealDS.MediaFile(ctx), + MediaFileRepository: ds.RealDS.MediaFile(), GetMissingAndMatchingError: errors.New("critical database error"), } ds.MockedMediaFile = mfRepo @@ -649,12 +649,12 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, false)).To(HaveOccurred()) // Check that error is recorded in scanner properties - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(ContainSubstring("critical database error")) // Scan type should still be recorded - scanType, _ := ds.Property(ctx).DefaultGet(consts.LastScanTypeKey, "") + scanType, _ := ds.Property().DefaultGet(ctx, consts.LastScanTypeKey, "") Expect(scanType).To(BeElementOf("incremental", "quick")) }) }) @@ -687,7 +687,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(warnings).ToNot(BeEmpty(), "Should have warnings for filesystem error") // Jazz library should scan completely successfully - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -695,13 +695,13 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(jazzFiles[0].Title).To(Equal("Chameleon")) // Jazz library statistics should be accurate - jazzLib, err := ds.Library(ctx).Get(lib2.ID) + jazzLib, err := ds.Library().Get(ctx, lib2.ID) Expect(err).ToNot(HaveOccurred()) Expect(jazzLib.TotalSongs).To(Equal(1)) Expect(jazzLib.TotalAlbums).To(Equal(1)) // Rock library may have partial content (depending on scanner implementation) - rockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) @@ -709,7 +709,7 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { _ = rockFiles // Error should be empty (warnings don't count as scan errors) - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "unset") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "unset") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(BeEmpty()) }) @@ -724,22 +724,22 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(warnings).ToNot(BeEmpty(), "Should have warnings for file corruption") // Verify that the working parts completed successfully - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) Expect(jazzFiles).To(HaveLen(1)) // Scanner properties should reflect successful completion despite warnings - scanType, _ := ds.Property(ctx).DefaultGet(consts.LastScanTypeKey, "") + scanType, _ := ds.Property().DefaultGet(ctx, consts.LastScanTypeKey, "") Expect(scanType).To(Equal("full")) // Start time should be recorded - startTimeStr, _ := ds.Property(ctx).DefaultGet(consts.LastScanStartTimeKey, "") + startTimeStr, _ := ds.Property().DefaultGet(ctx, consts.LastScanStartTimeKey, "") Expect(startTimeStr).ToNot(BeEmpty()) // Error should be empty (warnings don't count as scan errors) - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "unset") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "unset") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(BeEmpty()) }) @@ -780,30 +780,30 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(warnings).To(BeEmpty(), "Should have no warnings after error recovery") // Verify both libraries now have content (at least jazz should work) - rockFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + rockFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib1.ID}, }) Expect(err).ToNot(HaveOccurred()) // The scanner should recover and import both rock files Expect(len(rockFiles)).To(Equal(2)) - jazzFiles, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + jazzFiles, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) Expect(jazzFiles).To(HaveLen(1)) // Both libraries should have correct content counts - rockLib, err := ds.Library(ctx).Get(lib1.ID) + rockLib, err := ds.Library().Get(ctx, lib1.ID) Expect(err).ToNot(HaveOccurred()) Expect(rockLib.TotalSongs).To(Equal(2)) - jazzLib, err := ds.Library(ctx).Get(lib2.ID) + jazzLib, err := ds.Library().Get(ctx, lib2.ID) Expect(err).ToNot(HaveOccurred()) Expect(jazzLib.TotalSongs).To(Equal(1)) // Error should be empty (successful recovery) - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "unset") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "unset") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(BeEmpty()) }) @@ -822,15 +822,15 @@ var _ = Describe("Scanner - Multi-Library", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Validate properties - scanType, _ := ds.Property(ctx).DefaultGet(consts.LastScanTypeKey, "") + scanType, _ := ds.Property().DefaultGet(ctx, consts.LastScanTypeKey, "") Expect(scanType).To(Equal("full")) - startTimeStr, _ := ds.Property(ctx).DefaultGet(consts.LastScanStartTimeKey, "") + startTimeStr, _ := ds.Property().DefaultGet(ctx, consts.LastScanStartTimeKey, "") Expect(startTimeStr).ToNot(BeEmpty()) _, err := time.Parse(time.RFC3339, startTimeStr) Expect(err).ToNot(HaveOccurred()) - lastError, err := ds.Property(ctx).DefaultGet(consts.LastScanErrorKey, "unset") + lastError, err := ds.Property().DefaultGet(ctx, consts.LastScanErrorKey, "unset") Expect(err).ToNot(HaveOccurred()) Expect(lastError).To(BeEmpty()) }) diff --git a/scanner/scanner_selective_test.go b/scanner/scanner_selective_test.go index acaa8f850..2f27b74ce 100644 --- a/scanner/scanner_selective_test.go +++ b/scanner/scanner_selective_test.go @@ -63,13 +63,13 @@ var _ = Describe("ScanFolders", Ordered, func() { IsAdmin: true, NewPassword: "password", } - Expect(ds.User(ctx).Put(&adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) s = scanner.New(ctx, ds, events.NoopBroker(), playlists.NewPlaylists(ds, artwork.NewUploader(ds)), metrics.NewNoopInstance()) lib = model.Library{ID: 1, Name: "Fake Library", Path: "fake:///music"} - Expect(ds.Library(ctx).Put(&lib)).To(Succeed()) + Expect(ds.Library().Put(ctx, &lib)).To(Succeed()) // Initialize fake filesystem fsys = storagetest.FakeFS{} @@ -101,7 +101,7 @@ var _ = Describe("ScanFolders", Ordered, func() { Expect(warnings).To(BeEmpty()) // Verify all tracks in rock and jazz folders (including subdirectories) were imported - allFiles, err := ds.MediaFile(ctx).GetAll() + allFiles, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) // Should have 5 tracks (all rock and jazz tracks including subdirectories) @@ -123,7 +123,7 @@ var _ = Describe("ScanFolders", Ordered, func() { // Verify files in the pop folder were NOT scanned Expect(paths).ToNot(ContainElement("pop/track6.mp3")) - Expect(ds.Property(ctx).Get(consts.DBAnalyzePendingKey)).To(Equal("1")) + Expect(ds.Property().Get(ctx, consts.DBAnalyzePendingKey)).To(Equal("1")) }) }) @@ -135,26 +135,26 @@ var _ = Describe("ScanFolders", Ordered, func() { }) _, err := s.ScanAll(ctx, true) Expect(err).ToNot(HaveOccurred()) - Expect(ds.Property(ctx).Get(consts.DBAnalyzePendingKey)).To(Equal("0")) + Expect(ds.Property().Get(ctx, consts.DBAnalyzePendingKey)).To(Equal("0")) fsys.Add("rock/track2.mp3", rock(track(2, "Rock Track 2")), time.Now().Add(time.Second)) _, err = s.ScanAll(ctx, false) Expect(err).ToNot(HaveOccurred()) - Expect(ds.Property(ctx).Get(consts.DBAnalyzePendingKey)).To(Equal("0")) + Expect(ds.Property().Get(ctx, consts.DBAnalyzePendingKey)).To(Equal("0")) }) It("does not treat an interrupted scan in an untargeted library as a full scan", func() { otherLib := model.Library{ID: 2, Name: "Other Library", Path: "fake:///other"} - Expect(ds.Library(ctx).Put(&otherLib)).To(Succeed()) - Expect(ds.Library(ctx).ScanBegin(lib.ID, true)).To(Succeed()) + Expect(ds.Library().Put(ctx, &otherLib)).To(Succeed()) + Expect(ds.Library().ScanBegin(ctx, lib.ID, true)).To(Succeed()) lastAnalyze := "2026-07-09T12:00:00Z" - Expect(ds.Property(ctx).Put(consts.LastDBAnalyzeAtKey, lastAnalyze)).To(Succeed()) - Expect(ds.Property(ctx).Put(consts.DBAnalyzePendingKey, "0")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.LastDBAnalyzeAtKey, lastAnalyze)).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.DBAnalyzePendingKey, "0")).To(Succeed()) _, err := s.ScanFolders(ctx, false, []model.ScanTarget{{LibraryID: otherLib.ID, FolderPath: "."}}) Expect(err).ToNot(HaveOccurred()) - Expect(ds.Property(ctx).Get(consts.LastDBAnalyzeAtKey)).To(Equal(lastAnalyze)) + Expect(ds.Property().Get(ctx, consts.LastDBAnalyzeAtKey)).To(Equal(lastAnalyze)) }) }) @@ -187,7 +187,7 @@ var _ = Describe("ScanFolders", Ordered, func() { Expect(err).ToNot(HaveOccurred()) // Verify initial state - all folders exist - folders, err := ds.Folder(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"library_id": lib.ID}}) + folders, err := ds.Folder().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"library_id": lib.ID}}) Expect(err).ToNot(HaveOccurred()) Expect(folders).To(HaveLen(4)) // root, Artist, Album1, Album2 @@ -204,7 +204,7 @@ var _ = Describe("ScanFolders", Ordered, func() { } // Verify all tracks exist - allTracks, err := ds.MediaFile(ctx).GetAll() + allTracks, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(allTracks).To(HaveLen(4)) @@ -239,29 +239,29 @@ var _ = Describe("ScanFolders", Ordered, func() { Expect(err).ToNot(HaveOccurred()) // Verify the deleted child folder is now marked as missing - deletedFolder, err := ds.Folder(ctx).Get(album2FolderID) + deletedFolder, err := ds.Folder().Get(ctx, album2FolderID) Expect(err).ToNot(HaveOccurred()) Expect(deletedFolder.Missing).To(BeTrue(), "Deleted child folder should be marked as missing") // Verify the deleted folder's tracks are marked as missing for _, trackID := range album2TrackIDs { - track, err := ds.MediaFile(ctx).Get(trackID) + track, err := ds.MediaFile().Get(ctx, trackID) Expect(err).ToNot(HaveOccurred()) Expect(track.Missing).To(BeTrue(), "Track in deleted folder should be marked as missing") } // Verify the parent folder is still present and not marked as missing - parentFolder, err := ds.Folder(ctx).Get(artistFolderID) + parentFolder, err := ds.Folder().Get(ctx, artistFolderID) Expect(err).ToNot(HaveOccurred()) Expect(parentFolder.Missing).To(BeFalse(), "Parent folder should not be marked as missing") // Verify the sibling folder and its tracks are still present and not missing - siblingFolder, err := ds.Folder(ctx).Get(album1FolderID) + siblingFolder, err := ds.Folder().Get(ctx, album1FolderID) Expect(err).ToNot(HaveOccurred()) Expect(siblingFolder.Missing).To(BeFalse(), "Sibling folder should not be marked as missing") for _, trackID := range album1TrackIDs { - track, err := ds.MediaFile(ctx).Get(trackID) + track, err := ds.MediaFile().Get(ctx, trackID) Expect(err).ToNot(HaveOccurred()) Expect(track.Missing).To(BeFalse(), "Track in sibling folder should not be marked as missing") } @@ -283,7 +283,7 @@ var _ = Describe("ScanFolders", Ordered, func() { Expect(err).ToNot(HaveOccurred()) // Verify nested folders were created - allFolders, err := ds.Folder(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"library_id": lib.ID}}) + allFolders, err := ds.Folder().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"library_id": lib.ID}}) Expect(err).ToNot(HaveOccurred()) Expect(len(allFolders)).To(BeNumerically(">", 4), "Should have more folders with nested structure") @@ -301,7 +301,7 @@ var _ = Describe("ScanFolders", Ordered, func() { Expect(err).ToNot(HaveOccurred()) // Verify all Help! folders (including nested ones) are marked as missing - missingFolders, err := ds.Folder(ctx).GetAll(model.QueryOptions{ + missingFolders, err := ds.Folder().GetAll(ctx, model.QueryOptions{ Filters: squirrel.And{ squirrel.Eq{"library_id": lib.ID}, squirrel.Eq{"missing": true}, @@ -311,7 +311,7 @@ var _ = Describe("ScanFolders", Ordered, func() { Expect(len(missingFolders)).To(BeNumerically(">", 0), "At least one folder should be marked as missing") // Verify all tracks in deleted folders are marked as missing - allTracks, err := ds.MediaFile(ctx).GetAll() + allTracks, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(allTracks).To(HaveLen(6)) diff --git a/scanner/scanner_test.go b/scanner/scanner_test.go index 30f4a2b97..4ce8cce1e 100644 --- a/scanner/scanner_test.go +++ b/scanner/scanner_test.go @@ -78,7 +78,7 @@ var _ = Describe("Scanner", Ordered, func() { ds = &tests.MockDataStore{RealDS: persistence.New(db.Db())} mfRepo = &mockMediaFileRepo{ - MediaFileRepository: ds.RealDS.MediaFile(ctx), + MediaFileRepository: ds.RealDS.MediaFile(), } ds.MockedMediaFile = mfRepo @@ -90,13 +90,13 @@ var _ = Describe("Scanner", Ordered, func() { IsAdmin: true, NewPassword: "password", } - Expect(ds.User(ctx).Put(&adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) s = scanner.New(ctx, ds, events.NoopBroker(), playlists.NewPlaylists(ds, artwork.NewUploader(ds)), metrics.NewNoopInstance()) lib = model.Library{ID: 1, Name: "Fake Library", Path: "fake:///music"} - Expect(ds.Library(ctx).Put(&lib)).To(Succeed()) + Expect(ds.Library().Put(ctx, &lib)).To(Succeed()) }) runScanner := func(ctx context.Context, fullScan bool) error { @@ -108,14 +108,14 @@ var _ = Describe("Scanner", Ordered, func() { // so a later scan can only queue genuine reprocessing. resolveQueuedArtwork := func() []model.ArtworkQueueItem { GinkgoHelper() - queued, err := ds.ArtworkQueue(ctx).DequeueBatch(1000) + queued, err := ds.ArtworkQueue().DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) for _, it := range queued { - Expect(ds.Artwork(ctx).PutItemArtwork(&model.ItemArtwork{ + Expect(ds.Artwork().PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: it.ItemKind, ItemID: it.ItemID, ImageType: it.ImageType, Hash: "resolved", Source: "embedded", UpdatedAt: time.Now(), })).To(Succeed()) - Expect(ds.ArtworkQueue(ctx).DeleteIfUnchanged(it.ItemKind, it.ItemID, it.ImageType, it.RetryAt)).To(Succeed()) + Expect(ds.ArtworkQueue().DeleteIfUnchanged(ctx, it.ItemKind, it.ItemID, it.ImageType, it.RetryAt)).To(Succeed()) } return queued } @@ -140,7 +140,7 @@ var _ = Describe("Scanner", Ordered, func() { It("should import all folders", func() { Expect(runScanner(ctx, true)).To(Succeed()) - folders, _ := ds.Folder(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"library_id": lib.ID}}) + folders, _ := ds.Folder().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"library_id": lib.ID}}) paths := slice.Map(folders, func(f model.Folder) string { return f.Name }) Expect(paths).To(SatisfyAll( HaveLen(4), @@ -150,7 +150,7 @@ var _ = Describe("Scanner", Ordered, func() { It("should import all mediafiles", func() { Expect(runScanner(ctx, true)).To(Succeed()) - mfs, _ := ds.MediaFile(ctx).GetAll() + mfs, _ := ds.MediaFile().GetAll(ctx) paths := slice.Map(mfs, func(f model.MediaFile) string { return f.Title }) Expect(paths).To(SatisfyAll( HaveLen(7), @@ -163,7 +163,7 @@ var _ = Describe("Scanner", Ordered, func() { It("should import all albums", func() { Expect(runScanner(ctx, true)).To(Succeed()) - albums, _ := ds.Album(ctx).GetAll(model.QueryOptions{Sort: "name"}) + albums, _ := ds.Album().GetAll(ctx, model.QueryOptions{Sort: "name"}) Expect(albums).To(HaveLen(2)) Expect(albums[0]).To(SatisfyAll( HaveField("Name", Equal("Help!")), @@ -177,9 +177,9 @@ var _ = Describe("Scanner", Ordered, func() { It("should enqueue artwork resolution for the scanned albums and artists", func() { Expect(runScanner(ctx, true)).To(Succeed()) - albums, _ := ds.Album(ctx).GetAll() - artists, _ := ds.Artist(ctx).GetAll(model.QueryOptions{Filters: squirrel.NotEq{"name": consts.UnknownArtist}}) - queued, err := ds.ArtworkQueue(ctx).DequeueBatch(1000) + albums, _ := ds.Album().GetAll(ctx) + artists, _ := ds.Artist().GetAll(ctx, model.QueryOptions{Filters: squirrel.NotEq{"name": consts.UnknownArtist}}) + queued, err := ds.ArtworkQueue().DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) for _, al := range albums { @@ -204,7 +204,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) - requeued, err := ds.ArtworkQueue(ctx).DequeueBatch(1000) + requeued, err := ds.ArtworkQueue().DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(requeued).To(BeEmpty()) }) @@ -213,14 +213,14 @@ var _ = Describe("Scanner", Ordered, func() { It("should update the media_file", func() { Expect(runScanner(ctx, true)).To(Succeed()) - mf, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"title": "Help!"}}) + mf, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"title": "Help!"}}) Expect(err).ToNot(HaveOccurred()) Expect(mf[0].Tags).ToNot(HaveKey("barcode")) fsys.UpdateTags("The Beatles/Help!/01 - Help!.mp3", _t{"barcode": "123"}) Expect(runScanner(ctx, true)).To(Succeed()) - mf, err = ds.MediaFile(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"title": "Help!"}}) + mf, err = ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"title": "Help!"}}) Expect(err).ToNot(HaveOccurred()) Expect(mf[0].Tags).To(HaveKeyWithValue(model.TagName("barcode"), []string{"123"})) }) @@ -234,9 +234,9 @@ var _ = Describe("Scanner", Ordered, func() { fsys.UpdateTags("The Beatles/Help!/01 - Help!.mp3", _t{"producer": "George Martin"}) Expect(runScanner(ctx, false)).To(Succeed()) - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album.name": "Help!"}}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album.name": "Help!"}}) Expect(err).ToNot(HaveOccurred()) - requeued, err := ds.ArtworkQueue(ctx).DequeueBatch(1000) + requeued, err := ds.ArtworkQueue().DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(requeued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "al"), @@ -248,7 +248,7 @@ var _ = Describe("Scanner", Ordered, func() { tests.SkipOnWindows("path separator bug (#TBD-path-sep-scanner)") Expect(runScanner(ctx, true)).To(Succeed()) - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album.name": "Help!"}}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album.name": "Help!"}}) Expect(err).ToNot(HaveOccurred()) Expect(albums).ToNot(BeEmpty()) Expect(albums[0].Participants.First(model.RoleProducer).Name).To(BeEmpty()) @@ -257,7 +257,7 @@ var _ = Describe("Scanner", Ordered, func() { fsys.UpdateTags("The Beatles/Help!/01 - Help!.mp3", _t{"producer": "George Martin"}) Expect(runScanner(ctx, false)).To(Succeed()) - albums, err = ds.Album(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album.name": "Help!"}}) + albums, err = ds.Album().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album.name": "Help!"}}) Expect(err).ToNot(HaveOccurred()) Expect(albums[0].Participants.First(model.RoleProducer).Name).To(Equal("George Martin")) Expect(albums[0].SongCount).To(Equal(3)) @@ -266,12 +266,12 @@ var _ = Describe("Scanner", Ordered, func() { It("invalidates the media_file artwork state so new embedded art is picked up lazily", func() { Expect(runScanner(ctx, true)).To(Succeed()) - mf, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"title": "Help!"}}) + mf, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"title": "Help!"}}) Expect(err).ToNot(HaveOccurred()) Expect(mf).ToNot(BeEmpty()) trackID := mf[0].ID - Expect(ds.Artwork(ctx).PutItemArtwork(&model.ItemArtwork{ + Expect(ds.Artwork().PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "mf", ItemID: trackID, ImageType: model.ImageTypePrimary, Source: "embedded", Hash: "stalehash", })).To(Succeed()) @@ -279,7 +279,7 @@ var _ = Describe("Scanner", Ordered, func() { fsys.UpdateTags("The Beatles/Help!/01 - Help!.mp3", _t{"comment": "reimport"}) Expect(runScanner(ctx, true)).To(Succeed()) - _, err = ds.Artwork(ctx).GetItemArtwork(model.KindMediaFileArtwork, trackID, model.ImageTypePrimary) + _, err = ds.Artwork().GetItemArtwork(ctx, model.KindMediaFileArtwork, trackID, model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) }) }) @@ -291,21 +291,21 @@ var _ = Describe("Scanner", Ordered, func() { albumID := func(name string) string { GinkgoHelper() - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album.name": name}}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album.name": name}}) Expect(err).ToNot(HaveOccurred()) Expect(albums).To(HaveLen(1)) return albums[0].ID } artistID := func(name string) string { GinkgoHelper() - artists, err := ds.Artist(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"artist.name": name}}) + artists, err := ds.Artist().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"artist.name": name}}) Expect(err).ToNot(HaveOccurred()) Expect(artists).To(HaveLen(1)) return artists[0].ID } queuedItems := func() []model.ArtworkQueueItem { GinkgoHelper() - queued, err := ds.ArtworkQueue(ctx).DequeueBatch(1000) + queued, err := ds.ArtworkQueue().DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) return queued } @@ -464,7 +464,7 @@ var _ = Describe("Scanner", Ordered, func() { It("should not import the ignored file", func() { Expect(runScanner(ctx, true)).To(Succeed()) - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(1)) for _, mf := range mfs { @@ -486,11 +486,11 @@ var _ = Describe("Scanner", Ordered, func() { It("should import as one album", func() { Expect(runScanner(ctx, true)).To(Succeed()) - albums, err := ds.Album(ctx).GetAll() + albums, err := ds.Album().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(albums).To(HaveLen(1)) - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(2)) for _, mf := range mfs { @@ -512,7 +512,7 @@ var _ = Describe("Scanner", Ordered, func() { It("should import as two distinct albums", func() { Expect(runScanner(ctx, true)).To(Succeed()) - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Sort: "release_date"}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Sort: "release_date"}) Expect(err).ToNot(HaveOccurred()) Expect(albums).To(HaveLen(2)) Expect(albums[0]).To(SatisfyAll( @@ -551,7 +551,7 @@ var _ = Describe("Scanner", Ordered, func() { By("Doing a full scan") Expect(runScanner(ctx, true)).To(Succeed()) - Expect(ds.MediaFile(ctx).CountAll()).To(Equal(int64(4))) + Expect(ds.MediaFile().CountAll(ctx)).To(Equal(int64(4))) findByPath = createFindByPath(ctx, ds) }) @@ -559,7 +559,7 @@ var _ = Describe("Scanner", Ordered, func() { fsys.Add("The Beatles/Revolver/03 - I'm Only Sleeping.mp3", revolver(track(3, "I'm Only Sleeping"))) Expect(runScanner(ctx, false)).To(Succeed()) - Expect(ds.MediaFile(ctx).CountAll()).To(Equal(int64(5))) + Expect(ds.MediaFile().CountAll(ctx)).To(Equal(int64(5))) mf, err := findByPath("The Beatles/Revolver/03 - I'm Only Sleeping.mp3") Expect(err).ToNot(HaveOccurred()) Expect(mf.Title).To(Equal("I'm Only Sleeping")) @@ -569,7 +569,7 @@ var _ = Describe("Scanner", Ordered, func() { fsys.UpdateTags("The Beatles/Revolver/02 - Eleanor Rigby.mp3", _t{"title": "Eleanor Rigby (remix)"}) Expect(runScanner(ctx, false)).To(Succeed()) - Expect(ds.MediaFile(ctx).CountAll()).To(Equal(int64(4))) + Expect(ds.MediaFile().CountAll(ctx)).To(Equal(int64(4))) mf, _ := findByPath("The Beatles/Revolver/02 - Eleanor Rigby.mp3") Expect(mf.Title).To(Equal("Eleanor Rigby (remix)")) }) @@ -578,7 +578,7 @@ var _ = Describe("Scanner", Ordered, func() { fsys.Add("The Beatles/Revolver/01 - Taxman.mp3", revolver(track(1, "Taxman", _t{"bitrate": 640}))) Expect(runScanner(ctx, false)).To(Succeed()) - Expect(ds.MediaFile(ctx).CountAll()).To(Equal(int64(4))) + Expect(ds.MediaFile().CountAll(ctx)).To(Equal(int64(4))) mf, _ := findByPath("The Beatles/Revolver/01 - Taxman.mp3") Expect(mf.BitRate).To(Equal(640)) }) @@ -591,7 +591,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the file is marked as missing") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(3))) mf, err := findByPath("The Beatles/Revolver/02 - Eleanor Rigby.mp3") @@ -612,14 +612,14 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the old file is not in the library") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(4))) _, err = findByPath("The Beatles/Revolver/02 - Eleanor Rigby.mp3") Expect(err).To(MatchError(model.ErrNotFound)) By("Checking the new file is in the library") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, })).To(BeZero()) mf, err := findByPath("The Beatles/Help!/02 - Eleanor Rigby.mp3") @@ -641,7 +641,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(MatchError(ContainSubstring("I/O read error"))) By("Checking the both instances of the file are in the lib") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Taxman"}, })).To(Equal(int64(2))) @@ -650,7 +650,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the old file is not in the library") - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Taxman"}, }) Expect(err).ToNot(HaveOccurred()) @@ -671,14 +671,14 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the old file is not in the library") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, })).To(BeZero()) _, err = findByPath("The Beatles/Revolver/02 - Eleanor Rigby.mp3") Expect(err).To(MatchError(model.ErrNotFound)) By("Checking the new file is in the library") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(4))) mf, err := findByPath("The Beatles/Revolver/02 - Eleanor Rigby.flac") @@ -698,7 +698,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the file is marked as missing") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(3))) mf, err := findByPath("The Beatles/Revolver/02 - Eleanor Rigby.mp3") @@ -712,7 +712,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the file is not marked as missing") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(4))) mf, err = findByPath("The Beatles/Revolver/02 - Eleanor Rigby.mp3") @@ -737,7 +737,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking the file was found in the new folder") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(4))) mf, err = findByPath("The Beatles/Help!/02 - Eleanor Rigby.mp3") @@ -751,7 +751,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Verifying initial state has 5 tracks") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(5))) @@ -790,7 +790,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(mf.Missing).To(BeFalse()) By("Verifying only 2 non-missing tracks remain (Help! tracks)") - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{ + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": false}, })).To(Equal(int64(2))) }) @@ -815,7 +815,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) nonMissingArtists := func() []string { - aa, err := ds.Artist(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"missing": false}}) + aa, err := ds.Artist().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"missing": false}}) Expect(err).ToNot(HaveOccurred()) return slice.Map(aa, func(a model.Artist) string { return a.Name }) } @@ -860,7 +860,7 @@ var _ = Describe("Scanner", Ordered, func() { It("does not override artist fields when importing an undertagged file", func() { By("Making sure artist in the DB contains MBID and sort name") - aa, err := ds.Artist(ctx).GetAll(model.QueryOptions{ + aa, err := ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "The Beatles"}, }) Expect(err).ToNot(HaveOccurred()) @@ -887,7 +887,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(mf.SortArtistName).To(BeEmpty()) By("Makingsure the artist in the DB has not changed") - aa, err = ds.Artist(ctx).GetAll(model.QueryOptions{ + aa, err = ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "The Beatles"}, }) Expect(err).ToNot(HaveOccurred()) @@ -915,7 +915,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) By("Checking files are marked as missing but not deleted") - count, err := ds.MediaFile(ctx).CountAll(model.QueryOptions{ + count, err := ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, }) Expect(err).ToNot(HaveOccurred()) @@ -943,7 +943,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking missing files are deleted") - count, err := ds.MediaFile(ctx).CountAll(model.QueryOptions{ + count, err := ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, }) Expect(err).ToNot(HaveOccurred()) @@ -970,7 +970,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) By("Checking files are marked as missing but not deleted") - count, err := ds.MediaFile(ctx).CountAll(model.QueryOptions{ + count, err := ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, }) Expect(err).ToNot(HaveOccurred()) @@ -992,7 +992,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) By("Checking missing files are deleted") - count, err := ds.MediaFile(ctx).CountAll(model.QueryOptions{ + count, err := ds.MediaFile().CountAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"missing": true}, }) Expect(err).ToNot(HaveOccurred()) @@ -1020,10 +1020,10 @@ var _ = Describe("Scanner", Ordered, func() { simulateInterruptedScan := func(fullScan bool) { // Call ScanBegin to properly set LastScanStartedAt and FullScanInProgress // This simulates what would happen if a scan was interrupted (ScanBegin called but ScanEnd not) - Expect(ds.Library(ctx).ScanBegin(lib.ID, fullScan)).To(Succeed()) + Expect(ds.Library().ScanBegin(ctx, lib.ID, fullScan)).To(Succeed()) // Verify the update was persisted - reloaded, err := ds.Library(ctx).Get(lib.ID) + reloaded, err := ds.Library().Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(reloaded.LastScanStartedAt).ToNot(BeZero()) Expect(reloaded.FullScanInProgress).To(Equal(fullScan)) @@ -1035,7 +1035,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify files were imported - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(2)) @@ -1056,7 +1056,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify the comment was updated (which means the folder was processed and file re-imported) - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Help!"}, }) Expect(err).ToNot(HaveOccurred()) @@ -1071,7 +1071,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify files were imported - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(2)) @@ -1090,7 +1090,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) // Verify the comment was updated (folder was processed despite unchanged hash) - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Help!"}, }) Expect(err).ToNot(HaveOccurred()) @@ -1105,12 +1105,12 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify files were imported - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(mfs).To(HaveLen(2)) // Library should have LastScanStartedAt cleared after successful scan - updatedLib, err := ds.Library(ctx).Get(lib.ID) + updatedLib, err := ds.Library().Get(ctx, lib.ID) Expect(err).ToNot(HaveOccurred()) Expect(updatedLib.LastScanStartedAt).To(BeZero()) Expect(updatedLib.FullScanInProgress).To(BeFalse()) @@ -1125,7 +1125,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify the comment was updated - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Help!"}, }) Expect(err).ToNot(HaveOccurred()) @@ -1144,7 +1144,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, false)).To(Succeed()) // Verify the comment was NOT updated (folder was skipped) - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Help!"}, }) Expect(err).ToNot(HaveOccurred()) @@ -1163,7 +1163,7 @@ var _ = Describe("Scanner", Ordered, func() { refreshStatsCalls = nil // Create a mock artist repository that tracks RefreshStats calls - originalArtistRepo := ds.RealDS.Artist(ctx) + originalArtistRepo := ds.RealDS.Artist() ds.MockedArtist = &testArtistRepo{ ArtistRepository: originalArtistRepo, callTracker: &refreshStatsCalls, @@ -1209,7 +1209,7 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).To(Succeed()) // Verify initial artist stats - should have 1 album, 1 song - artists, err := ds.Artist(ctx).GetAll(model.QueryOptions{ + artists, err := ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "The Beatles"}, }) Expect(err).ToNot(HaveOccurred()) @@ -1228,7 +1228,7 @@ var _ = Describe("Scanner", Ordered, func() { By("Verifying artist stats were updated correctly") // Fetch the artist again to check updated stats - artists, err = ds.Artist(ctx).GetAll(model.QueryOptions{ + artists, err = ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "The Beatles"}, }) Expect(err).ToNot(HaveOccurred()) @@ -1277,8 +1277,8 @@ var _ = Describe("Scanner", Ordered, func() { Expect(runScanner(ctx, true)).ToNot(Succeed()) - Expect(ds.Folder(ctx).CountAll(model.QueryOptions{Filters: squirrel.Eq{"missing": true}})).To(BeZero()) - Expect(ds.MediaFile(ctx).CountAll(model.QueryOptions{Filters: squirrel.Eq{"missing": true}})).To(BeZero()) + Expect(ds.Folder().CountAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"missing": true}})).To(BeZero()) + Expect(ds.MediaFile().CountAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"missing": true}})).To(BeZero()) }) }) }) @@ -1298,7 +1298,7 @@ func (b *busyPersistDS) WithTxRetry(ctx context.Context, block func(context.Cont func createFindByPath(ctx context.Context, ds model.DataStore) func(string) (*model.MediaFile, error) { return func(path string) (*model.MediaFile, error) { - list, err := ds.MediaFile(ctx).FindByPaths([]string{path}) + list, err := ds.MediaFile().FindByPaths(ctx, []string{path}) if err != nil { return nil, err } @@ -1315,16 +1315,16 @@ type mockMediaFileRepo struct { cursorCalls atomic.Int32 } -func (m *mockMediaFileRepo) GetCursor(options ...model.QueryOptions) (model.MediaFileCursor, error) { +func (m *mockMediaFileRepo) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.MediaFileCursor, error) { m.cursorCalls.Add(1) - return m.MediaFileRepository.GetCursor(options...) + return m.MediaFileRepository.GetCursor(ctx, options...) } -func (m *mockMediaFileRepo) GetMissingAndMatching(libId int) (model.MediaFileCursor, error) { +func (m *mockMediaFileRepo) GetMissingAndMatching(ctx context.Context, libId int) (model.MediaFileCursor, error) { if m.GetMissingAndMatchingError != nil { return nil, m.GetMissingAndMatchingError } - return m.MediaFileRepository.GetMissingAndMatching(libId) + return m.MediaFileRepository.GetMissingAndMatching(ctx, libId) } type testArtistRepo struct { @@ -1332,7 +1332,7 @@ type testArtistRepo struct { callTracker *[]bool } -func (m *testArtistRepo) RefreshStats(allArtists bool) (int64, error) { +func (m *testArtistRepo) RefreshStats(ctx context.Context, allArtists bool) (int64, error) { *m.callTracker = append(*m.callTracker, allArtists) - return m.ArtistRepository.RefreshStats(allArtists) + return m.ArtistRepository.RefreshStats(ctx, allArtists) } diff --git a/scanner/watcher.go b/scanner/watcher.go index baf94b79b..1ac5468f0 100644 --- a/scanner/watcher.go +++ b/scanner/watcher.go @@ -23,7 +23,7 @@ type Watcher interface { } type watcher struct { - mainCtx context.Context + mainCtx context.Context //nolint:containedctx // watcher lifecycle ctx ds model.DataStore scanner model.Scanner triggerWait time.Duration @@ -60,7 +60,7 @@ func (w *watcher) Run(ctx context.Context) error { w.mainCtx = ctx // Start watchers for all existing libraries - libs, err := w.ds.Library(ctx).GetAll() + libs, err := w.ds.Library().GetAll(ctx) if err != nil { return fmt.Errorf("getting libraries: %w", err) } diff --git a/server/auth.go b/server/auth.go index 2aaa93e63..3e58359da 100644 --- a/server/auth.go +++ b/server/auth.go @@ -49,7 +49,7 @@ func login(ds model.DataStore) func(w http.ResponseWriter, r *http.Request) { } func doLogin(ds model.DataStore, username string, password string, w http.ResponseWriter, r *http.Request) { - user, err := validateLogin(ds.User(r.Context()), username, password) + user, err := validateLogin(r.Context(), ds.User(), username, password) if err != nil { _ = rest.RespondWithError(w, http.StatusInternalServerError, "Unknown error authentication user. Please try again") return @@ -127,7 +127,7 @@ func createAdmin(ds model.DataStore) func(w http.ResponseWriter, r *http.Request _ = rest.RespondWithError(w, http.StatusUnprocessableEntity, err.Error()) return } - c, err := ds.User(r.Context()).CountAll() + c, err := ds.User().CountAll(r.Context()) if err != nil { _ = rest.RespondWithError(w, http.StatusInternalServerError, err.Error()) return @@ -157,7 +157,7 @@ func createAdminUser(ctx context.Context, ds model.DataStore, username, password IsAdmin: true, LastLoginAt: new(time.Now()), } - err := ds.User(ctx).Put(&initialUser) + err := ds.User().Put(ctx, &initialUser) if err != nil { log.Error(ctx, "Could not create initial user", "user", initialUser.UserName, err) return fmt.Errorf("creating initial user: %w", err) @@ -165,8 +165,8 @@ func createAdminUser(ctx context.Context, ds model.DataStore, username, password return nil } -func validateLogin(userRepo model.UserRepository, userName, password string) (*model.User, error) { - u, err := userRepo.FindByUsernameWithPassword(userName) +func validateLogin(ctx context.Context, userRepo model.UserRepository, userName, password string) (*model.User, error) { + u, err := userRepo.FindByUsernameWithPassword(ctx, userName) if errors.Is(err, model.ErrNotFound) { return nil, nil } @@ -176,9 +176,9 @@ func validateLogin(userRepo model.UserRepository, userName, password string) (*m if u.Password != password { return nil, nil } - err = userRepo.UpdateLastLoginAt(u.ID) + err = userRepo.UpdateLastLoginAt(ctx, u.ID) if err != nil { - log.Error("Could not update LastLoginAt", "user", userName) + log.Error(ctx, "Could not update LastLoginAt", "user", userName) } return u, nil } @@ -244,7 +244,7 @@ func UsernameFromConfig(*http.Request) string { } func contextWithUser(ctx context.Context, ds model.DataStore, username string) (context.Context, error) { - user, err := ds.User(ctx).FindByUsername(username) + user, err := ds.User().FindByUsername(ctx, username) if err == nil { ctx = log.NewContext(ctx, "username", username) ctx = request.WithUsername(ctx, user.UserName) @@ -309,7 +309,7 @@ func tokenAllowed(ctx context.Context) bool { // epoch the handler bumped reaches the token the client stores. type refreshingWriter struct { http.ResponseWriter - ctx context.Context + ctx context.Context //nolint:containedctx // ResponseWriter wrapper defers work to Write, which has no ctx token jwt.Token once sync.Once } @@ -377,12 +377,13 @@ func handleLoginFromHeaders(ds model.DataStore, r *http.Request) map[string]any } } - userRepo := ds.User(r.Context()) - user, err := userRepo.FindByUsernameWithPassword(username) + ctx := r.Context() + userRepo := ds.User() + user, err := userRepo.FindByUsernameWithPassword(ctx, username) if user == nil || err != nil { log.Info(r, "User passed in header not found", "user", username) // Check if this is the first user being created - count, _ := userRepo.CountAll() + count, _ := userRepo.CountAll(ctx) isFirstUser := count == 0 newUser := model.User{ @@ -393,19 +394,19 @@ func handleLoginFromHeaders(ds model.DataStore, r *http.Request) map[string]any NewPassword: consts.PasswordAutogenPrefix + id.NewRandom(), IsAdmin: isFirstUser, // Make the first user an admin } - err := userRepo.Put(&newUser) + err := userRepo.Put(ctx, &newUser) if err != nil { log.Error(r, "Could not create new user", "user", username, err) return nil } - user, err = userRepo.FindByUsernameWithPassword(username) + user, err = userRepo.FindByUsernameWithPassword(ctx, username) if user == nil || err != nil { log.Error(r, "Created user but failed to fetch it", "user", username) return nil } } - err = userRepo.UpdateLastLoginAt(user.ID) + err = userRepo.UpdateLastLoginAt(ctx, user.ID) if err != nil { log.Error(r, "Could not update LastLoginAt", "user", username, err) return nil diff --git a/server/auth_test.go b/server/auth_test.go index abe144a12..1095fafc9 100644 --- a/server/auth_test.go +++ b/server/auth_test.go @@ -28,6 +28,12 @@ import ( ) var _ = Describe("Auth", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) + Describe("User login", func() { var ds model.DataStore var req *http.Request @@ -48,8 +54,8 @@ var _ = Describe("Auth", func() { }) It("creates an admin user with the specified password", func() { - usr := ds.User(context.Background()) - u, err := usr.FindByUsername("johndoe") + usr := ds.User() + u, err := usr.FindByUsername(ctx, "johndoe") Expect(err).To(BeNil()) Expect(u.Password).ToNot(BeEmpty()) Expect(u.IsAdmin).To(BeTrue()) @@ -99,8 +105,8 @@ var _ = Describe("Auth", func() { fs := os.DirFS("tests/fixtures") BeforeEach(func() { - usr := ds.User(context.Background()) - _ = usr.Put(&model.User{ID: "111", UserName: "janedoe", NewPassword: "abc123", Name: "Jane", IsAdmin: false}) + usr := ds.User() + _ = usr.Put(ctx, &model.User{ID: "111", UserName: "janedoe", NewPassword: "abc123", Name: "Jane", IsAdmin: false}) req = httptest.NewRequest("GET", "/index.html", nil) req.Header.Add("Remote-User", "janedoe") resp = httptest.NewRecorder() @@ -232,8 +238,8 @@ var _ = Describe("Auth", func() { }) It("logs in successfully if user exists", func() { - usr := ds.User(context.Background()) - _ = usr.Put(&model.User{ID: "111", UserName: "janedoe", NewPassword: "abc123", Name: "Jane", IsAdmin: false}) + usr := ds.User() + _ = usr.Put(ctx, &model.User{ID: "111", UserName: "janedoe", NewPassword: "abc123", Name: "Jane", IsAdmin: false}) login(ds)(resp, req) Expect(resp.Code).To(Equal(http.StatusOK)) @@ -397,14 +403,14 @@ var _ = Describe("Auth", func() { Expect(result["isAdmin"]).To(BeTrue()) // Verify user was created as admin - u, err := ds.User(context.Background()).FindByUsername("firstuser") + u, err := ds.User().FindByUsername(ctx, "firstuser") Expect(err).To(BeNil()) Expect(u.IsAdmin).To(BeTrue()) }) It("does not make subsequent users admins", func() { // Create the first user - _ = ds.User(context.Background()).Put(&model.User{ + _ = ds.User().Put(ctx, &model.User{ ID: "existing-user-id", UserName: "existinguser", Name: "Existing User", @@ -419,7 +425,7 @@ var _ = Describe("Auth", func() { Expect(result["isAdmin"]).To(BeFalse()) // Verify user was created as non-admin - u, err := ds.User(context.Background()).FindByUsername("seconduser") + u, err := ds.User().FindByUsername(ctx, "seconduser") Expect(err).To(BeNil()) Expect(u.IsAdmin).To(BeFalse()) }) @@ -434,9 +440,9 @@ var _ = Describe("Auth", func() { conf.Server.SessionTimeout = time.Hour ds = &tests.MockDataStore{} auth.Init(ds) - ur := ds.User(context.TODO()).(*tests.MockedUserRepo) + ur := ds.User().(*tests.MockedUserRepo) usr = &model.User{ID: "u1", UserName: "johndoe", NewPassword: "pw", TokenEpoch: 2} - Expect(ur.Put(usr)).To(Succeed()) + Expect(ur.Put(ctx, usr)).To(Succeed()) }) serve := func(token string) *httptest.ResponseRecorder { diff --git a/server/events/sse.go b/server/events/sse.go index 565d8c016..e4d6a05e5 100644 --- a/server/events/sse.go +++ b/server/events/sse.go @@ -34,7 +34,7 @@ type ( id uint64 event string data string - senderCtx context.Context + senderCtx context.Context //nolint:containedctx // queued message carries the sender ctx } messageChan chan message clientsChan chan client diff --git a/server/initial_setup.go b/server/initial_setup.go index e75220abe..be9e14ae9 100644 --- a/server/initial_setup.go +++ b/server/initial_setup.go @@ -17,23 +17,23 @@ import ( func initialSetup(ds model.DataStore) { ctx := context.TODO() err := ds.WithTx(func(tx model.DataStore) error { - if err := tx.Library(ctx).StoreMusicFolder(); err != nil { + if err := tx.Library().StoreMusicFolder(ctx); err != nil { return err } - properties := tx.Property(ctx) - _, err := properties.Get(consts.InitialSetupFlagKey) + properties := tx.Property() + _, err := properties.Get(ctx, consts.InitialSetupFlagKey) if err == nil { return nil } log.Info("Running initial setup") if conf.Server.DevAutoCreateAdminPassword != "" { - if err = createInitialAdminUser(tx, conf.Server.DevAutoCreateAdminPassword); err != nil { + if err = createInitialAdminUser(ctx, tx, conf.Server.DevAutoCreateAdminPassword); err != nil { return err } } - err = properties.Put(consts.InitialSetupFlagKey, time.Now().String()) + err = properties.Put(ctx, consts.InitialSetupFlagKey, time.Now().String()) return err }, "initial setup") if err != nil { @@ -42,9 +42,9 @@ func initialSetup(ds model.DataStore) { } // If the Dev Admin user is not present, create it -func createInitialAdminUser(ds model.DataStore, initialPassword string) error { - users := ds.User(context.TODO()) - c, err := users.CountAll(model.QueryOptions{Filters: squirrel.Eq{"user_name": consts.DevInitialUserName}}) +func createInitialAdminUser(ctx context.Context, ds model.DataStore, initialPassword string) error { + users := ds.User() + c, err := users.CountAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"user_name": consts.DevInitialUserName}}) if err != nil { return fmt.Errorf("could not access User table: %w", err) } @@ -60,7 +60,7 @@ func createInitialAdminUser(ds model.DataStore, initialPassword string) error { NewPassword: initialPassword, IsAdmin: true, } - if err := users.Put(&initialUser); err != nil { + if err := users.Put(ctx, &initialUser); err != nil { return fmt.Errorf("could not create initial admin user: %w", err) } } diff --git a/server/initial_setup_test.go b/server/initial_setup_test.go index 0ce8a39fa..0c85d9d0a 100644 --- a/server/initial_setup_test.go +++ b/server/initial_setup_test.go @@ -15,7 +15,7 @@ type failingPutUserRepo struct { err error } -func (r *failingPutUserRepo) Put(*model.User) error { return r.err } +func (r *failingPutUserRepo) Put(context.Context, *model.User) error { return r.err } func dsWithFailingPut(err error) model.DataStore { return &tests.MockDataStore{MockedUser: &failingPutUserRepo{UserRepository: tests.CreateMockUserRepo(), err: err}} @@ -23,37 +23,39 @@ func dsWithFailingPut(err error) model.DataStore { var _ = Describe("initial_setup", func() { var ds model.DataStore + var ctx context.Context BeforeEach(func() { ds = &tests.MockDataStore{} + ctx = GinkgoT().Context() }) Describe("createInitialAdminUser", func() { It("creates a new admin user with specified password if User table is empty", func() { - Expect(createInitialAdminUser(ds, "pass123")).To(BeNil()) - ur := ds.User(context.TODO()) - admin, err := ur.FindByUsername("admin") + Expect(createInitialAdminUser(ctx, ds, "pass123")).To(BeNil()) + ur := ds.User() + admin, err := ur.FindByUsername(ctx, "admin") Expect(err).To(BeNil()) Expect(admin.Password).To(Equal("pass123")) }) It("does not create a new admin user if User table is not empty", func() { - Expect(createInitialAdminUser(ds, "first")).To(BeNil()) - ur := ds.User(context.TODO()) - Expect(ur.CountAll()).To(Equal(int64(1))) - Expect(createInitialAdminUser(ds, "second")).To(BeNil()) - Expect(ur.CountAll()).To(Equal(int64(1))) + Expect(createInitialAdminUser(ctx, ds, "first")).To(BeNil()) + ur := ds.User() + Expect(ur.CountAll(ctx)).To(Equal(int64(1))) + Expect(createInitialAdminUser(ctx, ds, "second")).To(BeNil()) + Expect(ur.CountAll(ctx)).To(Equal(int64(1))) }) It("returns the error when the user cannot be stored", func() { boom := errors.New("db is down") - Expect(createInitialAdminUser(dsWithFailingPut(boom), "pass123")).To(MatchError(boom)) + Expect(createInitialAdminUser(ctx, dsWithFailingPut(boom), "pass123")).To(MatchError(boom)) }) It("returns the error when the user table cannot be read", func() { boom := errors.New("db is down") ds = &tests.MockDataStore{MockedUser: &tests.MockedUserRepo{Error: boom}} - Expect(createInitialAdminUser(ds, "pass123")).To(MatchError(boom)) + Expect(createInitialAdminUser(ctx, ds, "pass123")).To(MatchError(boom)) }) }) }) diff --git a/server/jellyfin/annotations.go b/server/jellyfin/annotations.go index ec84f65e0..a589798be 100644 --- a/server/jellyfin/annotations.go +++ b/server/jellyfin/annotations.go @@ -28,16 +28,16 @@ func (api *Router) resolveAnnotated(w http.ResponseWriter, r *http.Request, id s switch e := entity.(type) { case *model.Album: if u.HasLibraryAccess(e.LibraryID) { - return api.ds.Album(ctx), "album" + return api.ds.Album(), "album" } case *model.Artist: - return api.ds.Artist(ctx), "artist" + return api.ds.Artist(), "artist" case *model.MediaFile: if u.HasLibraryAccess(e.LibraryID) { - return api.ds.MediaFile(ctx), "song" + return api.ds.MediaFile(), "song" } case *model.Playlist: - return api.ds.Playlist(ctx), "playlist" + return api.ds.Playlist(), "playlist" } // Unknown ids, inaccessible-library items and non-annotatable entities (radios) all read as absent. http.Error(w, "Not Found", http.StatusNotFound) @@ -74,7 +74,7 @@ func (api *Router) setFavorite(w http.ResponseWriter, r *http.Request, starred b if repo == nil { return } - if err := repo.SetStar(starred, id); err != nil { + if err := repo.SetStar(r.Context(), starred, id); err != nil { api.internalError(w, r, err) return } @@ -97,7 +97,7 @@ func (api *Router) setItemRating(w http.ResponseWriter, r *http.Request, rating if repo == nil { return } - if err := repo.SetRating(rating, id); err != nil { + if err := repo.SetRating(r.Context(), rating, id); err != nil { api.internalError(w, r, err) return } diff --git a/server/jellyfin/annotations_test.go b/server/jellyfin/annotations_test.go index dd1487011..a436503ba 100644 --- a/server/jellyfin/annotations_test.go +++ b/server/jellyfin/annotations_test.go @@ -32,7 +32,7 @@ var _ = Describe("Annotations", func() { Describe("markFavorite / unmarkFavorite", func() { It("stars a song and returns IsFavorite=true", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/s1", nil).WithContext(ctxUser()) @@ -46,7 +46,7 @@ var _ = Describe("Annotations", func() { }) It("stars an album and returns IsFavorite=true", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxUser()) @@ -60,7 +60,7 @@ var _ = Describe("Annotations", func() { }) It("stars an artist without checking library access (artists span multiple libraries)", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() // alice only has access to library 1, but artists aren't gated per-library. @@ -75,7 +75,7 @@ var _ = Describe("Annotations", func() { }) It("stars a visible playlist", func() { - playlistRepo := ds.Playlist(context.Background()).(*tests.MockPlaylistRepo) + playlistRepo := ds.Playlist().(*tests.MockPlaylistRepo) playlistRepo.SetData(model.Playlists{{ID: testID("p1"), Name: "Mix", OwnerID: testID("u1")}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/"+dto.EncodeID(testID("p1")), nil).WithContext(ctxUser()) @@ -86,7 +86,7 @@ var _ = Describe("Annotations", func() { }) It("unstars a song and returns IsFavorite=false", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1, Annotations: model.Annotations{Starred: true}}}) w := httptest.NewRecorder() r := httptest.NewRequest("DELETE", "/Users/u1/FavoriteItems/s1", nil).WithContext(ctxUser()) @@ -100,7 +100,7 @@ var _ = Describe("Annotations", func() { }) It("returns 404 and does not star an album in a library the user can't access", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxUser()) // only has access to library 1 @@ -111,7 +111,7 @@ var _ = Describe("Annotations", func() { }) It("returns 404 and does not star a song in a library the user can't access", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/s1", nil).WithContext(ctxUser()) // only has access to library 1 @@ -130,7 +130,7 @@ var _ = Describe("Annotations", func() { }) It("returns 500 (not 404) when a repository lookup fails for a reason other than not-found", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetError(true) + ds.Album().(*tests.MockAlbumRepo).SetError(true) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/x1", nil).WithContext(ctxUser()) r = withChiURLParam(r, "itemId", dto.EncodeID(testID("x1"))) @@ -139,7 +139,7 @@ var _ = Describe("Annotations", func() { }) It("emits a refreshResource event when starring a song", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/s1", nil).WithContext(ctxUser()) @@ -150,7 +150,7 @@ var _ = Describe("Annotations", func() { }) It("emits a refreshResource event when starring an album", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxUser()) @@ -161,7 +161,7 @@ var _ = Describe("Annotations", func() { }) It("does not emit an event when the item is not accessible", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/FavoriteItems/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxUser()) @@ -174,7 +174,7 @@ var _ = Describe("Annotations", func() { Describe("setRating / removeRating", func() { It("maps a Jellyfin 0-10 rating to Navidrome's 0-5 scale", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=8", nil).WithContext(ctxUser()) @@ -189,7 +189,7 @@ var _ = Describe("Annotations", func() { }) It("rates an album", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/"+dto.EncodeID(testID("a1"))+"/Rating?Rating=10", nil).WithContext(ctxUser()) @@ -200,7 +200,7 @@ var _ = Describe("Annotations", func() { }) It("rates a visible playlist", func() { - playlistRepo := ds.Playlist(context.Background()).(*tests.MockPlaylistRepo) + playlistRepo := ds.Playlist().(*tests.MockPlaylistRepo) playlistRepo.SetData(model.Playlists{{ID: testID("p1"), Name: "Mix", OwnerID: testID("u1")}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/"+dto.EncodeID(testID("p1"))+"/Rating?Rating=8", nil).WithContext(ctxUser()) @@ -211,7 +211,7 @@ var _ = Describe("Annotations", func() { }) It("removes a rating", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1, Annotations: model.Annotations{Rating: 4}}}) w := httptest.NewRecorder() r := httptest.NewRequest("DELETE", "/Users/u1/Items/s1/Rating", nil).WithContext(ctxUser()) @@ -225,7 +225,7 @@ var _ = Describe("Annotations", func() { }) It("returns 404 and does not rate an album in a library the user can't access", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/"+dto.EncodeID(testID("a1"))+"/Rating?Rating=10", nil).WithContext(ctxUser()) // only has access to library 1 @@ -236,7 +236,7 @@ var _ = Describe("Annotations", func() { }) It("rounds an odd rating to the nearest star instead of truncating", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=9", nil).WithContext(ctxUser()) @@ -247,7 +247,7 @@ var _ = Describe("Annotations", func() { }) It("stores the minimum star for Rating=1 instead of clearing the rating", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1, Annotations: model.Annotations{Rating: 4}}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=1", nil).WithContext(ctxUser()) @@ -258,7 +258,7 @@ var _ = Describe("Annotations", func() { }) It("accepts a fractional rating (UserItemDataDto.Rating is a double)", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=7.5", nil).WithContext(ctxUser()) @@ -269,7 +269,7 @@ var _ = Describe("Annotations", func() { }) It("clamps a Rating above 10 to Navidrome's max (5)", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=100", nil).WithContext(ctxUser()) @@ -280,7 +280,7 @@ var _ = Describe("Annotations", func() { }) It("clamps a negative Rating to Navidrome's min (0)", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=-5", nil).WithContext(ctxUser()) @@ -291,7 +291,7 @@ var _ = Describe("Annotations", func() { }) It("emits a refreshResource event when rating a song", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/u1/Items/s1/Rating?Rating=8", nil).WithContext(ctxUser()) diff --git a/server/jellyfin/api_test.go b/server/jellyfin/api_test.go index a64dcfe8f..0c6ac9b98 100644 --- a/server/jellyfin/api_test.go +++ b/server/jellyfin/api_test.go @@ -1,6 +1,7 @@ package jellyfin import ( + "context" "net/http" "net/http/httptest" "strings" @@ -18,6 +19,12 @@ import ( ) var _ = Describe("Router", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = GinkgoT().Context() + }) + It("serves the public handshake through the mounted handler", func() { ds := &tests.MockDataStore{} api := New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) @@ -49,8 +56,8 @@ var _ = Describe("Router", func() { It("registers a player on a general authenticated request, not just playback reports", func() { ds := &tests.MockDataStore{} auth.Init(ds) - ur := ds.User(GinkgoT().Context()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) + ur := ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(ctx, &model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) token, err := auth.CreateToken(&model.User{ID: testID("u1"), UserName: "alice"}) Expect(err).ToNot(HaveOccurred()) @@ -95,7 +102,7 @@ var _ = Describe("Router", func() { ds := &tests.MockDataStore{} auth.Init(ds) usr := model.User{ID: testID("alice"), UserName: "alice"} - Expect(ds.User(GinkgoT().Context()).Put(&usr)).To(Succeed()) + Expect(ds.User().Put(ctx, &usr)).To(Succeed()) token, err := auth.CreateAPIToken(&usr, auth.AudienceJellyfin) Expect(err).ToNot(HaveOccurred()) api := New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, quickconnect.New()) diff --git a/server/jellyfin/auth.go b/server/jellyfin/auth.go index 32b4a1432..264d5b6ca 100644 --- a/server/jellyfin/auth.go +++ b/server/jellyfin/auth.go @@ -24,7 +24,7 @@ func (api *Router) authenticateByName(w http.ResponseWriter, r *http.Request) { } // Navidrome stores recoverable passwords; this mirrors Subsonic's validateCredentials plaintext path. - usr, err := api.ds.User(ctx).FindByUsernameWithPassword(body.Username) + usr, err := api.ds.User().FindByUsernameWithPassword(ctx, body.Username) if body.Pw == "" || err != nil || usr == nil || usr.Password != body.Pw { log.Warn(ctx, "Jellyfin API: invalid login", "username", body.Username, "remoteAddr", r.RemoteAddr) http.Error(w, "Unauthorized", http.StatusUnauthorized) @@ -37,7 +37,7 @@ func (api *Router) signIn(w http.ResponseWriter, r *http.Request, usr *model.Use ctx := r.Context() // Best-effort, like the web UI's validateLogin: without it, Jellyfin-only users show a // never/stale "Last Login" in the admin UI. - if err := api.ds.User(ctx).UpdateLastLoginAt(usr.ID); err != nil { + if err := api.ds.User().UpdateLastLoginAt(ctx, usr.ID); err != nil { log.Error(ctx, "Jellyfin API: could not update last login date", "username", usr.UserName, err) } diff --git a/server/jellyfin/auth_test.go b/server/jellyfin/auth_test.go index 0821614d6..dafad9244 100644 --- a/server/jellyfin/auth_test.go +++ b/server/jellyfin/auth_test.go @@ -17,13 +17,15 @@ import ( ) var _ = Describe("AuthenticateByName", func() { + var ctx context.Context var api *Router var ds *tests.MockDataStore BeforeEach(func() { + ctx = GinkgoT().Context() ds = &tests.MockDataStore{} auth.Init(ds) - ur := ds.User(context.Background()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) + ur := ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(ctx, &model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) api = &Router{ds: ds} }) @@ -105,15 +107,15 @@ var _ = Describe("AuthenticateByName", func() { api.authenticateByName(w, r) Expect(w.Code).To(Equal(http.StatusOK)) - ur := ds.User(context.Background()).(*tests.MockedUserRepo) - usr, err := ur.FindByUsername("alice") + ur := ds.User().(*tests.MockedUserRepo) + usr, err := ur.FindByUsername(ctx, "alice") Expect(err).ToNot(HaveOccurred()) Expect(usr.LastLoginAt).ToNot(BeNil()) }) It("reflects an administrator in the User.Policy", func() { - ur := ds.User(context.Background()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("admin1"), UserName: "root", NewPassword: "secret", IsAdmin: true})).To(Succeed()) + ur := ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(ctx, &model.User{ID: testID("admin1"), UserName: "root", NewPassword: "secret", IsAdmin: true})).To(Succeed()) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/AuthenticateByName", @@ -136,8 +138,8 @@ var _ = Describe("AuthenticateByName", func() { }) It("rejects an empty password even for a user with an empty stored password with 401", func() { - ur := ds.User(context.Background()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("e"), UserName: "empty", NewPassword: ""})).To(Succeed()) + ur := ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(ctx, &model.User{ID: testID("e"), UserName: "empty", NewPassword: ""})).To(Succeed()) w := httptest.NewRecorder() r := httptest.NewRequest("POST", "/Users/AuthenticateByName", diff --git a/server/jellyfin/browsing.go b/server/jellyfin/browsing.go index 25c44e748..baebbf0f3 100644 --- a/server/jellyfin/browsing.go +++ b/server/jellyfin/browsing.go @@ -77,7 +77,7 @@ func (api *Router) getStudios(w http.ResponseWriter, r *http.Request) { return } opts := model.QueryOptions{Sort: "tag_value", Filters: libraryScopeFilter(scope)} - labels, err := api.ds.Tag(ctx).GetAll(model.TagRecordLabel, opts) + labels, err := api.ds.Tag().GetAll(ctx, model.TagRecordLabel, opts) if err != nil { api.internalError(w, r, err) return @@ -97,12 +97,12 @@ func (api *Router) getQueryFiltersLegacy(w http.ResponseWriter, r *http.Request) return } genreOpts := model.QueryOptions{Sort: "name", Filters: libraryScopeFilter(scope)} - genres, err := api.ds.Genre(ctx).GetAll(genreOpts) + genres, err := api.ds.Genre().GetAll(ctx, genreOpts) if err != nil { api.internalError(w, r, err) return } - years, err := api.ds.Album(ctx).GetYears(scope...) + years, err := api.ds.Album().GetYears(ctx, scope...) if err != nil { api.internalError(w, r, err) return diff --git a/server/jellyfin/browsing_test.go b/server/jellyfin/browsing_test.go index 657c7a42c..48a62c4e5 100644 --- a/server/jellyfin/browsing_test.go +++ b/server/jellyfin/browsing_test.go @@ -33,7 +33,7 @@ var _ = Describe("Browsing", func() { Describe("getArtists", func() { It("lists artists via /Artists", func() { - ds.Artist(context.Background()).(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1"), Name: "A"}}) + ds.Artist().(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1"), Name: "A"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Artists", nil).WithContext(ctxUser(model.Libraries{{ID: 1}})) invoke(api.getArtists, w, r) @@ -45,7 +45,7 @@ var _ = Describe("Browsing", func() { }) It("handles /Artists/AlbumArtists the same way", func() { - ds.Artist(context.Background()).(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1"), Name: "A"}}) + ds.Artist().(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1"), Name: "A"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Artists/AlbumArtists", nil).WithContext(ctxUser(model.Libraries{{ID: 1}})) invoke(api.getArtists, w, r) @@ -56,7 +56,7 @@ var _ = Describe("Browsing", func() { }) It("scopes results to the user's accessible libraries", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} @@ -70,7 +70,7 @@ var _ = Describe("Browsing", func() { }) It("scopes to a single library when ParentId is an accessible library id", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} @@ -85,7 +85,7 @@ var _ = Describe("Browsing", func() { }) It("does not let ParentId= narrow the scope", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}} // no access to library 99 @@ -100,7 +100,7 @@ var _ = Describe("Browsing", func() { }) It("forwards SearchTerm to the repo's Search method", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Artists?SearchTerm=art", nil).WithContext(ctxUser(model.Libraries{{ID: 1}})) @@ -112,7 +112,7 @@ var _ = Describe("Browsing", func() { }) It("bounds a search the client left unbounded, and clamps an oversized one", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() @@ -129,7 +129,7 @@ var _ = Describe("Browsing", func() { }) It("forwards StartIndex/Limit as Offset/Max", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Artists?StartIndex=5&Limit=10", nil).WithContext(ctxUser(model.Libraries{{ID: 1}})) @@ -140,7 +140,7 @@ var _ = Describe("Browsing", func() { }) It("does not restrict results for an admin user", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Artists", nil).WithContext(ctxAdmin()) @@ -158,7 +158,7 @@ var _ = Describe("Browsing", func() { DescribeTable("restricts to favorites", func(url string, handler func(*Router) http.HandlerFunc) { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", url, nil).WithContext(ctxUser(model.Libraries{{ID: 1}})) @@ -180,7 +180,7 @@ var _ = Describe("Browsing", func() { ) It("404s a malformed ParentId instead of listing every library's artists", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Artists?ParentId=not-a-valid-id", nil).WithContext(ctxUser(model.Libraries{{ID: 1}})) @@ -210,7 +210,7 @@ var _ = Describe("Browsing", func() { Describe("getStudios", func() { It("scopes results to the user's accessible libraries", func() { - tagRepo := ds.Tag(context.Background()).(*tests.MockTagRepo) + tagRepo := ds.Tag().(*tests.MockTagRepo) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Studios", nil).WithContext(ctxUser(model.Libraries{{ID: 1}, {ID: 2}})) invoke(api.getStudios, w, r) @@ -224,7 +224,7 @@ var _ = Describe("Browsing", func() { // An empty scope (admin, or a non-admin with no explicit library grants) must be treated // as unrestricted, matching accessibleLibraryIDs' documented contract, not as "match nothing". It("does not restrict results for an admin user", func() { - tagRepo := ds.Tag(context.Background()).(*tests.MockTagRepo) + tagRepo := ds.Tag().(*tests.MockTagRepo) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Studios", nil).WithContext(ctxAdmin()) invoke(api.getStudios, w, r) @@ -242,7 +242,7 @@ var _ = Describe("Browsing", func() { Describe("getQueryFiltersLegacy", func() { It("scopes genres to the user's accessible libraries", func() { - genreRepo := ds.Genre(context.Background()).(*tests.MockedGenreRepo) + genreRepo := ds.Genre().(*tests.MockedGenreRepo) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/Filters", nil).WithContext(ctxUser(model.Libraries{{ID: 1}, {ID: 2}})) invoke(api.getQueryFiltersLegacy, w, r) @@ -254,7 +254,7 @@ var _ = Describe("Browsing", func() { }) It("does not restrict genres for an admin user", func() { - genreRepo := ds.Genre(context.Background()).(*tests.MockedGenreRepo) + genreRepo := ds.Genre().(*tests.MockedGenreRepo) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/Filters", nil).WithContext(ctxAdmin()) invoke(api.getQueryFiltersLegacy, w, r) diff --git a/server/jellyfin/e2e/auth_test.go b/server/jellyfin/e2e/auth_test.go index 156e79ce9..35daaf598 100644 --- a/server/jellyfin/e2e/auth_test.go +++ b/server/jellyfin/e2e/auth_test.go @@ -87,10 +87,10 @@ var _ = Describe("Authentication", func() { Expect(pw.Code).To(Equal(http.StatusOK)) // A real password change through the repository, which is what revokes in production. - admin, err := ds.User(ctx).Get(testID("admin-1")) + admin, err := ds.User().Get(ctx, testID("admin-1")) Expect(err).ToNot(HaveOccurred()) admin.NewPassword = "rotated" - Expect(ds.User(ctx).Put(admin)).To(Succeed()) + Expect(ds.User().Put(ctx, admin)).To(Succeed()) r = httptest.NewRequest("GET", "/Users/Me", nil) r.Header.Set("X-Emby-Token", res.AccessToken) diff --git a/server/jellyfin/e2e/e2e_suite_test.go b/server/jellyfin/e2e/e2e_suite_test.go index 5aa38cba1..4f3cd82b5 100644 --- a/server/jellyfin/e2e/e2e_suite_test.go +++ b/server/jellyfin/e2e/e2e_suite_test.go @@ -244,7 +244,7 @@ func enc(id string) string { return dto.EncodeID(id) } // guessing repository filter column names. func albumID(name string) string { - albums, err := ds.Album(ctx).GetAll() + albums, err := ds.Album().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) for _, a := range albums { if a.Name == name { @@ -256,7 +256,7 @@ func albumID(name string) string { } func songID(title string) string { - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) for _, mf := range mfs { if mf.Title == title { @@ -268,7 +268,7 @@ func songID(title string) string { } func artistID(name string) string { - artists, err := ds.Artist(ctx).GetAll() + artists, err := ds.Artist().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) for _, a := range artists { if a.Name == name { @@ -280,7 +280,7 @@ func artistID(name string) string { } func genreID(name string) string { - genres, err := ds.Genre(ctx).GetAll() + genres, err := ds.Genre().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) for _, g := range genres { if g.Name == name { @@ -399,7 +399,7 @@ func (f *fakeSonicProvider) FindSonicPath(context.Context, *model.MediaFile, *mo // songAgent looks a seeded track up by title (titles are unique in the seed) and builds an // agents.Song carrying its title+artist, so the matcher resolves it back to that MediaFile. func songAgent(title string) agents.Song { - mfs, err := ds.MediaFile(ctx).GetAll() + mfs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) for _, mf := range mfs { if mf.Title == title { diff --git a/server/jellyfin/e2e/multiuser_test.go b/server/jellyfin/e2e/multiuser_test.go index d3a521332..905d4e16f 100644 --- a/server/jellyfin/e2e/multiuser_test.go +++ b/server/jellyfin/e2e/multiuser_test.go @@ -19,8 +19,8 @@ var _ = Describe("Multi-user access control", func() { It("hides all content from a user with no library access", func() { noAccess := model.User{ID: testID("noaccess-1"), UserName: "noaccess", Name: "No Access", NewPassword: "password"} - Expect(ds.User(ctx).Put(&noAccess)).To(Succeed()) - loaded, err := ds.User(ctx).FindByUsername("noaccess") + Expect(ds.User().Put(ctx, &noAccess)).To(Succeed()) + loaded, err := ds.User().FindByUsername(ctx, "noaccess") Expect(err).ToNot(HaveOccurred()) q := queryResult(getAs(*loaded, "/Items?IncludeItemTypes=MusicAlbum&Recursive=true")) diff --git a/server/jellyfin/e2e/playlists_test.go b/server/jellyfin/e2e/playlists_test.go index 3dd53227f..49c8e5a8c 100644 --- a/server/jellyfin/e2e/playlists_test.go +++ b/server/jellyfin/e2e/playlists_test.go @@ -66,13 +66,13 @@ var _ = Describe("Playlists", func() { // dto.DecodeIDs is all-or-nothing: a malformed entry must 404 the whole request, not get // dropped while the well-formed entries are still used to create a playlist. It("404s when one of the Ids is malformed, without creating a playlist", func() { - before, err := ds.Playlist(ctx).CountAll() + before, err := ds.Playlist().CountAll(ctx) Expect(err).ToNot(HaveOccurred()) body := `{"Name":"ShouldNotExist","Ids":["` + enc(songID("So What")) + `","not-a-valid-id"]}` Expect(post("/Playlists", body).Code).To(Equal(http.StatusNotFound)) - after, err := ds.Playlist(ctx).CountAll() + after, err := ds.Playlist().CountAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(after).To(Equal(before)) }) @@ -318,14 +318,14 @@ var _ = Describe("Playlists", func() { Expect(upload(adminUser, "/Items/"+enc(plID)+"/Images/Primary", "image/jpeg", jpeg).Code). To(Equal(http.StatusNoContent)) - pls, err := ds.Playlist(ctx).Get(plID) + pls, err := ds.Playlist().Get(ctx, plID) Expect(err).ToNot(HaveOccurred()) Expect(pls.UploadedImage).ToNot(BeEmpty()) _, statErr := os.Stat(pls.UploadedImagePath()) Expect(statErr).ToNot(HaveOccurred(), "cover file should exist on disk") Expect(del("/Items/" + enc(plID) + "/Images/Primary").Code).To(Equal(http.StatusNoContent)) - pls, _ = ds.Playlist(ctx).Get(plID) + pls, _ = ds.Playlist().Get(ctx, plID) Expect(pls.UploadedImage).To(BeEmpty()) }) @@ -338,7 +338,7 @@ var _ = Describe("Playlists", func() { // from their tag-keyed cache until the next scan. It("clears the resolved image tag after a cover upload", func() { plID := createPlaylist("Cover Tag", nil) - Expect(ds.Artwork(ctx).PutItemArtwork(&model.ItemArtwork{ + Expect(ds.Artwork().PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: model.KindPlaylistArtwork.Prefix(), ItemID: plID, Hash: "1111111111111111", })).To(Succeed()) @@ -378,7 +378,7 @@ var _ = Describe("Playlists", func() { It("renames a playlist", func() { plID := createPlaylist("Old Name", nil) Expect(post("/Playlists/"+enc(plID), `{"Name":"New Name"}`).Code).To(Equal(http.StatusNoContent)) - pls, _ := ds.Playlist(ctx).Get(plID) + pls, _ := ds.Playlist().Get(ctx, plID) Expect(pls.Name).To(Equal("New Name")) }) @@ -410,7 +410,7 @@ var _ = Describe("Playlists", func() { q := playlistItems(plID) Expect(q.TotalRecordCount).To(Equal(1)) Expect(q.Items[0].Name).To(Equal("So What")) - pls, _ := ds.Playlist(ctx).Get(plID) + pls, _ := ds.Playlist().Get(ctx, plID) Expect(pls.Name).To(Equal("Combo Renamed")) Expect(pls.Public).To(BeTrue()) }) @@ -424,13 +424,13 @@ var _ = Describe("Playlists", func() { // An id that decodes to "" would tell Create to make a new playlist instead of updating one — // itemIDParam must 404 before that decode ever runs, not silently create one. It("404s for a malformed playlist id, without creating a playlist", func() { - before, err := ds.Playlist(ctx).CountAll() + before, err := ds.Playlist().CountAll(ctx) Expect(err).ToNot(HaveOccurred()) w := post("/Playlists/00000000000000000000000000000000", `{"Ids":["`+enc(songID("So What"))+`"]}`) Expect(w.Code).To(Equal(http.StatusNotFound)) - after, err := ds.Playlist(ctx).CountAll() + after, err := ds.Playlist().CountAll(ctx) Expect(err).ToNot(HaveOccurred()) Expect(after).To(Equal(before)) }) diff --git a/server/jellyfin/e2e/sessions_test.go b/server/jellyfin/e2e/sessions_test.go index c5c42cb2b..22c87b108 100644 --- a/server/jellyfin/e2e/sessions_test.go +++ b/server/jellyfin/e2e/sessions_test.go @@ -27,12 +27,12 @@ var _ = Describe("Sessions", func() { It("counts a play stopped past the threshold", func() { id := songID("So What") - mf, err := ds.MediaFile(ctx).Get(id) + mf, err := ds.MediaFile().Get(ctx, id) Expect(err).ToNot(HaveOccurred()) // Report a stop at the end of the track — comfortably past 50% / the 4-minute cap. Expect(post("/Sessions/Playing/Stopped", reportBody(id, ticks(int64(mf.Duration*1000)))).Code).To(Equal(http.StatusNoContent)) - mf, err = ds.MediaFile(ctx).Get(id) + mf, err = ds.MediaFile().Get(ctx, id) Expect(err).ToNot(HaveOccurred()) Expect(mf.PlayCount).To(BeNumerically(">=", 1)) }) @@ -48,7 +48,7 @@ var _ = Describe("Sessions", func() { id := songID("Help!") Expect(post("/Sessions/Playing/Stopped", reportBody(id, ticks(1000))).Code).To(Equal(http.StatusNoContent)) - mf, err := ds.MediaFile(ctx).Get(id) + mf, err := ds.MediaFile().Get(ctx, id) Expect(err).ToNot(HaveOccurred()) Expect(mf.PlayCount).To(Equal(int64(0))) }) diff --git a/server/jellyfin/e2e/similar_test.go b/server/jellyfin/e2e/similar_test.go index 06054f8bf..a5d2505ba 100644 --- a/server/jellyfin/e2e/similar_test.go +++ b/server/jellyfin/e2e/similar_test.go @@ -68,9 +68,9 @@ var _ = Describe("Similar", func() { // Seed an album in a second library the regular user has no access to, and point a // provider similar-song at it. otherLib := model.Library{ID: 2, Name: "Other Library", Path: "fake:///other"} - Expect(ds.Library(ctx).Put(&otherLib)).To(Succeed()) + Expect(ds.Library().Put(ctx, &otherLib)).To(Succeed()) otherAlbum := model.Album{ID: testID("other-album"), Name: "Other Album", LibraryID: 2} - Expect(ds.Album(ctx).Put(&otherAlbum)).To(Succeed()) + Expect(ds.Album().Put(ctx, &otherAlbum)).To(Succeed()) providerFake.similarSongs = model.MediaFiles{ {ID: testID("x1"), AlbumID: albumID("IV")}, // library 1 -> visible diff --git a/server/jellyfin/images.go b/server/jellyfin/images.go index 0c9e0380c..6af533dbb 100644 --- a/server/jellyfin/images.go +++ b/server/jellyfin/images.go @@ -82,16 +82,16 @@ func hashFromTag(r *http.Request) string { // resolveArtworkID maps a Jellyfin item id to a Navidrome ArtworkID, probing // album -> artist -> media file -> playlist. func (api *Router) resolveArtworkID(ctx context.Context, itemId string) string { - if al, err := api.ds.Album(ctx).Get(itemId); err == nil { + if al, err := api.ds.Album().Get(ctx, itemId); err == nil { return al.CoverArtID().String() } - if ar, err := api.ds.Artist(ctx).Get(itemId); err == nil { + if ar, err := api.ds.Artist().Get(ctx, itemId); err == nil { return ar.CoverArtID().String() } - if mf, err := api.ds.MediaFile(ctx).Get(itemId); err == nil { + if mf, err := api.ds.MediaFile().Get(ctx, itemId); err == nil { return mf.CoverArtID().String() } - if pl, err := api.ds.Playlist(ctx).Get(itemId); err == nil { + if pl, err := api.ds.Playlist().Get(ctx, itemId); err == nil { return pl.CoverArtID().String() } return (model.ArtworkID{}).String() diff --git a/server/jellyfin/images_test.go b/server/jellyfin/images_test.go index b7b262895..b7cf9b7e7 100644 --- a/server/jellyfin/images_test.go +++ b/server/jellyfin/images_test.go @@ -68,7 +68,7 @@ var _ = Describe("Images", func() { DescribeTable("derives the requested size from the Jellyfin size params", func(query string, wantSize int) { ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) fa := &fakeArtwork{} api := &Router{ds: ds, artwork: fa} @@ -97,7 +97,7 @@ var _ = Describe("Images", func() { It("streams album artwork", func() { ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) fa := &fakeArtwork{} api := &Router{ds: ds, artwork: fa} @@ -138,7 +138,7 @@ var _ = Describe("Images", func() { It("sniffs the Content-Type instead of hardcoding it", func() { ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) png := append([]byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}, make([]byte, 512)...) fa := &fakeArtwork{data: png} @@ -153,7 +153,7 @@ var _ = Describe("Images", func() { It("resolves a playlist's cover regardless of visibility, even for an anonymous caller", func() { ds := &tests.MockDataStore{} - ds.Playlist(context.Background()).(*tests.MockPlaylistRepo).SetData(model.Playlists{{ID: testID("pl1"), Name: "Mix", OwnerID: testID("someone")}}) + ds.Playlist().(*tests.MockPlaylistRepo).SetData(model.Playlists{{ID: testID("pl1"), Name: "Mix", OwnerID: testID("someone")}}) fa := &fakeArtwork{} api := &Router{ds: ds, artwork: fa} @@ -169,7 +169,7 @@ var _ = Describe("Images", func() { // silently falls back to the placeholder. It("resolves artwork under an elevated admin context", func() { ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) fa := &fakeArtwork{} api := &Router{ds: ds, artwork: fa} @@ -185,7 +185,7 @@ var _ = Describe("Images", func() { It("serves immutable when the tag param asserts the current hash", func() { const hash = "0123456789abcdef" ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) fa := &fakeArtwork{hash: hash} api := &Router{ds: ds, artwork: fa} @@ -203,7 +203,7 @@ var _ = Describe("Images", func() { It("revalidates via no-cache when no tag is provided", func() { const hash = "0123456789abcdef" ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) fa := &fakeArtwork{hash: hash} api := &Router{ds: ds, artwork: fa} diff --git a/server/jellyfin/items.go b/server/jellyfin/items.go index 23e35dd59..28de4c31e 100644 --- a/server/jellyfin/items.go +++ b/server/jellyfin/items.go @@ -382,7 +382,7 @@ func (api *Router) parseItemsQuery(ctx context.Context, r *http.Request) (itemsQ if q.parentId == dto.PlaylistsFolderID { // Browsing into the synthetic playlists folder lists the user's playlists. q.types = []string{"Playlist"} - } else if _, err := api.ds.Album(ctx).Get(q.parentId); err == nil { + } else if _, err := api.ds.Album().Get(ctx, q.parentId); err == nil { q.types = []string{"Audio"} } } @@ -412,7 +412,7 @@ func (api *Router) queryItems(ctx context.Context, r *http.Request) (itemsResult return materialized(result([]dto.BaseItemDto{playlistsFolder()}, 1, 0)), nil } if repo, ok := api.playlistTracksRepo(ctx, q); ok { - return api.playlistTrackPage(repo, q.fields, q.offset, q.limit) + return api.playlistTrackPage(ctx, repo, q.fields, q.offset, q.limit) } if q.search != "" { q.limit = clampLimit(q.limit, defaultSearchLimit, maxSearchLimit) @@ -702,7 +702,7 @@ func searchPage[S ~[]E, E any](opts model.QueryOptions, search func(model.QueryO func (api *Router) listAlbums(ctx context.Context, opts model.QueryOptions, q itemsQuery) (itemsResult, error) { toItem := func(al model.Album) dto.BaseItemDto { return dto.AlbumToBaseItem(al, q.fields) } - repo := api.ds.Album(ctx) + repo := api.ds.Album() filters := squirrel.And{} // For albums, ParentId (browse an artist) and AlbumArtistIds/ArtistIds both mean "this artist's // albums"; contributingArtistIds means "albums they only appear on" (Featured On). @@ -733,23 +733,23 @@ func (api *Router) listAlbums(ctx context.Context, opts model.QueryOptions, q it if q.search != "" { albums, total, err := searchPage(opts, func(o model.QueryOptions) (model.Albums, error) { - return repo.Search(q.search, o) + return repo.Search(ctx, q.search, o) }) if err != nil { return itemsResult{}, err } return materialized(result(slice.Map(albums, toItem), total, opts.Offset)), nil } - total, _ := repo.CountAll(model.QueryOptions{Filters: opts.Filters}) + total, _ := repo.CountAll(ctx, model.QueryOptions{Filters: opts.Filters}) open := streamCursor(func() (func(func(model.Album, error) bool), error) { - return repo.GetCursor(opts) + return repo.GetCursor(ctx, opts) }, toItem) return streamed(open, int(total), opts.Offset), nil } func (api *Router) listSongs(ctx context.Context, opts model.QueryOptions, q itemsQuery) (itemsResult, error) { toItem := func(mf model.MediaFile) dto.BaseItemDto { return dto.SongToBaseItem(mf, q.fields) } - repo := api.ds.MediaFile(ctx) + repo := api.ds.MediaFile() filters := squirrel.And{} // For songs, ArtistIds/AlbumArtistIds selects an artist's tracks; ParentId selects an album's. switch { @@ -782,7 +782,7 @@ func (api *Router) listSongs(ctx context.Context, opts model.QueryOptions, q ite if q.search != "" { mfs, total, err := searchPage(opts, func(o model.QueryOptions) (model.MediaFiles, error) { - return repo.Search(q.search, o) + return repo.Search(ctx, q.search, o) }) if err != nil { return itemsResult{}, err @@ -795,9 +795,9 @@ func (api *Router) listSongs(ctx context.Context, opts model.QueryOptions, q ite opts.Sort = filter.SongsByAlbum(q.entityParent).Sort } // A full-library request (Finamp's sync, with MediaSources) is tens of thousands of fat rows. - total, _ := repo.CountAll(model.QueryOptions{Filters: opts.Filters}) + total, _ := repo.CountAll(ctx, model.QueryOptions{Filters: opts.Filters}) open := streamCursor(func() (func(func(model.MediaFile, error) bool), error) { - return repo.GetCursorWithArtwork(opts) + return repo.GetCursorWithArtwork(ctx, opts) }, toItem) return streamed(open, int(total), opts.Offset), nil } @@ -806,7 +806,7 @@ func (api *Router) listSongs(ctx context.Context, opts model.QueryOptions, q ite // RoleArtist for performing artists (/Artists). Without the role filter both lists would be identical. // genreIds isn't applied to search — a name lookup, like role (see below). func (api *Router) listArtists(ctx context.Context, opts model.QueryOptions, q itemsQuery, role model.Role) (itemsResult, error) { - repo := api.ds.Artist(ctx) + repo := api.ds.Artist() toItem := func(ar model.Artist) dto.BaseItemDto { return dto.ArtistToBaseItem(ar, q.fields) } // Artist Search does its own library scoping: it consumes a sole Eq{"library_id": ...} filter as a @@ -818,7 +818,7 @@ func (api *Router) listArtists(ctx context.Context, opts model.QueryOptions, q i opts.Filters = squirrel.Eq{"library_id": q.scopeIDs} } artists, total, err := searchPage(opts, func(o model.QueryOptions) (model.Artists, error) { - return repo.Search(q.search, o) + return repo.Search(ctx, q.search, o) }) if err != nil { return itemsResult{}, err @@ -834,9 +834,9 @@ func (api *Router) listArtists(ctx context.Context, opts model.QueryOptions, q i opts.Filters = filters opts = filter.ArtistsByRole(opts, role) opts = filter.ApplyArtistLibraryFilter(opts, q.scopeIDs) - total, _ := repo.CountAll(model.QueryOptions{Filters: opts.Filters}) + total, _ := repo.CountAll(ctx, model.QueryOptions{Filters: opts.Filters}) open := streamCursor(func() (func(func(model.Artist, error) bool), error) { - return repo.GetCursor(opts) + return repo.GetCursor(ctx, opts) }, toItem) return streamed(open, int(total), opts.Offset), nil } @@ -845,7 +845,7 @@ func (api *Router) listArtists(ctx context.Context, opts model.QueryOptions, q i // the one listXxx that stays materialized: GenreRepository has no CountAll, so the total is the // length of the full list and paging is in-memory — nothing for a cursor to page over. func (api *Router) listGenres(ctx context.Context, opts model.QueryOptions) (itemsResult, error) { - genres, err := api.ds.Genre(ctx).GetAll(model.QueryOptions{Sort: opts.Sort, Order: opts.Order}) + genres, err := api.ds.Genre().GetAll(ctx, model.QueryOptions{Sort: opts.Sort, Order: opts.Order}) if err != nil { return itemsResult{}, err } @@ -859,13 +859,13 @@ func (api *Router) listPlaylists(ctx context.Context, opts model.QueryOptions, q if preds := q.filters.predicates(); len(preds) > 0 { opts.Filters = squirrel.And(preds) } - repo := api.ds.Playlist(ctx) - total, err := repo.CountAll(model.QueryOptions{Filters: opts.Filters}) + repo := api.ds.Playlist() + total, err := repo.CountAll(ctx, model.QueryOptions{Filters: opts.Filters}) if err != nil { return itemsResult{}, err } open := streamCursor(func() (func(func(model.Playlist, error) bool), error) { - return repo.GetCursor(opts) + return repo.GetCursor(ctx, opts) }, func(p model.Playlist) dto.BaseItemDto { return dto.PlaylistToBaseItem(p, q.fields) }) return streamed(open, int(total), opts.Offset), nil } @@ -882,21 +882,21 @@ func (api *Router) resolveItemByID(ctx context.Context, id string, fields dto.Fi // Finamp resolves a /UserViews entry (Id=library id) by fetching it as a plain item; without this // the home screen and library tabs 404. if libID, err := strconv.Atoi(id); err == nil && u.HasLibraryAccess(libID) { - if lib, err := api.ds.Library(ctx).Get(libID); err == nil { + if lib, err := api.ds.Library().Get(ctx, libID); err == nil { return dto.LibraryToBaseItem(*lib), true } } - if al, err := api.ds.Album(ctx).Get(id); err == nil { + if al, err := api.ds.Album().Get(ctx, id); err == nil { if !u.HasLibraryAccess(al.LibraryID) { return dto.BaseItemDto{}, false } return dto.AlbumToBaseItem(*al, fields), true } - if ar, err := api.ds.Artist(ctx).Get(id); err == nil { + if ar, err := api.ds.Artist().Get(ctx, id); err == nil { // Artist.Get already scopes to the user's libraries via library_artist. return dto.ArtistToBaseItem(*ar, fields), true } - if mf, err := api.ds.MediaFile(ctx).Get(id); err == nil { + if mf, err := api.ds.MediaFile().Get(ctx, id); err == nil { if !u.HasLibraryAccess(mf.LibraryID) { return dto.BaseItemDto{}, false } @@ -906,7 +906,7 @@ func (api *Router) resolveItemByID(ctx context.Context, id string, fields dto.Fi if pl, err := api.playlists.Get(ctx, id); err == nil { return dto.PlaylistToBaseItem(*pl, fields), true } - if g, err := api.ds.Genre(ctx).Get(id); err == nil { + if g, err := api.ds.Genre().Get(ctx, id); err == nil { return dto.GenreToBaseItem(*g), true } return dto.BaseItemDto{}, false @@ -917,7 +917,7 @@ func (api *Router) songsByIDs(ctx context.Context, ids []string) map[string]mode songs := make(map[string]model.MediaFile, len(ids)) // Chunked to stay under SQLITE_MAX_VARIABLE_NUMBER, like playqueue's loadTracks. for chunk := range slice.CollectChunks(slices.Values(ids), 500) { - mfs, err := api.ds.MediaFile(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"media_file.id": chunk}}) + mfs, err := api.ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"media_file.id": chunk}}) if err != nil { log.Error(ctx, "Jellyfin API: error fetching songs by id", err) continue @@ -997,9 +997,9 @@ func (api *Router) getLatest(w http.ResponseWriter, r *http.Request) { opts.Filters = squirrel.And{opts.Filters, filter.AlbumsByArtistID(parentID).Filters} } opts = filter.ApplyLibraryFilter(opts, scopeIDs) - repo := api.ds.Album(ctx) + repo := api.ds.Album() open := streamCursor(func() (func(func(model.Album, error) bool), error) { - return repo.GetCursor(opts) + return repo.GetCursor(ctx, opts) }, func(al model.Album) dto.BaseItemDto { return dto.AlbumToBaseItem(al, fields) }) api.writeItemsArray(w, r, streamed(open, 0, 0)) } diff --git a/server/jellyfin/items_test.go b/server/jellyfin/items_test.go index a8dfb7bb2..18789a34b 100644 --- a/server/jellyfin/items_test.go +++ b/server/jellyfin/items_test.go @@ -49,7 +49,7 @@ var _ = Describe("Items", func() { Describe("getItems", func() { It("lists albums when IncludeItemTypes=MusicAlbum", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&Recursive=true", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -62,8 +62,8 @@ var _ = Describe("Items", func() { }) It("lists an album's songs when ParentId is an album and type is Audio", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", AlbumID: testID("a1")}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", AlbumID: testID("a1")}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?ParentId="+dto.EncodeID(testID("a1"))+"&IncludeItemTypes=Audio", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -75,8 +75,8 @@ var _ = Describe("Items", func() { }) It("ignores IncludeItemTypes names that aren't Jellyfin item kinds", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", AlbumID: testID("a1")}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", AlbumID: testID("a1")}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?ParentId="+dto.EncodeID(testID("a1"))+"&IncludeItemTypes=music", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -126,8 +126,8 @@ var _ = Describe("Items", func() { It("falls through to the type dispatch when ParentId is not a playlist", func() { fp.getErr = model.ErrNotFound - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), AlbumID: testID("a1")}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), AlbumID: testID("a1")}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?ParentId="+dto.EncodeID(testID("a1"))+"&IncludeItemTypes=Audio", nil). WithContext(ctxUser()) @@ -140,7 +140,7 @@ var _ = Describe("Items", func() { }) It("returns 500 when the song cursor fails to open, instead of a truncated 200", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetError(true) + ds.MediaFile().(*tests.MockMediaFileRepo).SetError(true) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio&Recursive=true", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -151,8 +151,8 @@ var _ = Describe("Items", func() { // looking for tracks outside any album; answering with every track streams the whole library. Describe("Recursive=false", func() { BeforeEach(func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), AlbumID: testID("a1")}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), AlbumID: testID("a1")}}) }) It("returns no songs for a library parent, as tracks are never its direct children", func() { @@ -223,21 +223,21 @@ var _ = Describe("Items", func() { }) It("lists an artist's albums when ParentId is an artist and type is MusicAlbum", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", AlbumArtistID: testID("ar1")}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", AlbumArtistID: testID("ar1")}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?ParentId="+dto.EncodeID(testID("ar1"))+"&IncludeItemTypes=MusicAlbum", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) var res dto.QueryResult Expect(json.Unmarshal(w.Body.Bytes(), &res)).To(Succeed()) Expect(res.Items).To(HaveLen(1)) - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) sql, _, err := albumRepo.Options.Filters.ToSql() Expect(err).NotTo(HaveOccurred()) Expect(sql).To(ContainSubstring("album_artists")) }) It("lists artists when IncludeItemTypes=MusicArtist", func() { - ds.Artist(context.Background()).(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) + ds.Artist().(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicArtist", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -258,7 +258,7 @@ var _ = Describe("Items", func() { }) It("lists playlists when IncludeItemTypes=Playlist", func() { - ds.Playlist(context.Background()).(*tests.MockPlaylistRepo).SetData(model.Playlists{{ID: testID("p1"), Name: "My Mix", SongCount: 5}}) + ds.Playlist().(*tests.MockPlaylistRepo).SetData(model.Playlists{{ID: testID("p1"), Name: "My Mix", SongCount: 5}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Playlist", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -272,8 +272,8 @@ var _ = Describe("Items", func() { }) It("merges results from every requested type in IncludeItemTypes", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -287,11 +287,11 @@ var _ = Describe("Items", func() { }) It("merges favorite songs, albums, and playlists", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - playlistRepo := ds.Playlist(context.Background()).(*tests.MockPlaylistRepo) + playlistRepo := ds.Playlist().(*tests.MockPlaylistRepo) playlistRepo.SetData(model.Playlists{{ID: testID("p1"), Name: "My Mix", Annotations: model.Annotations{Starred: true}}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum,Playlist&Filters=IsFavorite", nil).WithContext(ctxUser()) @@ -312,8 +312,8 @@ var _ = Describe("Items", func() { }) It("applies StartIndex/Limit to the merged multi-type result set", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}, {ID: testID("s2"), Title: "Song2"}}) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}, {ID: testID("s2"), Title: "Song2"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum&StartIndex=1&Limit=2", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -326,9 +326,9 @@ var _ = Describe("Items", func() { }) It("caps each per-type query at StartIndex+Limit instead of fetching everything", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}, {ID: testID("s2"), Title: "Song2"}}) - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum&StartIndex=1&Limit=2", nil).WithContext(ctxUser()) @@ -341,7 +341,7 @@ var _ = Describe("Items", func() { DescribeTable("translates the Filters list and its standalone equivalents", func(query string, wantSQL, notWantSQL []string) { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&"+query, nil).WithContext(ctxUser()) @@ -381,8 +381,8 @@ var _ = Describe("Items", func() { // annotation predicate there is "no such column: starred" -> 500. DescribeTable("does not push annotation filters into a search", func(itemType, filters string) { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes="+itemType+"&SearchTerm=one&Filters="+filters, nil).WithContext(ctxUser()) @@ -390,9 +390,9 @@ var _ = Describe("Items", func() { Expect(w.Code).To(Equal(http.StatusOK)) var opts model.QueryOptions if itemType == "MusicAlbum" { - opts = ds.Album(context.Background()).(*tests.MockAlbumRepo).Options + opts = ds.Album().(*tests.MockAlbumRepo).Options } else { - opts = ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).Options + opts = ds.MediaFile().(*tests.MockMediaFileRepo).Options } if opts.Filters == nil { return @@ -410,7 +410,7 @@ var _ = Describe("Items", func() { ) It("forwards SearchTerm to the repo's Search method", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&SearchTerm=one", nil).WithContext(ctxUser()) @@ -422,7 +422,7 @@ var _ = Describe("Items", func() { }) It("caps a search the client left unbounded", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&SearchTerm=one", nil).WithContext(ctxUser()) @@ -432,7 +432,7 @@ var _ = Describe("Items", func() { }) It("honors an explicit search Limit up to the ceiling", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&SearchTerm=one&Limit=500", nil). @@ -443,7 +443,7 @@ var _ = Describe("Items", func() { }) It("clamps a search Limit that would materialize the library", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&SearchTerm=one&Limit=999999", nil). @@ -454,7 +454,7 @@ var _ = Describe("Items", func() { }) It("treats an all-whitespace SearchTerm as no search, streaming the unfiltered list", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&SearchTerm=%20%20", nil). @@ -472,8 +472,8 @@ var _ = Describe("Items", func() { for i := range songs { songs[i] = model.MediaFile{ID: testID(fmt.Sprintf("s%05d", i)), Title: "Song"} } - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(songs) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(songs) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum&SearchTerm=song&Limit=10", nil). WithContext(ctxUser()) @@ -486,9 +486,9 @@ var _ = Describe("Items", func() { }) It("bounds the multi-type search window however large StartIndex is", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum&SearchTerm=song&StartIndex=500000&Limit=1", nil). WithContext(ctxUser()) @@ -505,8 +505,8 @@ var _ = Describe("Items", func() { for i := range songs { songs[i] = model.MediaFile{ID: testID(fmt.Sprintf("s%05d", i)), Title: "Song"} } - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(songs) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(songs) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", fmt.Sprintf("/Items?IncludeItemTypes=Audio,MusicAlbum&SearchTerm=song&StartIndex=%d&Limit=1", maxSearchLimit), @@ -526,8 +526,8 @@ var _ = Describe("Items", func() { } // The mock repo returns rows sorted by ID; reorder to match so index-based assertions hold. slices.SortFunc(songs, func(a, b model.MediaFile) int { return strings.Compare(a.ID, b.ID) }) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(songs) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(songs) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", fmt.Sprintf("/Items?IncludeItemTypes=Audio,MusicAlbum&SearchTerm=song&StartIndex=%d&Limit=10", maxSearchLimit-1), @@ -546,8 +546,8 @@ var _ = Describe("Items", func() { for i := range songs { songs[i] = model.MediaFile{ID: testID(fmt.Sprintf("s%05d", i)), Title: "Song"} } - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(songs) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(songs) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum&SearchTerm=song", nil). WithContext(ctxUser()) @@ -565,8 +565,8 @@ var _ = Describe("Items", func() { } // The mock repo returns rows sorted by ID; reorder to match so index-based assertions hold. slices.SortFunc(songs, func(a, b model.MediaFile) int { return strings.Compare(a.ID, b.ID) }) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(songs) - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(songs) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", fmt.Sprintf("/Items?IncludeItemTypes=Audio,MusicAlbum&SearchTerm=song&StartIndex=%d", defaultSearchLimit+50), @@ -581,7 +581,7 @@ var _ = Describe("Items", func() { }) It("reports a search total beyond the fetched page instead of the page length", func() { - ds.Artist(context.Background()).(*tests.MockArtistRepo).SetData(model.Artists{ + ds.Artist().(*tests.MockArtistRepo).SetData(model.Artists{ {ID: testID("r1"), Name: "Alpha"}, {ID: testID("r2"), Name: "Beta"}, {ID: testID("r3"), Name: "Gamma"}, }) w := httptest.NewRecorder() @@ -595,7 +595,7 @@ var _ = Describe("Items", func() { }) It("forwards StartIndex/Limit as Offset/Max", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&StartIndex=5&Limit=10", nil).WithContext(ctxUser()) @@ -609,7 +609,7 @@ var _ = Describe("Items", func() { // Finamp's download/sync fetches a track's BaseItemDto via /Items?ids=; without // this, queryItems ignored Ids and returned the default type-dispatched list instead. It("returns exactly the requested item when Ids has a single id", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?Ids="+dto.EncodeID(testID("s1")), nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -623,8 +623,8 @@ var _ = Describe("Items", func() { }) It("returns items of different types for a lowercase ids param with multiple ids", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?ids="+dto.EncodeID(testID("a1"))+","+dto.EncodeID(testID("s1")), nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -640,7 +640,7 @@ var _ = Describe("Items", func() { }) It("resolves song ids with one batched IN query, not a Get per id", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 1}, {ID: testID("s2"), Title: "Song2", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?ids="+dto.EncodeID(testID("s1"))+","+dto.EncodeID(testID("s2")), nil).WithContext(ctxUser()) @@ -656,8 +656,8 @@ var _ = Describe("Items", func() { }) It("omits an id in a library the user can't access, without erroring the whole batch", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 2}}) // alice only has access to library 1 + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 2}}) // alice only has access to library 1 w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?Ids="+dto.EncodeID(testID("a1"))+","+dto.EncodeID(testID("s1")), nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -673,9 +673,9 @@ var _ = Describe("Items", func() { Describe("sorting", func() { DescribeTable("translates SortBy into the repo's sort keys", func(itemType, sortBy, want string) { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes="+itemType+"&SortBy="+sortBy, nil).WithContext(ctxUser()) @@ -710,7 +710,7 @@ var _ = Describe("Items", func() { // we honor the first value for all keys, matching Jellyfin's fallback for extra keys. DescribeTable("reads the first SortOrder value for the whole sort", func(sortOrder, want string) { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", @@ -728,7 +728,7 @@ var _ = Describe("Items", func() { Describe("library scoping", func() { It("scopes a MusicAlbum listing (no ParentId) to the user's accessible libraries", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} @@ -742,7 +742,7 @@ var _ = Describe("Items", func() { }) It("scopes a Audio listing (no ParentId) to the user's accessible libraries", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} @@ -756,7 +756,7 @@ var _ = Describe("Items", func() { }) It("scopes a MusicArtist listing to the user's accessible libraries", func() { - artistRepo := ds.Artist(context.Background()).(*tests.MockArtistRepo) + artistRepo := ds.Artist().(*tests.MockArtistRepo) artistRepo.SetData(model.Artists{{ID: testID("ar1"), Name: "Artist"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} @@ -770,7 +770,7 @@ var _ = Describe("Items", func() { }) It("treats a numeric ParentId matching an accessible library as a library scope, not an artist id", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} @@ -785,7 +785,7 @@ var _ = Describe("Items", func() { }) It("does not let ParentId= scope results to that library", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}} // no access to library 99 @@ -803,7 +803,7 @@ var _ = Describe("Items", func() { }) It("does not restrict a default MusicAlbum listing for an admin user", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}, {ID: testID("a2"), Name: "Two", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum", nil).WithContext(ctxAdmin()) @@ -824,7 +824,7 @@ var _ = Describe("Items", func() { // still reach the entity filter, not the unfiltered default. Describe("stale and malformed id filtering", func() { It("404s a malformed ParentId instead of listing every song", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio&ParentId=not-a-valid-id", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -832,7 +832,7 @@ var _ = Describe("Items", func() { }) It("404s a malformed AlbumArtistIds instead of listing every album", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&AlbumArtistIds=not-a-valid-id", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -840,7 +840,7 @@ var _ = Describe("Items", func() { }) It("404s a malformed ArtistIds instead of listing every song", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio&ArtistIds=not-a-valid-id", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -848,7 +848,7 @@ var _ = Describe("Items", func() { }) It("still applies the artist filter (rather than dropping it) for a well-formed but unknown AlbumArtistIds", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=MusicAlbum&AlbumArtistIds="+dto.EncodeID(testID("no-such-artist")), nil).WithContext(ctxUser()) @@ -860,7 +860,7 @@ var _ = Describe("Items", func() { }) It("still applies the album filter (rather than dropping it) for a well-formed but unknown ParentId", func() { - mfRepo := ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo) + mfRepo := ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio&ParentId="+dto.EncodeID(testID("no-such-album")), nil).WithContext(ctxUser()) @@ -875,8 +875,8 @@ var _ = Describe("Items", func() { Describe("mixed IncludeItemTypes merge", func() { BeforeEach(func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "S1"}, {ID: testID("s2"), Title: "S2"}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One"}, {ID: testID("a2"), Name: "Two"}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "S1"}, {ID: testID("s2"), Title: "S2"}}) }) It("returns a mix of both types, not all of one", func() { @@ -927,7 +927,7 @@ var _ = Describe("Items", func() { }) It("propagates a per-type query error", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetError(true) + ds.MediaFile().(*tests.MockMediaFileRepo).SetError(true) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items?IncludeItemTypes=Audio,MusicAlbum&Recursive=true&Limit=4", nil).WithContext(ctxUser()) invoke(api.getItems, w, r) @@ -938,7 +938,7 @@ var _ = Describe("Items", func() { Describe("getItem", func() { It("returns an album by id", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxUser()) r = withChiURLParam(r, "itemId", dto.EncodeID(testID("a1"))) @@ -959,7 +959,7 @@ var _ = Describe("Items", func() { }) It("returns 404 for an album in a library the user can't access", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxUser()) // only has access to library 1 r = withChiURLParam(r, "itemId", dto.EncodeID(testID("a1"))) @@ -968,7 +968,7 @@ var _ = Describe("Items", func() { }) It("returns 404 for a song in a library the user can't access", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 2}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1"), Title: "Song", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/"+dto.EncodeID(testID("s1")), nil).WithContext(ctxUser()) // only has access to library 1 r = withChiURLParam(r, "itemId", dto.EncodeID(testID("s1"))) @@ -977,7 +977,7 @@ var _ = Describe("Items", func() { }) It("returns an album to an admin even when it's outside their (empty) Libraries", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 2}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/"+dto.EncodeID(testID("a1")), nil).WithContext(ctxAdmin()) // admin, Libraries: nil r = withChiURLParam(r, "itemId", dto.EncodeID(testID("a1"))) @@ -993,7 +993,7 @@ var _ = Describe("Items", func() { It("resolves a library-view id (from /UserViews) as a CollectionFolder item", func() { w := httptest.NewRecorder() libs := model.Libraries{{ID: 1, Name: "Music Library"}} - ds.Library(context.Background()).(*tests.MockLibraryRepo).SetData(libs) + ds.Library().(*tests.MockLibraryRepo).SetData(libs) r := httptest.NewRequest("GET", "/Items/"+dto.EncodeLibraryID(1), nil).WithContext(ctxUserWithLibraries(libs)) r = withChiURLParam(r, "itemId", dto.EncodeLibraryID(1)) invoke(api.getItem, w, r) @@ -1043,7 +1043,7 @@ var _ = Describe("Items", func() { // Finamp's genre "See all" fetches the genre by id; a 404 white-screens it (see resolveItemByID). It("resolves a genre id as a MusicGenre item", func() { - Expect(ds.Genre(context.Background()).(*tests.MockedGenreRepo).Put(&model.Genre{ID: testID("g1"), Name: "Rock"})).To(Succeed()) + Expect(ds.Genre().(*tests.MockedGenreRepo).Put(&model.Genre{ID: testID("g1"), Name: "Rock"})).To(Succeed()) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/"+dto.EncodeID(testID("g1")), nil).WithContext(ctxUser()) r = withChiURLParam(r, "itemId", dto.EncodeID(testID("g1"))) @@ -1057,7 +1057,7 @@ var _ = Describe("Items", func() { }) It("resolves a library-view id for an admin even though their Libraries slice is empty", func() { - ds.Library(context.Background()).(*tests.MockLibraryRepo).SetData(model.Libraries{{ID: 1, Name: "Music Library"}}) + ds.Library().(*tests.MockLibraryRepo).SetData(model.Libraries{{ID: 1, Name: "Music Library"}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Items/"+dto.EncodeLibraryID(1), nil).WithContext(ctxAdmin()) r = withChiURLParam(r, "itemId", dto.EncodeLibraryID(1)) @@ -1073,7 +1073,7 @@ var _ = Describe("Items", func() { Describe("getLatest", func() { It("returns a bare array of the newest albums", func() { - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Users/u1/Items/Latest", nil).WithContext(ctxUser()) invoke(api.getLatest, w, r) @@ -1085,7 +1085,7 @@ var _ = Describe("Items", func() { }) It("scopes to the user's accessible libraries", func() { - albumRepo := ds.Album(context.Background()).(*tests.MockAlbumRepo) + albumRepo := ds.Album().(*tests.MockAlbumRepo) albumRepo.SetData(model.Albums{{ID: testID("a1"), Name: "One", LibraryID: 1}}) w := httptest.NewRecorder() libs := model.Libraries{{ID: 1}, {ID: 2}} diff --git a/server/jellyfin/lyrics_test.go b/server/jellyfin/lyrics_test.go index 2cf53a280..fb00563de 100644 --- a/server/jellyfin/lyrics_test.go +++ b/server/jellyfin/lyrics_test.go @@ -52,7 +52,7 @@ var _ = Describe("getLyrics", func() { BeforeEach(func() { ds = &tests.MockDataStore{} - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", LibraryID: 1}, {ID: testID("s2"), Title: "Silent Song", LibraryID: 1}, }) diff --git a/server/jellyfin/middlewares.go b/server/jellyfin/middlewares.go index 0ae4f6071..4788c84be 100644 --- a/server/jellyfin/middlewares.go +++ b/server/jellyfin/middlewares.go @@ -162,7 +162,7 @@ func (api *Router) userFromToken(r *http.Request) (model.User, bool) { if err != nil || claims.Subject == "" { return model.User{}, false } - usr, err := api.ds.User(r.Context()).FindByUsername(claims.Subject) + usr, err := api.ds.User().FindByUsername(r.Context(), claims.Subject) if err != nil { log.Warn(r.Context(), "Jellyfin API: token subject not found", "user", claims.Subject, err) return model.User{}, false diff --git a/server/jellyfin/middlewares_test.go b/server/jellyfin/middlewares_test.go index a3b88799b..2ab9f35eb 100644 --- a/server/jellyfin/middlewares_test.go +++ b/server/jellyfin/middlewares_test.go @@ -18,13 +18,15 @@ import ( ) var _ = Describe("authenticate middleware", func() { + var ctx context.Context var api *Router var ds *tests.MockDataStore BeforeEach(func() { + ctx = GinkgoT().Context() ds = &tests.MockDataStore{} auth.Init(ds) - ur := ds.User(context.Background()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) + ur := ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(ctx, &model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) api = &Router{ds: ds} }) @@ -100,9 +102,9 @@ var _ = Describe("authenticate middleware", func() { var usr *model.User BeforeEach(func() { - ur := ds.User(context.Background()).(*tests.MockedUserRepo) + ur := ds.User().(*tests.MockedUserRepo) usr = &model.User{ID: testID("u2"), UserName: "bob", NewPassword: "secret", TokenEpoch: 3} - Expect(ur.Put(usr)).To(Succeed()) + Expect(ur.Put(ctx, usr)).To(Succeed()) }) serve := func(token string) *httptest.ResponseRecorder { diff --git a/server/jellyfin/playlists.go b/server/jellyfin/playlists.go index 5fd2df8c9..5569f4f83 100644 --- a/server/jellyfin/playlists.go +++ b/server/jellyfin/playlists.go @@ -146,14 +146,14 @@ func (api *Router) clearPlaylist(ctx context.Context, id string) error { // playlistTrackPage streams one page of a playlist's tracks. Streams because a playlist can be the // whole library (a smart playlist matching everything) and clients may omit Limit. Excludes missing // tracks, and counts the same set, like GetWithTracks. -func (api *Router) playlistTrackPage(repo model.PlaylistTrackRepository, fields dto.Fields, offset, limit int) (itemsResult, error) { - total, err := repo.CountAll(model.QueryOptions{Filters: notMissing}) +func (api *Router) playlistTrackPage(ctx context.Context, repo model.PlaylistTrackRepository, fields dto.Fields, offset, limit int) (itemsResult, error) { + total, err := repo.CountAll(ctx, model.QueryOptions{Filters: notMissing}) if err != nil { return itemsResult{}, err } opts := model.QueryOptions{Sort: "id", Offset: offset, Max: limit, Filters: notMissing} open := streamCursor(func() (func(func(model.PlaylistTrack, error) bool), error) { - return repo.GetCursor(opts) + return repo.GetCursor(ctx, opts) }, func(t model.PlaylistTrack) dto.BaseItemDto { return trackToBaseItem(t, fields) }) return streamed(open, int(total), offset), nil } @@ -188,7 +188,7 @@ func (api *Router) getPlaylist(w http.ResponseWriter, r *http.Request) { return } // PlaylistInfo carries every track id, so this can't be paged — but it needs no track data. - trackIDs, err := repo.GetMediaFileIDs(model.QueryOptions{Sort: "id", Filters: notMissing}) + trackIDs, err := repo.GetMediaFileIDs(ctx, model.QueryOptions{Sort: "id", Filters: notMissing}) if err != nil { api.internalError(w, r, err) return @@ -216,7 +216,7 @@ func (api *Router) getPlaylistItems(w http.ResponseWriter, r *http.Request) { } p := req.Params(r) fields := dto.ParseFields(p.Strings("fields")...) - res, err := api.playlistTrackPage(repo, fields, p.IntOr("startindex", 0), p.IntOr("limit", 0)) + res, err := api.playlistTrackPage(ctx, repo, fields, p.IntOr("startindex", 0), p.IntOr("limit", 0)) if err != nil { api.internalError(w, r, err) return @@ -249,9 +249,9 @@ func (api *Router) expandContainerIDs(ctx context.Context, ids []string) []strin for _, id := range ids { if _, ok := songs[id]; ok { out = append(out, id) // already a song - } else if _, err := api.ds.Album(ctx).Get(id); err == nil { + } else if _, err := api.ds.Album().Get(ctx, id); err == nil { out = append(out, api.songIDs(ctx, filter.SongsByAlbum(id))...) - } else if _, err := api.ds.Artist(ctx).Get(id); err == nil { + } else if _, err := api.ds.Artist().Get(ctx, id); err == nil { out = append(out, api.songIDs(ctx, filter.SongsByArtistID(id))...) } else if pl, err := api.playlists.GetWithTracks(ctx, id); err == nil { out = append(out, slice.Map(pl.Tracks, func(t model.PlaylistTrack) string { return t.MediaFileID })...) @@ -263,7 +263,7 @@ func (api *Router) expandContainerIDs(ctx context.Context, ids []string) []strin } func (api *Router) songIDs(ctx context.Context, opts model.QueryOptions) []string { - mfs, err := api.ds.MediaFile(ctx).GetAll(opts) + mfs, err := api.ds.MediaFile().GetAll(ctx, opts) if err != nil { log.Error(ctx, "Jellyfin: error expanding container to tracks", err) return nil diff --git a/server/jellyfin/playlists_test.go b/server/jellyfin/playlists_test.go index edaa63dd1..d38ef3b44 100644 --- a/server/jellyfin/playlists_test.go +++ b/server/jellyfin/playlists_test.go @@ -299,27 +299,27 @@ var _ = Describe("Playlists", func() { } It("passes a bare song id through unchanged", func() { - ds.MediaFile(ctx).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1")}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1")}}) createWith(testID("s1")) Expect(fp.createdIds).To(Equal([]string{testID("s1")})) }) It("expands an album id into its songs, filtered by album", func() { - ds.Album(ctx).(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("al1")}}) - ds.MediaFile(ctx).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{{ID: testID("al1")}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), AlbumID: testID("al1")}, {ID: testID("s2"), AlbumID: testID("al1")}, }) createWith(testID("al1")) Expect(fp.createdIds).To(Equal([]string{testID("s1"), testID("s2")})) - Expect(ds.MediaFile(ctx).(*tests.MockMediaFileRepo).Options.Filters).To(Equal(filter.SongsByAlbum(testID("al1")).Filters)) + Expect(ds.MediaFile().(*tests.MockMediaFileRepo).Options.Filters).To(Equal(filter.SongsByAlbum(testID("al1")).Filters)) }) It("expands an artist id into its songs", func() { - ds.Artist(ctx).(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1")}}) - ds.MediaFile(ctx).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1")}, {ID: testID("s2")}}) + ds.Artist().(*tests.MockArtistRepo).SetData(model.Artists{{ID: testID("ar1")}}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{{ID: testID("s1")}, {ID: testID("s2")}}) createWith(testID("ar1")) Expect(fp.createdIds).To(Equal([]string{testID("s1"), testID("s2")})) - Expect(ds.MediaFile(ctx).(*tests.MockMediaFileRepo).Options.Filters).To(Equal(filter.SongsByArtistID(testID("ar1")).Filters)) + Expect(ds.MediaFile().(*tests.MockMediaFileRepo).Options.Filters).To(Equal(filter.SongsByArtistID(testID("ar1")).Filters)) }) It("expands a playlist id into its tracks' media file ids", func() { diff --git a/server/jellyfin/quickconnect.go b/server/jellyfin/quickconnect.go index c478b1894..f7fa1b979 100644 --- a/server/jellyfin/quickconnect.go +++ b/server/jellyfin/quickconnect.go @@ -77,7 +77,7 @@ func (api *Router) quickConnectAuthorize(w http.ResponseWriter, r *http.Request) http.Error(w, "Forbidden", http.StatusForbidden) return } - usr, err := api.ds.User(ctx).Get(userID) + usr, err := api.ds.User().Get(ctx, userID) if errors.Is(err, model.ErrNotFound) { http.Error(w, "Unknown user", http.StatusNotFound) return @@ -118,7 +118,7 @@ func (api *Router) authenticateWithQuickConnect(w http.ResponseWriter, r *http.R http.Error(w, "Unknown secret", http.StatusNotFound) return } - usr, err := api.ds.User(ctx).Get(userID) + usr, err := api.ds.User().Get(ctx, userID) if errors.Is(err, model.ErrNotFound) { log.Warn(ctx, "Jellyfin API: Quick Connect user not found", "userID", userID) http.Error(w, "Unauthorized", http.StatusUnauthorized) diff --git a/server/jellyfin/quickconnect_test.go b/server/jellyfin/quickconnect_test.go index 0478a356a..94573e6ff 100644 --- a/server/jellyfin/quickconnect_test.go +++ b/server/jellyfin/quickconnect_test.go @@ -1,7 +1,6 @@ package jellyfin import ( - "context" "encoding/json" "errors" "net/http" @@ -42,9 +41,9 @@ var _ = Describe("QuickConnect", func() { conf.Server.Jellyfin.QuickConnect = true ds = &tests.MockDataStore{} auth.Init(ds) - ur := ds.User(context.Background()).(*tests.MockedUserRepo) + ur := ds.User().(*tests.MockedUserRepo) for _, u := range []model.User{alice, bob, admin} { - Expect(ur.Put(&u)).To(Succeed()) + Expect(ur.Put(GinkgoT().Context(), &u)).To(Succeed()) } qc = quickconnect.New() api = &Router{ds: ds, quickConnect: qc} @@ -294,7 +293,7 @@ var _ = Describe("QuickConnect", func() { It("returns 500 when the user lookup fails", func() { req := initiate() _, _ = qc.Authorize(req.Code, alice.ID) - ds.User(context.Background()).(*tests.MockedUserRepo).Error = errors.New("db down") + ds.User().(*tests.MockedUserRepo).Error = errors.New("db down") Expect(redeemSecret(req.Secret).Code).To(Equal(http.StatusInternalServerError)) }) }) diff --git a/server/jellyfin/similar.go b/server/jellyfin/similar.go index 1bd887f43..50503160c 100644 --- a/server/jellyfin/similar.go +++ b/server/jellyfin/similar.go @@ -216,7 +216,7 @@ func (api *Router) similarAlbums(ctx context.Context, id string, limit int) dto. continue } seen[s.AlbumID] = true - if al, err := api.ds.Album(ctx).Get(s.AlbumID); err == nil && u.HasLibraryAccess(al.LibraryID) { + if al, err := api.ds.Album().Get(ctx, s.AlbumID); err == nil && u.HasLibraryAccess(al.LibraryID) { items = append(items, dto.AlbumToBaseItem(*al, nil)) if len(items) >= limit { break diff --git a/server/jellyfin/similar_test.go b/server/jellyfin/similar_test.go index cffcce0a8..42d04a632 100644 --- a/server/jellyfin/similar_test.go +++ b/server/jellyfin/similar_test.go @@ -122,7 +122,7 @@ var _ = Describe("getInstantMix", func() { DeferCleanup(func() { similarWait = old }) ds := &tests.MockDataStore{} - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Seed Song", LibraryID: 1}, }) release := make(chan struct{}) @@ -150,7 +150,7 @@ var _ = Describe("getInstantMix", func() { songs = append(songs, model.MediaFile{ID: testID(fmt.Sprintf("t%d", i)), Title: fmt.Sprintf("Track %d", i), LibraryID: 1}) } ds := &tests.MockDataStore{} - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(songs) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(songs) api := &Router{ds: ds, provider: &fakeSimilarProvider{songs: songs[1:]}} w := httptest.NewRecorder() @@ -191,7 +191,7 @@ var _ = Describe("getSimilarAlbums", func() { // With no external agent the provider falls back to the album's own tracks, which map // straight back to the requested album. ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{ + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{ {ID: testID("al-1"), Name: "Seed Album", LibraryID: 1}, }) api := &Router{ds: ds, provider: &fakeSimilarProvider{ @@ -212,7 +212,7 @@ var _ = Describe("getSimilarAlbums", func() { It("returns albums derived from the provider's similar songs", func() { ds := &tests.MockDataStore{} - ds.Album(context.Background()).(*tests.MockAlbumRepo).SetData(model.Albums{ + ds.Album().(*tests.MockAlbumRepo).SetData(model.Albums{ {ID: testID("al-2"), Name: "Other", LibraryID: 1}, }) api := &Router{ds: ds, provider: &fakeSimilarProvider{ diff --git a/server/jellyfin/socket_test.go b/server/jellyfin/socket_test.go index b4be0e244..79098aafd 100644 --- a/server/jellyfin/socket_test.go +++ b/server/jellyfin/socket_test.go @@ -1,7 +1,6 @@ package jellyfin import ( - "context" "net/http" "net/http/httptest" "strings" @@ -87,8 +86,8 @@ var _ = Describe("handleSocket", func() { BeforeEach(func() { ds = &tests.MockDataStore{} auth.Init(ds) - ur := ds.User(context.Background()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) + ur := ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(GinkgoT().Context(), &model.User{ID: testID("u1"), UserName: "alice", NewPassword: "secret"})).To(Succeed()) t, err := auth.CreateToken(&model.User{ID: testID("u1"), UserName: "alice"}) Expect(err).ToNot(HaveOccurred()) diff --git a/server/jellyfin/stream.go b/server/jellyfin/stream.go index b9809bd29..515a6fcb8 100644 --- a/server/jellyfin/stream.go +++ b/server/jellyfin/stream.go @@ -29,7 +29,7 @@ func (api *Router) mediaFileForRequest(w http.ResponseWriter, r *http.Request) ( if !ok { return nil, false } - mf, err := api.ds.MediaFile(ctx).Get(id) + mf, err := api.ds.MediaFile().Get(ctx, id) if err != nil { http.Error(w, "Not Found", http.StatusNotFound) return nil, false diff --git a/server/jellyfin/stream_test.go b/server/jellyfin/stream_test.go index 1a6d1badc..7a1112f37 100644 --- a/server/jellyfin/stream_test.go +++ b/server/jellyfin/stream_test.go @@ -42,7 +42,7 @@ var _ = Describe("Stream", func() { Describe("getPlaybackInfo", func() { It("returns a media source for an accessible track", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", Duration: 100, Size: 1000, LibraryID: 1}, }) w := httptest.NewRecorder() @@ -61,7 +61,7 @@ var _ = Describe("Stream", func() { }) It("returns 404 for a track in a library the user can't access", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 2}, }) w := httptest.NewRecorder() @@ -102,7 +102,7 @@ var _ = Describe("Stream", func() { } It("advertises a Lyric stream for plugin/sidecar-sourced lyrics not embedded in the file", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1}, }) api.lyrics = &fakeLyricsService{lyrics: map[string]model.LyricList{ @@ -113,7 +113,7 @@ var _ = Describe("Stream", func() { }) It("advertises no Lyric stream when the pipeline finds nothing", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1}, }) @@ -121,7 +121,7 @@ var _ = Describe("Stream", func() { }) It("advertises no Lyric stream when the lyrics endpoint would 404 (main lyric has no lines)", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1}, }) api.lyrics = &fakeLyricsService{lyrics: map[string]model.LyricList{ @@ -132,7 +132,7 @@ var _ = Describe("Stream", func() { }) It("doesn't duplicate the Lyric stream when lyrics are already embedded", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1, Lyrics: `[{"lang":"xxx","line":[]}]`}, }) api.lyrics = &fakeLyricsService{lyrics: map[string]model.LyricList{ @@ -145,7 +145,7 @@ var _ = Describe("Stream", func() { It("still returns 200 with a valid MediaSource and no Lyric stream when the lyrics pipeline errors", func() { // Own ID: an erroring loader isn't cached, but a shared ID could still pick up // another test's cached (non-error) result and mask this assertion. - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s-err"), Title: "Song", Suffix: "mp3", Duration: 100, Size: 1000, LibraryID: 1}, }) api.lyrics = &fakeLyricsService{err: errors.New("boom")} @@ -166,7 +166,7 @@ var _ = Describe("Stream", func() { Describe("streamAudio", func() { It("invokes the transcode decider and streamer for an accessible track", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1}, }) streamer.content = "audio-bytes" @@ -182,7 +182,7 @@ var _ = Describe("Stream", func() { }) It("returns 404 for a track in a library the user can't access, without invoking the streamer or decider", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 2}, }) w := httptest.NewRecorder() @@ -207,7 +207,7 @@ var _ = Describe("Stream", func() { }) It("converts the bps audioBitRate param to kbps", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "flac", LibraryID: 1}, }) w := httptest.NewRecorder() @@ -219,7 +219,7 @@ var _ = Describe("Stream", func() { }) It("uses the audioCodec param as target format when no container is given", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "flac", LibraryID: 1}, }) w := httptest.NewRecorder() @@ -231,7 +231,7 @@ var _ = Describe("Stream", func() { }) It("returns 500 and logs when the streamer fails", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1}, }) streamer.err = errors.New("boom") @@ -246,7 +246,7 @@ var _ = Describe("Stream", func() { Describe("HEAD requests", func() { head := func(query string) *httptest.ResponseRecorder { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "flac", LibraryID: 1}, }) streamer.content = "audio-bytes" @@ -276,7 +276,7 @@ var _ = Describe("Stream", func() { Describe("streamUniversal", func() { universal := func(query string) { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Suffix: "mp3", LibraryID: 1}, }) w := httptest.NewRecorder() @@ -315,7 +315,7 @@ var _ = Describe("Stream", func() { Describe("streamHls", func() { BeforeEach(func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "dsf", Duration: 100.5, LibraryID: 1}, }) }) @@ -370,21 +370,21 @@ var _ = Describe("Stream", func() { }) It("returns 404 for a track in a library the user can't access", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "dsf", LibraryID: 2}, }) Expect(hls("", ctxUser()).Code).To(Equal(http.StatusNotFound)) }) It("returns 404 when the id doesn't match any media file", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{}) + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{}) Expect(hls("", ctxUser()).Code).To(Equal(http.StatusNotFound)) }) }) Describe("streamFile", func() { It("invokes the decider with a raw/direct-play request and the streamer for an accessible track", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 1}, }) streamer.content = "audio-bytes" @@ -401,7 +401,7 @@ var _ = Describe("Stream", func() { }) It("returns 404 for a track in a library the user can't access, without invoking the streamer or decider", func() { - ds.MediaFile(context.Background()).(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ {ID: testID("s1"), Title: "Song", Suffix: "mp3", LibraryID: 2}, }) w := httptest.NewRecorder() diff --git a/server/jellyfin/system.go b/server/jellyfin/system.go index 49e41470c..2aae4c9c0 100644 --- a/server/jellyfin/system.go +++ b/server/jellyfin/system.go @@ -51,11 +51,11 @@ func resolveServerID(ctx context.Context, ds model.DataStore, cached *string) st *cached = newServerID() return *cached } - id, err := ds.Property(ctx).Get(consts.JellyfinServerIDKey) + id, err := ds.Property().Get(ctx, consts.JellyfinServerIDKey) switch { case errors.Is(err, model.ErrNotFound): id = newServerID() - if err := ds.Property(ctx).Put(consts.JellyfinServerIDKey, id); err != nil { + if err := ds.Property().Put(ctx, consts.JellyfinServerIDKey, id); err != nil { log.Error(ctx, "Jellyfin API: could not persist server id", err) return id } diff --git a/server/jellyfin/system_test.go b/server/jellyfin/system_test.go index 2350a7342..d339043a7 100644 --- a/server/jellyfin/system_test.go +++ b/server/jellyfin/system_test.go @@ -182,10 +182,10 @@ var _ = Describe("System", func() { }) It("does not overwrite or pin over a stored id when the property read fails transiently", func() { - Expect(ds.Property(ctx).Put(consts.JellyfinServerIDKey, "6ba7b8109dad11d180b400c04fd430c8")).To(Succeed()) + Expect(ds.Property().Put(ctx, consts.JellyfinServerIDKey, "6ba7b8109dad11d180b400c04fd430c8")).To(Succeed()) r := &Router{ds: ds} - props := ds.Property(ctx).(*tests.MockedPropertyRepo) + props := ds.Property().(*tests.MockedPropertyRepo) props.Error = errors.New("database is locked") degraded := r.serverID(ctx) Expect(degraded).ToNot(BeEmpty()) @@ -194,7 +194,7 @@ var _ = Describe("System", func() { // Once the DB recovers, the stored id is intact and served again. Expect(r.serverID(ctx)).To(Equal("6ba7b8109dad11d180b400c04fd430c8")) - stored, err := ds.Property(ctx).Get(consts.JellyfinServerIDKey) + stored, err := ds.Property().Get(ctx, consts.JellyfinServerIDKey) Expect(err).ToNot(HaveOccurred()) Expect(stored).To(Equal("6ba7b8109dad11d180b400c04fd430c8")) }) @@ -205,7 +205,7 @@ var _ = Describe("System", func() { }) It("strips dashes from an already-persisted id", func() { - Expect(ds.Property(ctx).Put( + Expect(ds.Property().Put(ctx, consts.JellyfinServerIDKey, "1b4e28ba-2fa1-11d2-883f-0016d3cca427")).To(Succeed()) r := &Router{ds: ds} Expect(r.serverID(ctx)).To(Equal("1b4e28ba2fa111d2883f0016d3cca427")) diff --git a/server/jellyfin/users.go b/server/jellyfin/users.go index dddcd99da..d0a54ee63 100644 --- a/server/jellyfin/users.go +++ b/server/jellyfin/users.go @@ -17,7 +17,7 @@ func (api *Router) getUserViews(w http.ResponseWriter, r *http.Request) { u, _ := request.UserFrom(ctx) // u.Libraries comes from a projection without counts or stats, and clients hide a library that // looks empty, so the rows are re-read in full here. - libs, err := api.ds.Library(ctx).GetAll() + libs, err := api.ds.Library().GetAll(ctx) if err != nil { api.internalError(w, r, err) return @@ -55,7 +55,7 @@ func (api *Router) getPublicUsers(w http.ResponseWriter, r *http.Request) { continue } seen[key] = true - usr, err := api.ds.User(ctx).FindByUsername(name) + usr, err := api.ds.User().FindByUsername(ctx, name) if err != nil { log.Warn(ctx, "Jellyfin API: configured public user not found", "username", name, err) continue diff --git a/server/jellyfin/users_test.go b/server/jellyfin/users_test.go index c6bb99f71..793a64a29 100644 --- a/server/jellyfin/users_test.go +++ b/server/jellyfin/users_test.go @@ -18,10 +18,11 @@ import ( ) var _ = Describe("Users", func() { + var ctx context.Context var api *Router // The repo holds the full rows; the user carries the id/name-only copy its projection returns. authedWithLibraries := func(r *http.Request, libs model.Libraries) *http.Request { - api.ds.Library(context.Background()).(*tests.MockLibraryRepo).SetData(libs) + api.ds.Library().(*tests.MockLibraryRepo).SetData(libs) stripped := make(model.Libraries, len(libs)) for i, lib := range libs { stripped[i] = model.Library{ID: lib.ID, Name: lib.Name} @@ -29,7 +30,10 @@ var _ = Describe("Users", func() { ctx := request.WithUser(context.Background(), model.User{ID: testID("u1"), UserName: "alice", Libraries: stripped}) return r.WithContext(ctx) } - BeforeEach(func() { api = &Router{ds: &tests.MockDataStore{}} }) + BeforeEach(func() { + ctx = GinkgoT().Context() + api = &Router{ds: &tests.MockDataStore{}} + }) Describe("getUserViews", func() { It("returns one view per accessible library", func() { @@ -117,9 +121,9 @@ var _ = Describe("Users", func() { BeforeEach(func() { DeferCleanup(configtest.SetupConfig()) - ur = api.ds.User(context.Background()).(*tests.MockedUserRepo) - Expect(ur.Put(&model.User{ID: testID("u1"), UserName: "alice"})).To(Succeed()) - Expect(ur.Put(&model.User{ID: testID("u2"), UserName: "bob"})).To(Succeed()) + ur = api.ds.User().(*tests.MockedUserRepo) + Expect(ur.Put(ctx, &model.User{ID: testID("u1"), UserName: "alice"})).To(Succeed()) + Expect(ur.Put(ctx, &model.User{ID: testID("u2"), UserName: "bob"})).To(Succeed()) }) It("returns an empty list when the config is unset", func() { diff --git a/server/middlewares.go b/server/middlewares.go index 674337e92..b65a2d6e1 100644 --- a/server/middlewares.go +++ b/server/middlewares.go @@ -380,7 +380,7 @@ func UpdateLastAccessMiddleware(ds model.DataStore) func(next http.Handler) http ctx, cancel := context.WithTimeout(ctx, time.Second) defer cancel() - err := ds.User(ctx).UpdateLastAccessAt(usr.ID) + err := ds.User().UpdateLastAccessAt(ctx, usr.ID) if err != nil { log.Warn(ctx, "Could not update user's lastAccessAt", "username", usr.UserName, "elapsed", time.Since(start), err) diff --git a/server/middlewares_test.go b/server/middlewares_test.go index a9b0bc99e..15cf70341 100644 --- a/server/middlewares_test.go +++ b/server/middlewares_test.go @@ -381,7 +381,7 @@ var _ = Describe("middlewares", func() { id = uuid.NewString() ds = &tests.MockDataStore{} lastAccessTime = time.Now() - Expect(ds.User(ctx).Put(&model.User{ID: id, UserName: "johndoe", LastAccessAt: &lastAccessTime})). + Expect(ds.User().Put(ctx, &model.User{ID: id, UserName: "johndoe", LastAccessAt: &lastAccessTime})). To(Succeed()) middleware = UpdateLastAccessMiddleware(ds) @@ -407,14 +407,14 @@ var _ = Describe("middlewares", func() { callMiddleware(req) - user, _ := ds.MockedUser.FindByUsername("johndoe") + user, _ := ds.MockedUser.FindByUsername(ctx, "johndoe") Expect(*user.LastAccessAt).To(BeTemporally(">", lastAccessTime, time.Second)) }) It("skip fast successive requests", func() { // First request callMiddleware(req) - user, _ := ds.MockedUser.FindByUsername("johndoe") + user, _ := ds.MockedUser.FindByUsername(ctx, "johndoe") lastAccessTime = *user.LastAccessAt // Store the last access time // Second request @@ -422,7 +422,7 @@ var _ = Describe("middlewares", func() { callMiddleware(req) // The second request should not have changed the last access time - user, _ = ds.MockedUser.FindByUsername("johndoe") + user, _ = ds.MockedUser.FindByUsername(ctx, "johndoe") Expect(user.LastAccessAt).To(Equal(&lastAccessTime)) }) }) @@ -431,7 +431,7 @@ var _ = Describe("middlewares", func() { req = req.WithContext(context.Background()) callMiddleware(req) - usr, _ := ds.MockedUser.FindByUsername("johndoe") + usr, _ := ds.MockedUser.FindByUsername(ctx, "johndoe") Expect(usr.LastAccessAt).To(Equal(&lastAccessTime)) }) }) diff --git a/server/nativeapi/artists.go b/server/nativeapi/artists.go index 193f88eda..91508c825 100644 --- a/server/nativeapi/artists.go +++ b/server/nativeapi/artists.go @@ -15,14 +15,12 @@ import ( ) func (api *Router) addArtistRoute(r chi.Router) { - constructor := func(ctx context.Context) rest.Repository { - return api.ds.Resource(ctx, model.Artist{}) - } + repo := api.ds.Artist() r.Route("/artist", func(r chi.Router) { - r.Get("/", rest.GetAll(constructor)) + r.Get("/", rest.GetAll(repo)) r.Route("/{id}", func(r chi.Router) { r.Use(server.URLParamsMiddleware) - r.Get("/", rest.Get(constructor)) + r.Get("/", rest.Get(repo)) r.Post("/image", api.uploadArtistImage()) r.Delete("/image", api.deleteArtistImage()) }) @@ -32,7 +30,7 @@ func (api *Router) addArtistRoute(r chi.Router) { func (api *Router) uploadArtistImage() http.HandlerFunc { return handleImageUpload(func(ctx context.Context, reader io.Reader, ext string) error { artistID := chi.URLParamFromCtx(ctx, "id") - ar, err := api.ds.Artist(ctx).Get(artistID) + ar, err := api.ds.Artist().Get(ctx, artistID) if err != nil { if errors.Is(err, model.ErrNotFound) { return model.ErrNotFound @@ -46,7 +44,7 @@ func (api *Router) uploadArtistImage() http.HandlerFunc { } ar.UploadedImage = filename ar.UpdatedAt = new(time.Now()) - if err := api.ds.Artist(ctx).Put(ar, "uploaded_image", "updated_at"); err != nil { + if err := api.ds.Artist().Put(ctx, ar, "uploaded_image", "updated_at"); err != nil { return err } api.imgUpload.EnqueueArtwork(ctx, consts.EntityArtist, ar.ID) @@ -57,7 +55,7 @@ func (api *Router) uploadArtistImage() http.HandlerFunc { func (api *Router) deleteArtistImage() http.HandlerFunc { return handleImageDelete(func(ctx context.Context) error { artistID := chi.URLParamFromCtx(ctx, "id") - ar, err := api.ds.Artist(ctx).Get(artistID) + ar, err := api.ds.Artist().Get(ctx, artistID) if err != nil { if errors.Is(err, model.ErrNotFound) { return model.ErrNotFound @@ -69,7 +67,7 @@ func (api *Router) deleteArtistImage() http.HandlerFunc { } ar.UploadedImage = "" ar.UpdatedAt = new(time.Now()) - if err := api.ds.Artist(ctx).Put(ar, "uploaded_image", "updated_at"); err != nil { + if err := api.ds.Artist().Put(ctx, ar, "uploaded_image", "updated_at"); err != nil { return err } api.imgUpload.EnqueueArtwork(ctx, consts.EntityArtist, ar.ID) diff --git a/server/nativeapi/config_test.go b/server/nativeapi/config_test.go index 7c4f00fdd..6ac41f07a 100644 --- a/server/nativeapi/config_test.go +++ b/server/nativeapi/config_test.go @@ -11,6 +11,7 @@ import ( "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/server" "github.com/navidrome/navidrome/tests" @@ -19,17 +20,19 @@ import ( ) var _ = Describe("Config API", func() { + var ctx context.Context var ds model.DataStore var router http.Handler var adminUser, regularUser model.User BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) conf.Server.EnableSharing = false conf.Server.DevUIShowConfig = true // Enable config endpoint for tests ds = &tests.MockDataStore{} auth.Init(ds) - nativeRouter := New(ds, nil, nil, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) + nativeRouter := New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) router = server.JWTVerifier(nativeRouter) // Create test users @@ -49,8 +52,8 @@ var _ = Describe("Config API", func() { } // Store in mock datastore - Expect(ds.User(context.TODO()).Put(&adminUser)).To(Succeed()) - Expect(ds.User(context.TODO()).Put(®ularUser)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, ®ularUser)).To(Succeed()) }) Describe("GET /api/config", func() { diff --git a/server/nativeapi/inspect.go b/server/nativeapi/inspect.go index 7c96312ed..f1e6c4539 100644 --- a/server/nativeapi/inspect.go +++ b/server/nativeapi/inspect.go @@ -13,7 +13,7 @@ import ( ) func doInspect(ctx context.Context, ds model.DataStore, id string) (*core.InspectOutput, error) { - file, err := ds.MediaFile(ctx).Get(id) + file, err := ds.MediaFile().Get(ctx, id) if err != nil { return nil, err } diff --git a/server/nativeapi/library_test.go b/server/nativeapi/library_test.go index cef2e06ad..cc05a30e5 100644 --- a/server/nativeapi/library_test.go +++ b/server/nativeapi/library_test.go @@ -13,6 +13,7 @@ import ( "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/server" "github.com/navidrome/navidrome/tests" @@ -21,17 +22,19 @@ import ( ) var _ = Describe("Library API", func() { + var ctx context.Context var ds model.DataStore var router http.Handler var adminUser, regularUser model.User var library1, library2 model.Library BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) conf.Server.EnableSharing = false ds = &tests.MockDataStore{} auth.Init(ds) - nativeRouter := New(ds, nil, nil, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) + nativeRouter := New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) router = server.JWTVerifier(nativeRouter) // Create test users @@ -63,10 +66,10 @@ var _ = Describe("Library API", func() { } // Store in mock datastore - Expect(ds.User(context.TODO()).Put(&adminUser)).To(Succeed()) - Expect(ds.User(context.TODO()).Put(®ularUser)).To(Succeed()) - Expect(ds.Library(context.TODO()).Put(&library1)).To(Succeed()) - Expect(ds.Library(context.TODO()).Put(&library2)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, ®ularUser)).To(Succeed()) + Expect(ds.Library().Put(ctx, &library1)).To(Succeed()) + Expect(ds.Library().Put(ctx, &library2)).To(Succeed()) }) Describe("Library CRUD Operations", func() { @@ -293,7 +296,7 @@ var _ = Describe("Library API", func() { Describe("GET /api/user/{id}/library", func() { It("returns user's libraries", func() { // Set up user libraries - err := ds.User(context.TODO()).SetUserLibraries(regularUser.ID, []int{1, 2}) + err := ds.User().SetUserLibraries(ctx, regularUser.ID, []int{1, 2}) Expect(err).ToNot(HaveOccurred()) req := createAuthenticatedRequest("GET", fmt.Sprintf("/user/%s/library", regularUser.ID), nil, adminToken) diff --git a/server/nativeapi/metadata_test.go b/server/nativeapi/metadata_test.go index ae6e301a8..294a26efe 100644 --- a/server/nativeapi/metadata_test.go +++ b/server/nativeapi/metadata_test.go @@ -11,6 +11,7 @@ import ( "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/core/auth" "github.com/navidrome/navidrome/core/external" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/server" "github.com/navidrome/navidrome/tests" @@ -38,6 +39,7 @@ func (f *fakeProvider) calls() []string { } var _ = Describe("Metadata API", func() { + var ctx context.Context var ds *tests.MockDataStore var artRepo *tests.MockArtworkRepo var queueRepo *tests.MockArtworkQueueRepo @@ -47,6 +49,7 @@ var _ = Describe("Metadata API", func() { var adminToken, userToken string BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) conf.Server.EnableSharing = false artRepo = tests.CreateMockArtworkRepo() @@ -54,9 +57,9 @@ var _ = Describe("Metadata API", func() { albumRepo = tests.CreateMockAlbumRepo() artistRepo := tests.CreateMockArtistRepo() playlistRepo := tests.CreateMockPlaylistRepo() - Expect(albumRepo.Put(&model.Album{ID: "al-1", Name: "Kid A"})).To(Succeed()) - Expect(artistRepo.Put(&model.Artist{ID: "ar-1", Name: "Radiohead"})).To(Succeed()) - Expect(playlistRepo.Put(&model.Playlist{ID: "pl-1", Name: "My Playlist"})).To(Succeed()) + Expect(albumRepo.Put(ctx, &model.Album{ID: "al-1", Name: "Kid A"})).To(Succeed()) + Expect(artistRepo.Put(ctx, &model.Artist{ID: "ar-1", Name: "Radiohead"})).To(Succeed()) + Expect(playlistRepo.Put(ctx, &model.Playlist{ID: "pl-1", Name: "My Playlist"})).To(Succeed()) ds = &tests.MockDataStore{ MockedArtwork: artRepo, MockedArtworkQueue: queueRepo, @@ -66,13 +69,13 @@ var _ = Describe("Metadata API", func() { } auth.Init(ds) provider = &fakeProvider{} - nativeRouter := New(ds, nil, nil, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, provider, nil) + nativeRouter := New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, provider, nil) router = server.JWTVerifier(nativeRouter) adminUser := model.User{ID: "admin-1", UserName: "admin", IsAdmin: true, NewPassword: "adminpass"} regularUser := model.User{ID: "user-1", UserName: "regular", IsAdmin: false, NewPassword: "userpass"} - Expect(ds.User(context.TODO()).Put(&adminUser)).To(Succeed()) - Expect(ds.User(context.TODO()).Put(®ularUser)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, ®ularUser)).To(Succeed()) var err error adminToken, err = auth.CreateToken(&adminUser) @@ -83,7 +86,7 @@ var _ = Describe("Metadata API", func() { Describe("POST /api/metadata/{kind}/{id}/refresh", func() { It("clears state and enqueues a Bump for admins", func() { - Expect(artRepo.PutItemArtwork(&model.ItemArtwork{ + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ ItemKind: "al", ItemID: "al-1", Hash: "oldhash", Source: "external", })).To(Succeed()) @@ -93,10 +96,10 @@ var _ = Describe("Metadata API", func() { Expect(w.Code).To(Equal(http.StatusNoContent)) - _, err := artRepo.GetItemArtwork(model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) + _, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "al-1", model.ImageTypePrimary) Expect(err).To(MatchError(model.ErrNotFound)) - queued, err := queueRepo.DequeueBatch(1000) + queued, err := queueRepo.DequeueBatch(ctx, 1000) Expect(err).ToNot(HaveOccurred()) Expect(queued).To(ContainElement(SatisfyAll( HaveField("ItemKind", "al"), diff --git a/server/nativeapi/missing.go b/server/nativeapi/missing.go index de6595e00..a906e36b9 100644 --- a/server/nativeapi/missing.go +++ b/server/nativeapi/missing.go @@ -14,24 +14,21 @@ import ( ) type missingRepository struct { - model.ResourceRepository + rest.Repository[model.MediaFile] mfRepo model.MediaFileRepository } -func newMissingRepository(ds model.DataStore) rest.RepositoryConstructor { - return func(ctx context.Context) rest.Repository { - return &missingRepository{mfRepo: ds.MediaFile(ctx), ResourceRepository: ds.Resource(ctx, model.MediaFile{})} - } +func newMissingRepository(ds model.DataStore) rest.Repository[model.MediaFile] { + mf := ds.MediaFile() + return &missingRepository{Repository: mf, mfRepo: mf} } -func (r *missingRepository) Count(options ...rest.QueryOptions) (int64, error) { - opt := r.parseOptions(options) - return r.ResourceRepository.Count(opt) +func (r *missingRepository) Count(ctx context.Context, options ...rest.QueryOptions) (int64, error) { + return r.Repository.Count(ctx, r.parseOptions(options)) } -func (r *missingRepository) ReadAll(options ...rest.QueryOptions) (any, error) { - opt := r.parseOptions(options) - return r.ResourceRepository.ReadAll(opt) +func (r *missingRepository) ReadAll(ctx context.Context, options ...rest.QueryOptions) ([]model.MediaFile, error) { + return r.Repository.ReadAll(ctx, r.parseOptions(options)) } func (r *missingRepository) parseOptions(options []rest.QueryOptions) rest.QueryOptions { @@ -44,8 +41,8 @@ func (r *missingRepository) parseOptions(options []rest.QueryOptions) rest.Query return opt } -func (r *missingRepository) Read(id string) (any, error) { - mf, err := r.mfRepo.Get(id) +func (r *missingRepository) Read(ctx context.Context, id string) (*model.MediaFile, error) { + mf, err := r.mfRepo.Get(ctx, id) if err != nil { return nil, err } @@ -55,10 +52,6 @@ func (r *missingRepository) Read(id string) (any, error) { return mf, nil } -func (r *missingRepository) EntityName() string { - return "missing_files" -} - func deleteMissingFiles(maintenance core.Maintenance) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -86,5 +79,3 @@ func deleteMissingFiles(maintenance core.Maintenance) http.HandlerFunc { writeDeleteManyResponse(w, r, ids) } } - -var _ model.ResourceRepository = &missingRepository{} diff --git a/server/nativeapi/missing_test.go b/server/nativeapi/missing_test.go index 9d7575a0c..a53f13fc3 100644 --- a/server/nativeapi/missing_test.go +++ b/server/nativeapi/missing_test.go @@ -9,6 +9,7 @@ import ( "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/server" "github.com/navidrome/navidrome/tests" @@ -35,12 +36,12 @@ var _ = Describe("Missing Files Endpoint", func() { auth.Init(ds) user := model.User{ID: "user-1", UserName: "user", NewPassword: "pass"} - Expect(userRepo.Put(&user)).To(Succeed()) + Expect(userRepo.Put(GinkgoT().Context(), &user)).To(Succeed()) var err error token, err = auth.CreateToken(&user) Expect(err).ToNot(HaveOccurred()) - router = server.JWTVerifier(New(ds, nil, nil, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil)) + router = server.JWTVerifier(New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil)) }) DescribeTable("GET /missing/{id}", diff --git a/server/nativeapi/native_api.go b/server/nativeapi/native_api.go index a7c53df09..97ad14be2 100644 --- a/server/nativeapi/native_api.go +++ b/server/nativeapi/native_api.go @@ -60,25 +60,25 @@ func (api *Router) routes() http.Handler { r := chi.NewRouter() // Public - api.RX(r, "/translation", newTranslationRepository, false) + rx(r, "/translation", newTranslationRepository(), false) // Protected r.Group(func(r chi.Router) { r.Use(server.Authenticator(api.ds)) r.Use(server.JWTRefresher) r.Use(server.UpdateLastAccessMiddleware(api.ds)) - api.RX(r, "/user", api.users.NewRepository, true) - api.R(r, "/song", model.MediaFile{}, false) - api.R(r, "/album", model.Album{}, false) + rx(r, "/user", api.users.Repository(), true) + rx(r, "/song", api.ds.MediaFile(), false) + rx(r, "/album", api.ds.Album(), false) api.addArtistRoute(r) - api.R(r, "/genre", model.Genre{}, false) - api.R(r, "/player", model.Player{}, true) - api.R(r, "/transcoding", model.Transcoding{}, conf.Server.EnableTranscodingConfig) + rx(r, "/genre", api.ds.Genre(), false) + rx(r, "/player", api.ds.Player(), true) + rx(r, "/transcoding", api.ds.Transcoding(), conf.Server.EnableTranscodingConfig) api.addRadioRoute(r) - api.R(r, "/tag", model.Tag{}, false) - api.R(r, "/scrobble", model.Scrobble{}, false) + rx(r, "/tag", api.ds.Tag(), false) + rx(r, "/scrobble", api.ds.Scrobble(), false) if conf.Server.EnableSharing { - api.RX(r, "/share", api.share.NewRepository, true) + rx(r, "/share", api.share.Repository(), true) } api.addPlaylistRoute(r) @@ -96,47 +96,38 @@ func (api *Router) routes() http.Handler { api.addUserLibraryRoute(r) api.addPluginRoute(r) api.addMetadataRoute(r) - api.RX(r, "/library", api.libs.NewRepository, true) + rx(r, "/library", api.libs.Repository(), true) }) }) return r } -func (api *Router) R(r chi.Router, pathPrefix string, model any, persistable bool) { - constructor := func(ctx context.Context) rest.Repository { - return api.ds.Resource(ctx, model) - } - api.RX(r, pathPrefix, constructor, persistable) -} - -func (api *Router) RX(r chi.Router, pathPrefix string, constructor rest.RepositoryConstructor, persistable bool) { +func rx[T any](r chi.Router, pathPrefix string, repo rest.Repository[T], persistable bool) { r.Route(pathPrefix, func(r chi.Router) { - r.Get("/", rest.GetAll(constructor)) + r.Get("/", rest.GetAll(repo)) if persistable { - r.Post("/", rest.Post(constructor)) + r.Post("/", rest.Post(repo)) } r.Route("/{id}", func(r chi.Router) { r.Use(server.URLParamsMiddleware) - r.Get("/", rest.Get(constructor)) + r.Get("/", rest.Get(repo)) if persistable { - r.Put("/", rest.Put(constructor)) - r.Delete("/", rest.Delete(constructor)) + r.Put("/", rest.Put(repo)) + r.Delete("/", rest.Delete(repo)) } }) }) } func (api *Router) addPlaylistRoute(r chi.Router) { - constructor := func(ctx context.Context) rest.Repository { - return api.playlists.NewRepository(ctx) - } + repo := api.playlists.Repository() r.Route("/playlist", func(r chi.Router) { - r.Get("/", rest.GetAll(constructor)) + r.Get("/", rest.GetAll(repo)) r.Post("/", func(w http.ResponseWriter, r *http.Request) { if r.Header.Get("Content-type") == "application/json" { - rest.Post(constructor)(w, r) + rest.Post(repo)(w, r) return } createPlaylistFromM3U(api.playlists)(w, r) @@ -144,9 +135,9 @@ func (api *Router) addPlaylistRoute(r chi.Router) { r.Route("/{id}", func(r chi.Router) { r.Use(server.URLParamsMiddleware) - r.Get("/", rest.Get(constructor)) - r.Put("/", rest.Put(constructor)) - r.Delete("/", rest.Delete(constructor)) + r.Get("/", rest.Get(repo)) + r.Put("/", rest.Put(repo)) + r.Delete("/", rest.Delete(repo)) r.Post("/image", uploadPlaylistImage(api.playlists)) r.Delete("/image", deletePlaylistImage(api.playlists)) }) @@ -198,7 +189,7 @@ func (api *Router) addQueueRoute(r chi.Router) { func (api *Router) addMissingFilesRoute(r chi.Router) { r.Route("/missing", func(r chi.Router) { - api.RX(r, "/", newMissingRepository(api.ds), false) + rx(r, "/", newMissingRepository(api.ds), false) r.Delete("/", deleteMissingFiles(api.maintenance)) }) } diff --git a/server/nativeapi/native_api_song_test.go b/server/nativeapi/native_api_song_test.go index 954c872a7..f151b1d72 100644 --- a/server/nativeapi/native_api_song_test.go +++ b/server/nativeapi/native_api_song_test.go @@ -2,6 +2,7 @@ package nativeapi import ( "bytes" + "context" "encoding/json" "net/http" "net/http/httptest" @@ -12,6 +13,7 @@ import ( "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/server" "github.com/navidrome/navidrome/tests" @@ -21,6 +23,7 @@ import ( var _ = Describe("Song Endpoints", func() { var ( + ctx context.Context router http.Handler ds *tests.MockDataStore mfRepo *tests.MockMediaFileRepo @@ -31,6 +34,7 @@ var _ = Describe("Song Endpoints", func() { ) BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) conf.Server.EnableSharing = false conf.Server.SessionTimeout = time.Minute @@ -56,7 +60,7 @@ var _ = Describe("Song Endpoints", func() { IsAdmin: false, NewPassword: "testpass", } - err := userRepo.Put(&testUser) + err := userRepo.Put(ctx, &testUser) Expect(err).ToNot(HaveOccurred()) // Create test songs @@ -95,7 +99,7 @@ var _ = Describe("Song Endpoints", func() { mfRepo.SetData(testSongs) // Create the native API router and wrap it with the JWTVerifier middleware - nativeRouter := New(ds, nil, nil, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) + nativeRouter := New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) router = server.JWTVerifier(nativeRouter) w = httptest.NewRecorder() }) @@ -369,7 +373,7 @@ var _ = Describe("Song Endpoints", func() { IsAdmin: true, NewPassword: "adminpass", } - err := userRepo.Put(&adminUser) + err := userRepo.Put(ctx, &adminUser) Expect(err).ToNot(HaveOccurred()) // Create JWT token for admin user @@ -392,7 +396,7 @@ var _ = Describe("Song Endpoints", func() { IsAdmin: false, NewPassword: "userpass", } - err := userRepo.Put(®ularUser) + err := userRepo.Put(ctx, ®ularUser) Expect(err).ToNot(HaveOccurred()) // Create JWT token for regular user diff --git a/server/nativeapi/playlists.go b/server/nativeapi/playlists.go index 00c1575a5..82f138492 100644 --- a/server/nativeapi/playlists.go +++ b/server/nativeapi/playlists.go @@ -19,8 +19,6 @@ import ( "github.com/navidrome/navidrome/utils/str" ) -type restHandler = func(rest.RepositoryConstructor, ...rest.Logger) http.HandlerFunc - // writePlaylistError maps a playlist service error to an HTTP status, or defaultStatus if unknown. func writePlaylistError(w http.ResponseWriter, err error, defaultStatus int) { switch { @@ -35,7 +33,7 @@ func writePlaylistError(w http.ResponseWriter, err error, defaultStatus int) { } } -func playlistTracksHandler(pls playlists.Playlists, handler restHandler, refreshSmartPlaylist func(*http.Request) bool) http.HandlerFunc { +func playlistTracksHandler(pls playlists.Playlists, handler func(rest.Repository[model.PlaylistTrack]) http.HandlerFunc, refreshSmartPlaylist func(*http.Request) bool) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { plsId := chi.URLParam(r, "playlistId") tracks := pls.TracksRepository(r.Context(), plsId, refreshSmartPlaylist(r)) @@ -43,12 +41,12 @@ func playlistTracksHandler(pls playlists.Playlists, handler restHandler, refresh http.Error(w, "not found", http.StatusNotFound) return } - handler(func(ctx context.Context) rest.Repository { return tracks }).ServeHTTP(w, r) + handler(tracks).ServeHTTP(w, r) } } func getPlaylist(pls playlists.Playlists) http.HandlerFunc { - handler := playlistTracksHandler(pls, rest.GetAll, func(r *http.Request) bool { + handler := playlistTracksHandler(pls, rest.GetAll[model.PlaylistTrack], func(r *http.Request) bool { return req.Params(r).Int64Or("_start", 0) == 0 }) return func(w http.ResponseWriter, r *http.Request) { @@ -61,7 +59,7 @@ func getPlaylist(pls playlists.Playlists) http.HandlerFunc { } func getPlaylistTrack(pls playlists.Playlists) http.HandlerFunc { - return playlistTracksHandler(pls, rest.Get, func(*http.Request) bool { return true }) + return playlistTracksHandler(pls, rest.Get[model.PlaylistTrack], func(*http.Request) bool { return true }) } func createPlaylistFromM3U(pls playlists.Playlists) http.HandlerFunc { diff --git a/server/nativeapi/playlists_test.go b/server/nativeapi/playlists_test.go index 349b4a662..82e3bc86a 100644 --- a/server/nativeapi/playlists_test.go +++ b/server/nativeapi/playlists_test.go @@ -97,7 +97,7 @@ var _ = Describe("Playlist Tracks Endpoint", func() { IsAdmin: false, NewPassword: "testpass", } - err := userRepo.Put(&testUser) + err := userRepo.Put(GinkgoT().Context(), &testUser) Expect(err).ToNot(HaveOccurred()) nativeRouter := New(ds, nil, plsSvc, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) @@ -235,23 +235,15 @@ type mockPlaylistTrackRepo struct { tracks model.PlaylistTracks } -func (m *mockPlaylistTrackRepo) Count(...rest.QueryOptions) (int64, error) { +func (m *mockPlaylistTrackRepo) Count(context.Context, ...rest.QueryOptions) (int64, error) { return int64(len(m.tracks)), nil } -func (m *mockPlaylistTrackRepo) ReadAll(...rest.QueryOptions) (any, error) { +func (m *mockPlaylistTrackRepo) ReadAll(context.Context, ...rest.QueryOptions) ([]model.PlaylistTrack, error) { return m.tracks, nil } -func (m *mockPlaylistTrackRepo) EntityName() string { - return "playlist_track" -} - -func (m *mockPlaylistTrackRepo) NewInstance() any { - return &model.PlaylistTrack{} -} - -func (m *mockPlaylistTrackRepo) Read(id string) (any, error) { +func (m *mockPlaylistTrackRepo) Read(_ context.Context, id string) (*model.PlaylistTrack, error) { for _, t := range m.tracks { if t.ID == id { return &t, nil @@ -262,7 +254,8 @@ func (m *mockPlaylistTrackRepo) Read(id string) (any, error) { type mockPlaylistsService struct { playlists.Playlists - tracksRepo rest.Repository + repo rest.Repository[model.Playlist] + tracksRepo rest.Repository[model.PlaylistTrack] playlist *model.Playlist removeImageFn func(ctx context.Context, id string) error setImageFn func(ctx context.Context, id string, reader io.Reader, ext string) error @@ -282,6 +275,10 @@ func (m *mockPlaylistsService) SetImage(ctx context.Context, id string, reader i return model.ErrNotFound } +func (m *mockPlaylistsService) Repository() rest.Repository[model.Playlist] { + return m.repo +} + func (m *mockPlaylistsService) GetWithTracks(_ context.Context, _ string) (*model.Playlist, error) { if m.playlist == nil { return nil, model.ErrNotFound @@ -289,6 +286,6 @@ func (m *mockPlaylistsService) GetWithTracks(_ context.Context, _ string) (*mode return m.playlist, nil } -func (m *mockPlaylistsService) TracksRepository(_ context.Context, _ string, _ bool) rest.Repository { +func (m *mockPlaylistsService) TracksRepository(_ context.Context, _ string, _ bool) rest.Repository[model.PlaylistTrack] { return m.tracksRepo } diff --git a/server/nativeapi/plugin.go b/server/nativeapi/plugin.go index a7d261681..d34bf23fc 100644 --- a/server/nativeapi/plugin.go +++ b/server/nativeapi/plugin.go @@ -15,17 +15,15 @@ import ( ) func (api *Router) addPluginRoute(r chi.Router) { - constructor := func(ctx context.Context) rest.Repository { - return api.ds.Plugin(ctx) - } + repo := api.ds.Plugin() r.Route("/plugin", func(r chi.Router) { r.Use(pluginsEnabledMiddleware) - r.Get("/", rest.GetAll(constructor)) + r.Get("/", rest.GetAll(repo)) r.Post("/rescan", api.rescanPlugins) r.Route("/{id}", func(r chi.Router) { r.Use(server.URLParamsMiddleware) - r.Get("/", rest.Get(constructor)) + r.Get("/", rest.Get(repo)) r.Put("/", api.updatePlugin) }) }) @@ -68,10 +66,10 @@ type PluginUpdateRequest struct { func (api *Router) updatePlugin(w http.ResponseWriter, r *http.Request) { id := chi.URLParam(r, "id") ctx := r.Context() - repo := api.ds.Plugin(ctx) + repo := api.ds.Plugin() // Get existing plugin to verify it exists - if _, err := repo.Get(id); err != nil { + if _, err := repo.Get(ctx, id); err != nil { if errors.Is(err, rest.ErrPermissionDenied) { http.Error(w, "Access denied: admin privileges required", http.StatusForbidden) return @@ -123,7 +121,7 @@ func (api *Router) updatePlugin(w http.ResponseWriter, r *http.Request) { if enableErr := api.pluginManager.EnablePlugin(ctx, id); enableErr != nil { log.Error(ctx, "Error enabling plugin", "id", id, enableErr) // Refresh plugin from DB to get the error - plugin, err := repo.Get(id) + plugin, err := repo.Get(ctx, id) if err != nil { log.Error(ctx, "Error getting updated plugin after enable failure", "id", id, err) http.Error(w, "Internal server error", http.StatusInternalServerError) @@ -153,7 +151,7 @@ func (api *Router) updatePlugin(w http.ResponseWriter, r *http.Request) { } // Refresh and return updated plugin - plugin, err := repo.Get(id) + plugin, err := repo.Get(ctx, id) if err != nil { log.Error(ctx, "Error getting updated plugin", "id", id, err) http.Error(w, "Internal server error", http.StatusInternalServerError) @@ -204,7 +202,7 @@ func validateAndUpdateConfig(ctx context.Context, pm PluginManager, id, configJS // Returns an error if validation or update fails (error response already written). func validateAndUpdateUsers(ctx context.Context, pm PluginManager, repo model.PluginRepository, id string, req PluginUpdateRequest, w http.ResponseWriter) error { // Get current values if not provided in request - plugin, err := repo.Get(id) + plugin, err := repo.Get(ctx, id) if err != nil { log.Error(ctx, "Error getting plugin for users update", "id", id, err) http.Error(w, "Internal server error", http.StatusInternalServerError) @@ -237,7 +235,7 @@ func validateAndUpdateUsers(ctx context.Context, pm PluginManager, repo model.Pl // Returns an error if validation or update fails (error response already written). func validateAndUpdateLibraries(ctx context.Context, pm PluginManager, repo model.PluginRepository, id string, req PluginUpdateRequest, w http.ResponseWriter) error { // Get current values if not provided in request - plugin, err := repo.Get(id) + plugin, err := repo.Get(ctx, id) if err != nil { log.Error(ctx, "Error getting plugin for libraries update", "id", id, err) http.Error(w, "Internal server error", http.StatusInternalServerError) diff --git a/server/nativeapi/plugin_test.go b/server/nativeapi/plugin_test.go index 4e45ddb92..c18d61e65 100644 --- a/server/nativeapi/plugin_test.go +++ b/server/nativeapi/plugin_test.go @@ -12,6 +12,7 @@ import ( "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server" @@ -21,6 +22,7 @@ import ( ) var _ = Describe("Plugin API", func() { + var ctx context.Context var ds *tests.MockDataStore var mockManager *tests.MockPluginManager var router http.Handler @@ -28,13 +30,14 @@ var _ = Describe("Plugin API", func() { var testPlugin1, testPlugin2 model.Plugin BeforeEach(func() { + ctx = GinkgoT().Context() DeferCleanup(configtest.SetupConfig()) conf.Server.EnableSharing = false conf.Server.Plugins.Enabled = true ds = &tests.MockDataStore{} mockManager = &tests.MockPluginManager{} auth.Init(ds) - nativeRouter := New(ds, nil, nil, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, mockManager, nil, nil, nil) + nativeRouter := New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, mockManager, nil, nil, nil) router = server.JWTVerifier(nativeRouter) // Create test users @@ -71,8 +74,8 @@ var _ = Describe("Plugin API", func() { } // Store users in mock datastore - Expect(ds.User(GinkgoT().Context()).Put(&adminUser)).To(Succeed()) - Expect(ds.User(GinkgoT().Context()).Put(®ularUser)).To(Succeed()) + Expect(ds.User().Put(ctx, &adminUser)).To(Succeed()) + Expect(ds.User().Put(ctx, ®ularUser)).To(Succeed()) }) Context("when plugins are disabled", func() { @@ -104,10 +107,9 @@ var _ = Describe("Plugin API", func() { Expect(err).ToNot(HaveOccurred()) // Store test plugins as admin - ctx := GinkgoT().Context() adminCtx := request.WithUser(ctx, adminUser) - Expect(ds.Plugin(adminCtx).Put(&testPlugin1)).To(Succeed()) - Expect(ds.Plugin(adminCtx).Put(&testPlugin2)).To(Succeed()) + Expect(ds.Plugin().Put(adminCtx, &testPlugin1)).To(Succeed()) + Expect(ds.Plugin().Put(adminCtx, &testPlugin2)).To(Succeed()) }) Describe("GET /api/plugin", func() { @@ -160,9 +162,9 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when EnablePlugin is called mockManager.EnablePluginFn = func(ctx context.Context, id string) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Enabled = true - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"enabled":true}`) @@ -186,9 +188,9 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when UpdatePluginConfig is called mockManager.UpdatePluginConfigFn = func(ctx context.Context, id, configJSON string) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Config = configJSON - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"config":"{\"key\":\"value\"}"}`) @@ -226,9 +228,9 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when UpdatePluginConfig is called mockManager.UpdatePluginConfigFn = func(ctx context.Context, id, configJSON string) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Config = configJSON - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"config":""}`) @@ -251,10 +253,10 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when UpdatePluginUsers is called mockManager.UpdatePluginUsersFn = func(ctx context.Context, id, usersJSON string, allUsers bool) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Users = usersJSON p.AllUsers = allUsers - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"users":"[\"user1\",\"user2\"]"}`) @@ -279,10 +281,10 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when UpdatePluginUsers is called mockManager.UpdatePluginUsersFn = func(ctx context.Context, id, usersJSON string, allUsers bool) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Users = usersJSON p.AllUsers = allUsers - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"allUsers":true}`) @@ -307,10 +309,10 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when UpdatePluginUsers is called mockManager.UpdatePluginUsersFn = func(ctx context.Context, id, usersJSON string, allUsers bool) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Users = usersJSON p.AllUsers = allUsers - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"users":"[\"user1\"]","allUsers":false}`) @@ -348,10 +350,10 @@ var _ = Describe("Plugin API", func() { // Configure mock to update the repo when UpdatePluginUsers is called mockManager.UpdatePluginUsersFn = func(ctx context.Context, id, usersJSON string, allUsers bool) error { adminCtx := request.WithUser(ctx, adminUser) - p, _ := ds.Plugin(adminCtx).Get(id) + p, _ := ds.Plugin().Get(adminCtx, id) p.Users = usersJSON p.AllUsers = allUsers - return ds.Plugin(adminCtx).Put(p) + return ds.Plugin().Put(adminCtx, p) } body := bytes.NewBufferString(`{"users":""}`) diff --git a/server/nativeapi/queue.go b/server/nativeapi/queue.go index a7700c02c..05188106c 100644 --- a/server/nativeapi/queue.go +++ b/server/nativeapi/queue.go @@ -32,7 +32,7 @@ func validateCurrentIndex(w http.ResponseWriter, current int, itemsLength int) b // retrieveExistingQueue retrieves an existing play queue for a user with proper error handling. // Returns the queue (nil if not found) and false if an error occurred and response was sent. func retrieveExistingQueue(ctx context.Context, w http.ResponseWriter, ds model.DataStore, userID string) (*model.PlayQueue, bool) { - existing, err := ds.PlayQueue(ctx).Retrieve(userID) + existing, err := ds.PlayQueue().Retrieve(ctx, userID) if err != nil && !errors.Is(err, model.ErrNotFound) { log.Error(ctx, "Error retrieving queue", err) http.Error(w, err.Error(), http.StatusInternalServerError) @@ -70,8 +70,8 @@ func getQueue(ds model.DataStore) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() user, _ := request.UserFrom(ctx) - repo := ds.PlayQueue(ctx) - pq, err := repo.RetrieveWithMediaFiles(user.ID) + repo := ds.PlayQueue() + pq, err := repo.RetrieveWithMediaFiles(ctx, user.ID) if err != nil && !errors.Is(err, model.ErrNotFound) { log.Error(ctx, "Error retrieving queue", err) http.Error(w, err.Error(), http.StatusInternalServerError) @@ -112,7 +112,7 @@ func saveQueue(ds model.DataStore) http.HandlerFunc { ChangedBy: client, Items: items, } - if err := ds.PlayQueue(ctx).Store(pq); err != nil { + if err := ds.PlayQueue().Store(ctx, pq); err != nil { log.Error(ctx, "Error saving queue", err) http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -191,7 +191,7 @@ func updateQueue(ds model.DataStore) http.HandlerFunc { } // Perform partial update of the specified columns only - if err := ds.PlayQueue(ctx).Store(pq, cols...); err != nil { + if err := ds.PlayQueue().Store(ctx, pq, cols...); err != nil { log.Error(ctx, "Error updating queue", err) http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -204,7 +204,7 @@ func clearQueue(ds model.DataStore) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() user, _ := request.UserFrom(ctx) - if err := ds.PlayQueue(ctx).Clear(user.ID); err != nil { + if err := ds.PlayQueue().Clear(ctx, user.ID); err != nil { log.Error(ctx, "Error clearing queue", err) http.Error(w, err.Error(), http.StatusInternalServerError) return diff --git a/server/nativeapi/queue_test.go b/server/nativeapi/queue_test.go index 0aad09718..6e68fec0e 100644 --- a/server/nativeapi/queue_test.go +++ b/server/nativeapi/queue_test.go @@ -25,7 +25,7 @@ var _ = Describe("Queue Endpoints", func() { repo = &tests.MockPlayQueueRepo{} user = model.User{ID: "u1", UserName: "user"} userRepo = tests.CreateMockUserRepo() - _ = userRepo.Put(&user) + _ = userRepo.Put(GinkgoT().Context(), &user) ds = &tests.MockDataStore{MockedPlayQueue: repo, MockedUser: userRepo, MockedProperty: &tests.MockedPropertyRepo{}} }) diff --git a/server/nativeapi/radios.go b/server/nativeapi/radios.go index 3e88af287..4bacbea7a 100644 --- a/server/nativeapi/radios.go +++ b/server/nativeapi/radios.go @@ -14,17 +14,15 @@ import ( ) func (api *Router) addRadioRoute(r chi.Router) { - constructor := func(ctx context.Context) rest.Repository { - return api.ds.Resource(ctx, model.Radio{}) - } + repo := api.ds.Radio() r.Route("/radio", func(r chi.Router) { - r.Get("/", rest.GetAll(constructor)) - r.Post("/", rest.Post(constructor)) + r.Get("/", rest.GetAll(repo)) + r.Post("/", rest.Post(repo)) r.Route("/{id}", func(r chi.Router) { r.Use(server.URLParamsMiddleware) - r.Get("/", rest.Get(constructor)) - r.Put("/", rest.Put(constructor)) - r.Delete("/", rest.Delete(constructor)) + r.Get("/", rest.Get(repo)) + r.Put("/", rest.Put(repo)) + r.Delete("/", rest.Delete(repo)) r.Post("/image", api.uploadRadioImage()) r.Delete("/image", api.deleteRadioImage()) }) @@ -34,7 +32,7 @@ func (api *Router) addRadioRoute(r chi.Router) { func (api *Router) uploadRadioImage() http.HandlerFunc { return handleImageUpload(func(ctx context.Context, reader io.Reader, ext string) error { radioID := chi.URLParamFromCtx(ctx, "id") - radio, err := api.ds.Radio(ctx).Get(radioID) + radio, err := api.ds.Radio().Get(ctx, radioID) if err != nil { if errors.Is(err, model.ErrNotFound) { return model.ErrNotFound @@ -47,7 +45,7 @@ func (api *Router) uploadRadioImage() http.HandlerFunc { return err } radio.UploadedImage = filename - if err := api.ds.Radio(ctx).Put(radio, "UploadedImage"); err != nil { + if err := api.ds.Radio().Put(ctx, radio, "UploadedImage"); err != nil { return err } api.imgUpload.EnqueueArtwork(ctx, consts.EntityRadio, radio.ID) @@ -58,7 +56,7 @@ func (api *Router) uploadRadioImage() http.HandlerFunc { func (api *Router) deleteRadioImage() http.HandlerFunc { return handleImageDelete(func(ctx context.Context) error { radioID := chi.URLParamFromCtx(ctx, "id") - radio, err := api.ds.Radio(ctx).Get(radioID) + radio, err := api.ds.Radio().Get(ctx, radioID) if err != nil { if errors.Is(err, model.ErrNotFound) { return model.ErrNotFound @@ -69,7 +67,7 @@ func (api *Router) deleteRadioImage() http.HandlerFunc { return err } radio.UploadedImage = "" - if err := api.ds.Radio(ctx).Put(radio, "UploadedImage"); err != nil { + if err := api.ds.Radio().Put(ctx, radio, "UploadedImage"); err != nil { return err } api.imgUpload.EnqueueArtwork(ctx, consts.EntityRadio, radio.ID) diff --git a/server/nativeapi/translations.go b/server/nativeapi/translations.go index 39d071279..fc4651c4c 100644 --- a/server/nativeapi/translations.go +++ b/server/nativeapi/translations.go @@ -23,28 +23,28 @@ type translation struct { TermCount int `json:"termCount"` } -func newTranslationRepository(context.Context) rest.Repository { +func newTranslationRepository() rest.Repository[translation] { return &translationRepository{} } type translationRepository struct{} -func (r *translationRepository) Read(id string) (any, error) { +func (r *translationRepository) Read(_ context.Context, id string) (*translation, error) { translations, _ := loadTranslations() if t, ok := translations[id]; ok { - return t, nil + return &t, nil } return nil, rest.ErrNotFound } // Count simple implementation, does not support any `options` -func (r *translationRepository) Count(...rest.QueryOptions) (int64, error) { +func (r *translationRepository) Count(context.Context, ...rest.QueryOptions) (int64, error) { _, count := loadTranslations() return count, nil } // ReadAll simple implementation, only returns IDs. Does not support any `options` -func (r *translationRepository) ReadAll(...rest.QueryOptions) (any, error) { +func (r *translationRepository) ReadAll(context.Context, ...rest.QueryOptions) ([]translation, error) { translations, _ := loadTranslations() var result []translation for _, t := range translations { @@ -54,14 +54,6 @@ func (r *translationRepository) ReadAll(...rest.QueryOptions) (any, error) { return result, nil } -func (r *translationRepository) EntityName() string { - return "translation" -} - -func (r *translationRepository) NewInstance() any { - return &translation{} -} - var loadTranslations = sync.OnceValues(func() (map[string]translation, int64) { translations := make(map[string]translation) fsys := resources.FS() @@ -140,4 +132,4 @@ func countTranslatedTerms(obj map[string]any) int { return count } -var _ rest.Repository = (*translationRepository)(nil) +var _ rest.Repository[translation] = (*translationRepository)(nil) diff --git a/server/nativeapi/user_password_token_refresh_test.go b/server/nativeapi/user_password_token_refresh_test.go index 27454fc7d..81f28893f 100644 --- a/server/nativeapi/user_password_token_refresh_test.go +++ b/server/nativeapi/user_password_token_refresh_test.go @@ -14,6 +14,7 @@ import ( "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/persistence" @@ -29,10 +30,12 @@ func (noopPluginUnloader) UnloadDisabledPlugins(context.Context) {} // Pins that the token-epoch handoff survives a real request through the real middleware chain. var _ = Describe("PUT /user/{id}: token refresh on self password change", func() { + var ctx context.Context var ds model.DataStore var router http.Handler BeforeEach(func() { + ctx = GinkgoT().Context() // db.Db() is a process-wide singleton that this DeferCleanup closes for the whole binary; keep this the only real-DB spec in this package. DeferCleanup(configtest.SetupConfig()) conf.Server.EnableUserEditing = true @@ -45,13 +48,13 @@ var _ = Describe("PUT /user/{id}: token refresh on self password change", func() auth.Init(ds) userService := core.NewUser(ds, noopPluginUnloader{}) - nativeRouter := New(ds, nil, nil, nil, tests.NewMockLibraryService(), userService, nil, nil, nil, nil, nil) + nativeRouter := New(ds, nil, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), userService, nil, nil, nil, nil, nil) router = server.JWTVerifier(nativeRouter) }) It("carries the bumped epoch in the refreshed token, not the epoch the token was minted with", func() { usr := model.User{UserName: "selfchanger", Name: "Self Changer", NewPassword: "old-password"} - Expect(ds.User(GinkgoT().Context()).Put(&usr)).To(Succeed()) + Expect(ds.User().Put(ctx, &usr)).To(Succeed()) token, err := auth.CreateToken(&usr) Expect(err).ToNot(HaveOccurred()) @@ -72,7 +75,7 @@ var _ = Describe("PUT /user/{id}: token refresh on self password change", func() claims, err := auth.Validate(refreshed) Expect(err).ToNot(HaveOccurred()) - reloaded, err := ds.User(GinkgoT().Context()).Get(usr.ID) + reloaded, err := ds.User().Get(ctx, usr.ID) Expect(err).ToNot(HaveOccurred()) Expect(reloaded.TokenEpoch).To(Equal(1)) Expect(claims.Epoch).To(Equal(reloaded.TokenEpoch)) diff --git a/server/public/handle_streams.go b/server/public/handle_streams.go index 3d624f661..37ae56c2b 100644 --- a/server/public/handle_streams.go +++ b/server/public/handle_streams.go @@ -26,7 +26,7 @@ func (pub *Router) handleStream(w http.ResponseWriter, r *http.Request) { return } - share, err := pub.ds.Share(ctx).Get(info.shareID) + share, err := pub.ds.Share().Get(ctx, info.shareID) if err != nil { checkShareError(ctx, w, err, info.shareID) return @@ -35,14 +35,14 @@ func (pub *Router) handleStream(w http.ResponseWriter, r *http.Request) { checkShareError(ctx, w, model.ErrExpired, info.shareID) return } - shareOwner, err := pub.ds.User(ctx).Get(share.UserID) + shareOwner, err := pub.ds.User().Get(ctx, share.UserID) if err != nil { log.Error(ctx, "Error retrieving share owner for shared stream", "share", info.shareID, "owner", share.UserID, err) http.Error(w, "internal error", http.StatusInternalServerError) return } - mf, err := pub.ds.MediaFile(ctx).Get(info.id) + mf, err := pub.ds.MediaFile().Get(ctx, info.id) if err != nil { if errors.Is(err, model.ErrNotFound) { http.Error(w, "not found", http.StatusNotFound) diff --git a/server/public/handle_streams_test.go b/server/public/handle_streams_test.go index 870dfa8ef..4b4a3545b 100644 --- a/server/public/handle_streams_test.go +++ b/server/public/handle_streams_test.go @@ -107,12 +107,14 @@ var _ = Describe("encodeMediafileShare", func() { }) var _ = Describe("handleStream", func() { + var ctx context.Context var ds *tests.MockDataStore var shareRepo *tests.MockShareRepo var streamer *mockStreamer var pub *Router BeforeEach(func() { + ctx = GinkgoT().Context() auth.PublicTokenAuth = jwtauth.New("HS256", []byte("test-secret"), nil) ds = &tests.MockDataStore{} shareRepo = &tests.MockShareRepo{} @@ -132,7 +134,7 @@ var _ = Describe("handleStream", func() { shareRepo.ID = "share123" shareRepo.Entity = &model.Share{ID: "share123", UserID: owner.ID, Tracks: model.MediaFiles{mf}} userRepo := tests.CreateMockUserRepo() - Expect(userRepo.Put(&owner)).To(Succeed()) + Expect(userRepo.Put(ctx, &owner)).To(Succeed()) ds.MockedUser = userRepo mfRepo := tests.CreateMockMediaFileRepo() mfRepo.SetData(model.MediaFiles{mf}) @@ -171,7 +173,7 @@ var _ = Describe("handleStream", func() { It("returns 404 when the track is not a member of the share", func() { owner := model.User{ID: "owner1", UserName: "owner1", IsAdmin: true} userRepo := tests.CreateMockUserRepo() - Expect(userRepo.Put(&owner)).To(Succeed()) + Expect(userRepo.Put(ctx, &owner)).To(Succeed()) ds.MockedUser = userRepo mfRepo := tests.CreateMockMediaFileRepo() mfRepo.SetData(model.MediaFiles{{ID: "mf-shared"}, {ID: "mf-other"}}) diff --git a/server/serve_index.go b/server/serve_index.go index 651c8c907..4b093b953 100644 --- a/server/serve_index.go +++ b/server/serve_index.go @@ -32,7 +32,7 @@ func IndexWithShare(ds model.DataStore, fs fs.FS, shareInfo *model.Share) http.H // Injects the config in the `index.html` template func serveIndex(ds model.DataStore, fs fs.FS, shareInfo *model.Share) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - c, err := ds.User(r.Context()).CountAll() + c, err := ds.User().CountAll(r.Context()) firstTime := c == 0 && err == nil t, err := getIndexTemplate(r, fs) diff --git a/server/serve_index_test.go b/server/serve_index_test.go index 23215a5ef..e2df55c4b 100644 --- a/server/serve_index_test.go +++ b/server/serve_index_test.go @@ -1,6 +1,7 @@ package server import ( + "context" "encoding/json" "fmt" "net/http" @@ -341,7 +342,7 @@ type mockedUserRepo struct { empty bool } -func (u *mockedUserRepo) CountAll(...model.QueryOptions) (int64, error) { +func (u *mockedUserRepo) CountAll(context.Context, ...model.QueryOptions) (int64, error) { if u.empty { return 0, nil } diff --git a/server/subsonic/album_lists.go b/server/subsonic/album_lists.go index ba5f88a68..0a0c65ec7 100644 --- a/server/subsonic/album_lists.go +++ b/server/subsonic/album_lists.go @@ -71,14 +71,14 @@ func (api *Router) getAlbumList(r *http.Request) (model.Albums, int64, error) { opts.Offset = p.IntOr("offset", 0) opts.Max = min(p.IntOr("size", 10), 500) - albums, err := api.ds.Album(r.Context()).GetAll(opts) + albums, err := api.ds.Album().GetAll(r.Context(), opts) if err != nil { log.Error(r, "Error retrieving albums", err) return nil, 0, newError(responses.ErrorGeneric, "internal error") } - count, err := api.ds.Album(r.Context()).CountAll(opts) + count, err := api.ds.Album().CountAll(r.Context(), opts) if err != nil { log.Error(r, "Error counting albums", err) return nil, 0, newError(responses.ErrorGeneric, "internal error") @@ -137,7 +137,7 @@ func (api *Router) getStarredItems(r *http.Request) (model.Artists, model.Albums func() error { artistOpts := filter.ApplyArtistLibraryFilter(filter.ArtistsByStarred(), musicFolderIds) var err error - artists, err = api.ds.Artist(ctx).GetAll(artistOpts) + artists, err = api.ds.Artist().GetAll(ctx, artistOpts) if err != nil { log.Error(r, "Error retrieving starred artists", err) } @@ -147,7 +147,7 @@ func (api *Router) getStarredItems(r *http.Request) (model.Artists, model.Albums func() error { albumOpts := filter.ApplyLibraryFilter(filter.ByStarred(), musicFolderIds) var err error - albums, err = api.ds.Album(ctx).GetAll(albumOpts) + albums, err = api.ds.Album().GetAll(ctx, albumOpts) if err != nil { log.Error(r, "Error retrieving starred albums", err) } @@ -157,7 +157,7 @@ func (api *Router) getStarredItems(r *http.Request) (model.Artists, model.Albums func() error { mediaFileOpts := filter.ApplyLibraryFilter(filter.ByStarred(), musicFolderIds) var err error - mediaFiles, err = api.ds.MediaFile(ctx).GetAll(mediaFileOpts) + mediaFiles, err = api.ds.MediaFile().GetAll(ctx, mediaFileOpts) if err != nil { log.Error(r, "Error retrieving starred mediaFiles", err) } @@ -244,7 +244,7 @@ func (api *Router) GetRandomSongs(r *http.Request) (*responses.Subsonic, error) opts = filter.ApplyLibraryFilter(opts, musicFolderIds) opts.Max = size - songs, err := api.ds.MediaFile(r.Context()).GetRandom(opts) + songs, err := api.ds.MediaFile().GetRandom(r.Context(), opts) if err != nil { log.Error(r, "Error retrieving random songs", err) return nil, err @@ -286,5 +286,5 @@ func (api *Router) GetSongsByGenre(r *http.Request) (*responses.Subsonic, error) func (api *Router) getSongs(ctx context.Context, offset, size int, opts filter.Options) (model.MediaFiles, error) { opts.Offset = offset opts.Max = size - return api.ds.MediaFile(ctx).GetAll(opts) + return api.ds.MediaFile().GetAll(ctx, opts) } diff --git a/server/subsonic/album_lists_test.go b/server/subsonic/album_lists_test.go index 220376b15..c4a8847d4 100644 --- a/server/subsonic/album_lists_test.go +++ b/server/subsonic/album_lists_test.go @@ -6,7 +6,6 @@ import ( "net/http/httptest" "github.com/navidrome/navidrome/core/auth" - "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server/subsonic/responses" @@ -21,12 +20,11 @@ var _ = Describe("Album Lists", func() { var ds model.DataStore var mockRepo *tests.MockAlbumRepo var w *httptest.ResponseRecorder - ctx := log.NewContext(context.TODO()) BeforeEach(func() { ds = &tests.MockDataStore{} auth.Init(ds) - mockRepo = ds.Album(ctx).(*tests.MockAlbumRepo) + mockRepo = ds.Album().(*tests.MockAlbumRepo) router = New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) w = httptest.NewRecorder() }) @@ -236,7 +234,7 @@ var _ = Describe("Album Lists", func() { var mockMediaFileRepo *tests.MockMediaFileRepo BeforeEach(func() { - mockMediaFileRepo = ds.MediaFile(ctx).(*tests.MockMediaFileRepo) + mockMediaFileRepo = ds.MediaFile().(*tests.MockMediaFileRepo) }) It("should return random songs", func() { @@ -328,7 +326,7 @@ var _ = Describe("Album Lists", func() { var mockMediaFileRepo *tests.MockMediaFileRepo BeforeEach(func() { - mockMediaFileRepo = ds.MediaFile(ctx).(*tests.MockMediaFileRepo) + mockMediaFileRepo = ds.MediaFile().(*tests.MockMediaFileRepo) }) It("should return songs by genre", func() { @@ -422,9 +420,9 @@ var _ = Describe("Album Lists", func() { var mockMediaFileRepo *tests.MockMediaFileRepo BeforeEach(func() { - mockArtistRepo = ds.Artist(ctx).(*tests.MockArtistRepo) - mockAlbumRepo = ds.Album(ctx).(*tests.MockAlbumRepo) - mockMediaFileRepo = ds.MediaFile(ctx).(*tests.MockMediaFileRepo) + mockArtistRepo = ds.Artist().(*tests.MockArtistRepo) + mockAlbumRepo = ds.Album().(*tests.MockAlbumRepo) + mockMediaFileRepo = ds.MediaFile().(*tests.MockMediaFileRepo) }) It("should return starred items", func() { @@ -484,9 +482,9 @@ var _ = Describe("Album Lists", func() { var mockMediaFileRepo *tests.MockMediaFileRepo BeforeEach(func() { - mockArtistRepo = ds.Artist(ctx).(*tests.MockArtistRepo) - mockAlbumRepo = ds.Album(ctx).(*tests.MockAlbumRepo) - mockMediaFileRepo = ds.MediaFile(ctx).(*tests.MockMediaFileRepo) + mockArtistRepo = ds.Artist().(*tests.MockArtistRepo) + mockAlbumRepo = ds.Album().(*tests.MockAlbumRepo) + mockMediaFileRepo = ds.MediaFile().(*tests.MockMediaFileRepo) }) It("should return starred items in ID3 format", func() { diff --git a/server/subsonic/bookmarks.go b/server/subsonic/bookmarks.go index 7ac492ca8..6a4c4962d 100644 --- a/server/subsonic/bookmarks.go +++ b/server/subsonic/bookmarks.go @@ -15,8 +15,8 @@ import ( func (api *Router) GetBookmarks(r *http.Request) (*responses.Subsonic, error) { user, _ := request.UserFrom(r.Context()) - repo := api.ds.MediaFile(r.Context()) - bookmarks, err := repo.GetBookmarks() + repo := api.ds.MediaFile() + bookmarks, err := repo.GetBookmarks(r.Context()) if err != nil { return nil, err } @@ -46,8 +46,8 @@ func (api *Router) CreateBookmark(r *http.Request) (*responses.Subsonic, error) comment, _ := p.String("comment") position := p.Int64Or("position", 0) - repo := api.ds.MediaFile(r.Context()) - ok, err := repo.Exists(id) + repo := api.ds.MediaFile() + ok, err := repo.Exists(r.Context(), id) if err != nil { return nil, err } @@ -55,7 +55,7 @@ func (api *Router) CreateBookmark(r *http.Request) (*responses.Subsonic, error) return nil, newError(responses.ErrorDataNotFound, "Song not found") } - err = repo.AddBookmark(id, comment, position) + err = repo.AddBookmark(r.Context(), id, comment, position) if err != nil { return nil, err } @@ -69,8 +69,8 @@ func (api *Router) DeleteBookmark(r *http.Request) (*responses.Subsonic, error) return nil, err } - repo := api.ds.MediaFile(r.Context()) - err = repo.DeleteBookmark(id) + repo := api.ds.MediaFile() + err = repo.DeleteBookmark(r.Context(), id) if err != nil { return nil, err } @@ -80,8 +80,8 @@ func (api *Router) DeleteBookmark(r *http.Request) (*responses.Subsonic, error) func (api *Router) GetPlayQueue(r *http.Request) (*responses.Subsonic, error) { user, _ := request.UserFrom(r.Context()) - repo := api.ds.PlayQueue(r.Context()) - pq, err := repo.RetrieveWithMediaFiles(user.ID) + repo := api.ds.PlayQueue() + pq, err := repo.RetrieveWithMediaFiles(r.Context(), user.ID) if err != nil && !errors.Is(err, model.ErrNotFound) { return nil, err } @@ -140,8 +140,8 @@ func (api *Router) SavePlayQueue(r *http.Request) (*responses.Subsonic, error) { UpdatedAt: time.Time{}, } - repo := api.ds.PlayQueue(r.Context()) - err := repo.Store(pq) + repo := api.ds.PlayQueue() + err := repo.Store(r.Context(), pq) if err != nil { return nil, err } @@ -151,8 +151,8 @@ func (api *Router) SavePlayQueue(r *http.Request) (*responses.Subsonic, error) { func (api *Router) GetPlayQueueByIndex(r *http.Request) (*responses.Subsonic, error) { user, _ := request.UserFrom(r.Context()) - repo := api.ds.PlayQueue(r.Context()) - pq, err := repo.RetrieveWithMediaFiles(user.ID) + repo := api.ds.PlayQueue() + pq, err := repo.RetrieveWithMediaFiles(r.Context(), user.ID) if err != nil && !errors.Is(err, model.ErrNotFound) { return nil, err } @@ -215,8 +215,8 @@ func (api *Router) SavePlayQueueByIndex(r *http.Request) (*responses.Subsonic, e UpdatedAt: time.Time{}, } - repo := api.ds.PlayQueue(r.Context()) - err = repo.Store(pq) + repo := api.ds.PlayQueue() + err = repo.Store(r.Context(), pq) if err != nil { return nil, err } diff --git a/server/subsonic/bookmarks_test.go b/server/subsonic/bookmarks_test.go index 0fcb81ab9..387ba6ab2 100644 --- a/server/subsonic/bookmarks_test.go +++ b/server/subsonic/bookmarks_test.go @@ -21,7 +21,7 @@ var _ = Describe("Bookmarks", func() { ds = &tests.MockDataStore{} router = New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) ctx = request.WithUser(context.Background(), model.User{ID: "u1", UserName: "u1"}) - mfRepo = ds.MediaFile(ctx).(*tests.MockMediaFileRepo) + mfRepo = ds.MediaFile().(*tests.MockMediaFileRepo) mfRepo.SetData(model.MediaFiles{{ID: "visible"}}) }) diff --git a/server/subsonic/browsing.go b/server/subsonic/browsing.go index ac0edb69f..d74468940 100644 --- a/server/subsonic/browsing.go +++ b/server/subsonic/browsing.go @@ -33,7 +33,7 @@ func (api *Router) GetMusicFolders(r *http.Request) (*responses.Subsonic, error) func (api *Router) getArtist(r *http.Request, libIds []int, ifModifiedSince time.Time) (model.ArtistIndexes, int64, error) { ctx := r.Context() - lastScanStr, err := api.ds.Property(ctx).DefaultGet(consts.LastScanStartTimeKey, "") + lastScanStr, err := api.ds.Property().DefaultGet(ctx, consts.LastScanStartTimeKey, "") if err != nil { log.Error(ctx, "Error retrieving last scan start time", err) return nil, 0, err @@ -45,7 +45,7 @@ func (api *Router) getArtist(r *http.Request, libIds []int, ifModifiedSince time var indexes model.ArtistIndexes if lastScan.After(ifModifiedSince) { - indexes, err = api.ds.Artist(ctx).GetIndex(false, libIds, model.RoleAlbumArtist) + indexes, err = api.ds.Artist().GetIndex(ctx, false, libIds, model.RoleAlbumArtist) if err != nil { log.Error(ctx, "Error retrieving Indexes", err) return nil, 0, err @@ -167,7 +167,7 @@ func (api *Router) GetArtist(r *http.Request) (*responses.Subsonic, error) { id, _ := p.String("id") ctx := r.Context() - artist, err := api.ds.Artist(ctx).Get(id) + artist, err := api.ds.Artist().Get(ctx, id) if errors.Is(err, model.ErrNotFound) { log.Error(ctx, "Requested ArtistID not found ", "id", id) return nil, newError(responses.ErrorDataNotFound, "Artist not found") @@ -191,7 +191,7 @@ func (api *Router) GetAlbum(r *http.Request) (*responses.Subsonic, error) { ctx := r.Context() - album, err := api.ds.Album(ctx).Get(id) + album, err := api.ds.Album().Get(ctx, id) if errors.Is(err, model.ErrNotFound) { log.Error(ctx, "Requested AlbumID not found ", "id", id) return nil, newError(responses.ErrorDataNotFound, "Album not found") @@ -201,7 +201,7 @@ func (api *Router) GetAlbum(r *http.Request) (*responses.Subsonic, error) { return nil, err } - mfs, err := api.ds.MediaFile(ctx).GetAll(filter.SongsByAlbum(id)) + mfs, err := api.ds.MediaFile().GetAll(ctx, filter.SongsByAlbum(id)) if err != nil { log.Error(ctx, "Error retrieving tracks from album", "id", id, "name", album.Name, err) return nil, err @@ -247,7 +247,7 @@ func (api *Router) GetSong(r *http.Request) (*responses.Subsonic, error) { id, _ := p.String("id") ctx := r.Context() - mf, err := api.ds.MediaFile(ctx).Get(id) + mf, err := api.ds.MediaFile().Get(ctx, id) if errors.Is(err, model.ErrNotFound) { log.Error(r, "Requested MediaFileID not found ", "id", id) return nil, newError(responses.ErrorDataNotFound, "Song not found") @@ -264,7 +264,7 @@ func (api *Router) GetSong(r *http.Request) (*responses.Subsonic, error) { func (api *Router) GetGenres(r *http.Request) (*responses.Subsonic, error) { ctx := r.Context() - genres, err := api.ds.Genre(ctx).GetAll(model.QueryOptions{Sort: "song_count, album_count, name desc", Order: "desc"}) + genres, err := api.ds.Genre().GetAll(ctx, model.QueryOptions{Sort: "song_count, album_count, name desc", Order: "desc"}) if err != nil { log.Error(r, err) return nil, err @@ -421,7 +421,7 @@ func (api *Router) buildArtistDirectory(ctx context.Context, artist *model.Artis dir.Starred = artist.StarredAt } - albums, err := api.ds.Album(ctx).GetAll(filter.AlbumsByArtistID(artist.ID)) + albums, err := api.ds.Album().GetAll(ctx, filter.AlbumsByArtistID(artist.ID)) if err != nil { return nil, err } @@ -435,7 +435,7 @@ func (api *Router) buildArtist(r *http.Request, artist *model.Artist) (*response a := &responses.ArtistWithAlbumsID3{} a.ArtistID3 = toArtistID3(r, *artist) - albums, err := api.ds.Album(ctx).GetAll(filter.AlbumsByArtistID(artist.ID)) + albums, err := api.ds.Album().GetAll(ctx, filter.AlbumsByArtistID(artist.ID)) if err != nil { return nil, err } @@ -463,7 +463,7 @@ func (api *Router) buildAlbumDirectory(ctx context.Context, album *model.Album) dir.Starred = album.StarredAt } - mfs, err := api.ds.MediaFile(ctx).GetAll(filter.SongsByAlbum(album.ID)) + mfs, err := api.ds.MediaFile().GetAll(ctx, filter.SongsByAlbum(album.ID)) if err != nil { return nil, err } diff --git a/server/subsonic/browsing_test.go b/server/subsonic/browsing_test.go index d34da2f37..71d758cfd 100644 --- a/server/subsonic/browsing_test.go +++ b/server/subsonic/browsing_test.go @@ -83,14 +83,14 @@ var _ = Describe("Browsing", func() { ctx = contextWithUser(ctx, "user-id", 2, 3) // Setup minimal mock library data for working tests - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(model.Libraries{ {ID: 2, Name: "Test Library 2", Path: "/music/library2"}, {ID: 3, Name: "Test Library 3", Path: "/music/library3"}, }) // Setup mock artist data - mockArtistRepo := ds.Artist(ctx).(*tests.MockArtistRepo) + mockArtistRepo := ds.Artist().(*tests.MockArtistRepo) mockArtistRepo.SetData(model.Artists{ {ID: "1", Name: "Test Artist 1"}, {ID: "2", Name: "Test Artist 2"}, @@ -132,14 +132,14 @@ var _ = Describe("Browsing", func() { ctx = contextWithUser(ctx, "user-id", 1, 2) // Setup minimal mock library data for working tests - mockLibRepo := ds.Library(ctx).(*tests.MockLibraryRepo) + mockLibRepo := ds.Library().(*tests.MockLibraryRepo) mockLibRepo.SetData(model.Libraries{ {ID: 1, Name: "Test Library 1", Path: "/music/library1"}, {ID: 2, Name: "Test Library 2", Path: "/music/library2"}, }) // Setup mock artist data - mockArtistRepo := ds.Artist(ctx).(*tests.MockArtistRepo) + mockArtistRepo := ds.Artist().(*tests.MockArtistRepo) mockArtistRepo.SetData(model.Artists{ {ID: "1", Name: "Test Artist 1"}, {ID: "2", Name: "Test Artist 2"}, diff --git a/server/subsonic/e2e/e2e_suite_test.go b/server/subsonic/e2e/e2e_suite_test.go index 58e877b0d..fdfb562fd 100644 --- a/server/subsonic/e2e/e2e_suite_test.go +++ b/server/subsonic/e2e/e2e_suite_test.go @@ -222,10 +222,10 @@ func createUser(id, username, name string, isAdmin bool) model.User { IsAdmin: isAdmin, NewPassword: "password", } - Expect(ds.User(ctx).Put(&user)).To(Succeed()) - Expect(ds.User(ctx).SetUserLibraries(user.ID, []int{lib.ID})).To(Succeed()) + Expect(ds.User().Put(ctx, &user)).To(Succeed()) + Expect(ds.User().SetUserLibraries(ctx, user.ID, []int{lib.ID})).To(Succeed()) - loadedUser, err := ds.User(ctx).FindByUsername(user.UserName) + loadedUser, err := ds.User().FindByUsername(ctx, user.UserName) Expect(err).ToNot(HaveOccurred()) user.Libraries = loadedUser.Libraries return user diff --git a/server/subsonic/e2e/subsonic_album_lists_test.go b/server/subsonic/e2e/subsonic_album_lists_test.go index 6d32a3c88..e1be3f133 100644 --- a/server/subsonic/e2e/subsonic_album_lists_test.go +++ b/server/subsonic/e2e/subsonic_album_lists_test.go @@ -142,7 +142,7 @@ var _ = Describe("Album List Endpoints", func() { setupTestDB() // Star an album so the starred filter returns results - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -166,7 +166,7 @@ var _ = Describe("Album List Endpoints", func() { setupTestDB() // Rate an album so the highest filter returns results - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Kind of Blue"}, }) Expect(err).ToNot(HaveOccurred()) diff --git a/server/subsonic/e2e/subsonic_artwork_test.go b/server/subsonic/e2e/subsonic_artwork_test.go index 9324ea9e3..9394c8830 100644 --- a/server/subsonic/e2e/subsonic_artwork_test.go +++ b/server/subsonic/e2e/subsonic_artwork_test.go @@ -108,8 +108,8 @@ var _ = Describe("Artwork Serving", Ordered, func() { // harness's MaxOpenConns=1, so wipe the golden content and import this library fresh. wipeScannedContent() artLib := model.Library{Name: "Artwork Library", Path: musicDir} - Expect(ds.Library(ctx).Put(&artLib)).To(Succeed()) - Expect(ds.User(ctx).SetUserLibraries(adminUser.ID, []int{artLib.ID})).To(Succeed()) + Expect(ds.Library().Put(ctx, &artLib)).To(Succeed()) + Expect(ds.User().SetUserLibraries(ctx, adminUser.ID, []int{artLib.ID})).To(Succeed()) s := scanner.New(ctx, ds, events.NoopBroker(), playlists.NewPlaylists(ds, artwork.NewUploader(ds)), metrics.NewNoopInstance()) @@ -148,20 +148,20 @@ var _ = Describe("Artwork Serving", Ordered, func() { It("drains the queue: folder art is acquired, the artless album settles absent", func() { // Enqueues the way the serving paths do, so the drain is driven by a plain queue row. for _, id := range []string{artfulID, artlessID} { - Expect(ds.ArtworkQueue(ctx).EnqueuePreservingBackoff(model.ArtworkQueueItem{ + Expect(ds.ArtworkQueue().EnqueuePreservingBackoff(ctx, model.ArtworkQueueItem{ ItemKind: model.KindAlbumArtwork.Prefix(), ItemID: id, ImageType: model.ImageTypePrimary, Priority: model.ArtworkPriorityBump, })).To(Succeed()) } runWorkerUntil(ctx, worker, func() bool { - found, err := ds.Artwork(ctx).GetItemArtwork(model.KindAlbumArtwork, artfulID, model.ImageTypePrimary) + found, err := ds.Artwork().GetItemArtwork(ctx, model.KindAlbumArtwork, artfulID, model.ImageTypePrimary) if err != nil || found.Hash == "" { return false } - absent, err := ds.Artwork(ctx).GetItemArtwork(model.KindAlbumArtwork, artlessID, model.ImageTypePrimary) + absent, err := ds.Artwork().GetItemArtwork(ctx, model.KindAlbumArtwork, artlessID, model.ImageTypePrimary) return err == nil && absent.Hash == "" }) - ia, err := ds.Artwork(ctx).GetItemArtwork(model.KindAlbumArtwork, artfulID, model.ImageTypePrimary) + ia, err := ds.Artwork().GetItemArtwork(ctx, model.KindAlbumArtwork, artfulID, model.ImageTypePrimary) Expect(err).ToNot(HaveOccurred()) Expect(ia.Source).To(Equal("folder")) artfulHash = ia.Hash @@ -277,7 +277,7 @@ func wipeScannedContent() { func albumIDByName(name string) string { GinkgoHelper() - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Filters: squirrel.Eq{"album.name": name}}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"album.name": name}}) Expect(err).ToNot(HaveOccurred()) Expect(albums).To(HaveLen(1), "expected exactly one album named %q", name) return albums[0].ID diff --git a/server/subsonic/e2e/subsonic_bookmarks_test.go b/server/subsonic/e2e/subsonic_bookmarks_test.go index 726b41743..e2d659b9b 100644 --- a/server/subsonic/e2e/subsonic_bookmarks_test.go +++ b/server/subsonic/e2e/subsonic_bookmarks_test.go @@ -19,7 +19,7 @@ var _ = Describe("Bookmark and PlayQueue Endpoints", Ordered, func() { BeforeAll(func() { // Get a media file ID from the database to use for bookmarks - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1}) + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1}) Expect(err).ToNot(HaveOccurred()) Expect(mfs).ToNot(BeEmpty()) trackID = mfs[0].ID @@ -69,7 +69,7 @@ var _ = Describe("Bookmark and PlayQueue Endpoints", Ordered, func() { BeforeAll(func() { // Get multiple media file IDs from the database - mfs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 3, Sort: "title"}) + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 3, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(len(mfs)).To(BeNumerically(">=", 2)) for _, mf := range mfs { diff --git a/server/subsonic/e2e/subsonic_browsing_test.go b/server/subsonic/e2e/subsonic_browsing_test.go index 992f9e0fb..8aa93ee2e 100644 --- a/server/subsonic/e2e/subsonic_browsing_test.go +++ b/server/subsonic/e2e/subsonic_browsing_test.go @@ -14,7 +14,7 @@ var _ = Describe("Browsing Endpoints", func() { }) getBeatlesId := func() string { - artists, err := ds.Artist(ctx).GetAll(model.QueryOptions{ + artists, err := ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "The Beatles"}, }) Expect(err).ToNot(HaveOccurred()) @@ -105,7 +105,7 @@ var _ = Describe("Browsing Endpoints", func() { }) It("returns an album directory with its tracks as children", func() { - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -159,7 +159,7 @@ var _ = Describe("Browsing Endpoints", func() { }) It("returns artist with a single album", func() { - artists, err := ds.Artist(ctx).GetAll(model.QueryOptions{ + artists, err := ds.Artist().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"name": "Led Zeppelin"}, }) Expect(err).ToNot(HaveOccurred()) @@ -177,7 +177,7 @@ var _ = Describe("Browsing Endpoints", func() { Describe("getAlbum", func() { It("returns album with its tracks", func() { - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -193,7 +193,7 @@ var _ = Describe("Browsing Endpoints", func() { }) It("includes correct track metadata", func() { - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -210,7 +210,7 @@ var _ = Describe("Browsing Endpoints", func() { }) It("returns album with correct artist and year", func() { - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Kind of Blue"}, }) Expect(err).ToNot(HaveOccurred()) @@ -236,7 +236,7 @@ var _ = Describe("Browsing Endpoints", func() { Describe("getSong", func() { It("returns a song by its ID", func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Come Together"}, }) Expect(err).ToNot(HaveOccurred()) @@ -260,7 +260,7 @@ var _ = Describe("Browsing Endpoints", func() { }) It("returns correct metadata for a jazz track", func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "So What"}, }) Expect(err).ToNot(HaveOccurred()) @@ -343,7 +343,7 @@ var _ = Describe("Browsing Endpoints", func() { Describe("getAlbumInfo", func() { It("returns album info for a valid album", func() { - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -359,7 +359,7 @@ var _ = Describe("Browsing Endpoints", func() { Describe("getAlbumInfo2", func() { It("returns album info for a valid album", func() { - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -434,7 +434,7 @@ var _ = Describe("Browsing Endpoints", func() { Describe("getSimilarSongs", func() { It("returns a response for a valid song ID", func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Come Together"}, }) Expect(err).ToNot(HaveOccurred()) @@ -451,7 +451,7 @@ var _ = Describe("Browsing Endpoints", func() { Describe("getSimilarSongs2", func() { It("returns a response for a valid song ID", func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Come Together"}, }) Expect(err).ToNot(HaveOccurred()) diff --git a/server/subsonic/e2e/subsonic_media_annotation_test.go b/server/subsonic/e2e/subsonic_media_annotation_test.go index 74b5238f2..4b90dd143 100644 --- a/server/subsonic/e2e/subsonic_media_annotation_test.go +++ b/server/subsonic/e2e/subsonic_media_annotation_test.go @@ -17,19 +17,19 @@ var _ = Describe("Media Annotation Endpoints", Ordered, func() { BeforeAll(func() { // Look up a song from the scanned data - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "title"}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) songID = songs[0].ID // Look up an album - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "name"}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "name"}) Expect(err).ToNot(HaveOccurred()) Expect(albums).ToNot(BeEmpty()) albumID = albums[0].ID // Look up an artist - artists, err := ds.Artist(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "name"}) + artists, err := ds.Artist().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "name"}) Expect(err).ToNot(HaveOccurred()) Expect(artists).ToNot(BeEmpty()) artistID = artists[0].ID @@ -97,12 +97,12 @@ var _ = Describe("Media Annotation Endpoints", Ordered, func() { var songID, albumID string BeforeAll(func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "title"}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) songID = songs[0].ID - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "name"}) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "name"}) Expect(err).ToNot(HaveOccurred()) Expect(albums).ToNot(BeEmpty()) albumID = albums[0].ID @@ -141,7 +141,7 @@ var _ = Describe("Media Annotation Endpoints", Ordered, func() { Describe("Scrobble", func() { It("submits a scrobble for a song", func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "title"}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) @@ -163,7 +163,7 @@ var _ = Describe("Media Annotation Endpoints", Ordered, func() { var songID string BeforeAll(func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "title"}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) songID = songs[0].ID diff --git a/server/subsonic/e2e/subsonic_media_retrieval_test.go b/server/subsonic/e2e/subsonic_media_retrieval_test.go index 268d93b82..b0ef4fd72 100644 --- a/server/subsonic/e2e/subsonic_media_retrieval_test.go +++ b/server/subsonic/e2e/subsonic_media_retrieval_test.go @@ -21,7 +21,7 @@ var _ = Describe("Media Retrieval Endpoints", Ordered, func() { BeforeAll(func() { // All test tracks are mp3 at 320kbps - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "title"}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) trackID = songs[0].ID @@ -110,7 +110,7 @@ var _ = Describe("Media Retrieval Endpoints", Ordered, func() { BeforeAll(func() { // All test tracks are mp3 at 320kbps - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Max: 1, Sort: "title"}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Max: 1, Sort: "title"}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) trackID = songs[0].ID diff --git a/server/subsonic/e2e/subsonic_multilibrary_test.go b/server/subsonic/e2e/subsonic_multilibrary_test.go index 18e8c6391..8158be1de 100644 --- a/server/subsonic/e2e/subsonic_multilibrary_test.go +++ b/server/subsonic/e2e/subsonic_multilibrary_test.go @@ -44,10 +44,10 @@ var _ = Describe("Multi-Library Support", Ordered, func() { // Create the second library in the DB (Put auto-assigns admin users) lib2 = model.Library{ID: 2, Name: "Classical Library", Path: "fake2:///classical"} - Expect(ds.Library(ctx).Put(&lib2)).To(Succeed()) + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) // Reload admin user to get both libraries in the Libraries field - loadedAdmin, err := ds.User(ctx).FindByUsername(adminUser.UserName) + loadedAdmin, err := ds.User().FindByUsername(ctx, adminUser.UserName) Expect(err).ToNot(HaveOccurred()) adminWithLibs = *loadedAdmin @@ -65,10 +65,10 @@ var _ = Describe("Multi-Library Support", Ordered, func() { IsAdmin: false, NewPassword: "password", } - Expect(ds.User(ctx).Put(&userLib1Only)).To(Succeed()) - Expect(ds.User(ctx).SetUserLibraries(userLib1Only.ID, []int{lib.ID})).To(Succeed()) + Expect(ds.User().Put(ctx, &userLib1Only)).To(Succeed()) + Expect(ds.User().SetUserLibraries(ctx, userLib1Only.ID, []int{lib.ID})).To(Succeed()) - loadedUser, err := ds.User(ctx).FindByUsername(userLib1Only.UserName) + loadedUser, err := ds.User().FindByUsername(ctx, userLib1Only.UserName) Expect(err).ToNot(HaveOccurred()) userLib1Only.Libraries = loadedUser.Libraries }) @@ -181,7 +181,7 @@ var _ = Describe("Multi-Library Support", Ordered, func() { BeforeAll(func() { // Look up one song from each library - lib1Songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + lib1Songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"media_file.library_id": lib.ID}, Max: 1, Sort: "title", }) @@ -189,7 +189,7 @@ var _ = Describe("Multi-Library Support", Ordered, func() { Expect(lib1Songs).ToNot(BeEmpty()) lib1SongID = lib1Songs[0].ID - lib2Songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + lib2Songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"media_file.library_id": lib2.ID}, Max: 1, Sort: "title", }) @@ -248,7 +248,7 @@ var _ = Describe("Multi-Library Support", Ordered, func() { var lib2AlbumID string BeforeAll(func() { - lib2Albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + lib2Albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.library_id": lib2.ID}, }) Expect(err).ToNot(HaveOccurred()) diff --git a/server/subsonic/e2e/subsonic_playlists_test.go b/server/subsonic/e2e/subsonic_playlists_test.go index 467535df7..7a6d7df5d 100644 --- a/server/subsonic/e2e/subsonic_playlists_test.go +++ b/server/subsonic/e2e/subsonic_playlists_test.go @@ -19,7 +19,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { setupTestDB() // Look up song IDs from scanned data for playlist operations - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Sort: "title", Max: 6}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Sort: "title", Max: 6}) Expect(err).ToNot(HaveOccurred()) Expect(len(songs)).To(BeNumerically(">=", 5)) for _, s := range songs { @@ -244,7 +244,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { BeforeAll(func() { setupTestDB() - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Sort: "title", Max: 6}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Sort: "title", Max: 6}) Expect(err).ToNot(HaveOccurred()) Expect(len(songs)).To(BeNumerically(">=", 3)) for _, s := range songs { @@ -438,7 +438,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { setupTestDB() // Look up a song ID for mutation tests - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Sort: "title", Max: 1}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Sort: "title", Max: 1}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) songID = songs[0].ID @@ -450,7 +450,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { Public: false, Rules: &criteria.Criteria{Expression: criteria.Contains{"title": ""}}, } - Expect(ds.Playlist(ctx).Put(smartPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, smartPls)).To(Succeed()) smartPlaylistID = smartPls.ID }) @@ -525,7 +525,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { BeforeAll(func() { setupTestDB() - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{Sort: "title", Max: 1}) + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Sort: "title", Max: 1}) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) songID = songs[0].ID @@ -543,7 +543,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.Is{"loved": true}}}, } - Expect(ds.Playlist(ctx).Put(boolPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, boolPls)).To(Succeed()) boolPlaylistID = boolPls.ID // Create smart playlist with string "true" @@ -552,7 +552,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.Is{"loved": "true"}}}, } - Expect(ds.Playlist(ctx).Put(stringPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, stringPls)).To(Succeed()) stringPlaylistID = stringPls.ID // Create smart playlist with string "true" in nested any group (exact issue #4826 scenario) @@ -565,7 +565,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { }, }}, } - Expect(ds.Playlist(ctx).Put(nestedPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, nestedPls)).To(Succeed()) nestedPlaylistID = nestedPls.ID }) @@ -607,7 +607,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.IsPresent{"genre": "true"}}}, } - Expect(ds.Playlist(ctx).Put(pls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, pls)).To(Succeed()) resp := doReq("getPlaylist", "id", pls.ID) Expect(resp.Status).To(Equal(responses.StatusOK)) @@ -620,7 +620,7 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.IsMissing{"genre": "true"}}}, } - Expect(ds.Playlist(ctx).Put(pls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, pls)).To(Succeed()) resp := doReq("getPlaylist", "id", pls.ID) Expect(resp.Status).To(Equal(responses.StatusOK)) @@ -633,14 +633,14 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.IsMissing{"genre": true}}}, } - Expect(ds.Playlist(ctx).Put(boolPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, boolPls)).To(Succeed()) stringPls := &model.Playlist{ Name: "Genre Missing String2", OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.IsMissing{"genre": "true"}}}, } - Expect(ds.Playlist(ctx).Put(stringPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, stringPls)).To(Succeed()) boolResp := doReq("getPlaylist", "id", boolPls.ID) stringResp := doReq("getPlaylist", "id", stringPls.ID) @@ -654,19 +654,19 @@ var _ = Describe("Playlist Endpoints", Ordered, func() { OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.Contains{"title": ""}}, } - Expect(ds.Playlist(ctx).Put(allPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, allPls)).To(Succeed()) missingPls := &model.Playlist{ Name: "Missing " + fieldName, OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.IsMissing{fieldName: true}}}, } - Expect(ds.Playlist(ctx).Put(missingPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, missingPls)).To(Succeed()) presentPls := &model.Playlist{ Name: "Present " + fieldName, OwnerID: adminUser.ID, Rules: &criteria.Criteria{Expression: criteria.All{criteria.IsPresent{fieldName: true}}}, } - Expect(ds.Playlist(ctx).Put(presentPls)).To(Succeed()) + Expect(ds.Playlist().Put(ctx, presentPls)).To(Succeed()) allResp := doReq("getPlaylist", "id", allPls.ID) missingResp := doReq("getPlaylist", "id", missingPls.ID) diff --git a/server/subsonic/e2e/subsonic_sharing_test.go b/server/subsonic/e2e/subsonic_sharing_test.go index 1ae68dc1a..cf3e55d74 100644 --- a/server/subsonic/e2e/subsonic_sharing_test.go +++ b/server/subsonic/e2e/subsonic_sharing_test.go @@ -18,14 +18,14 @@ var _ = Describe("Sharing Endpoints", Ordered, func() { conf.Server.EnableSharing = true setupTestDB() - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) Expect(albums).ToNot(BeEmpty()) albumID = albums[0].ID - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Come Together"}, }) Expect(err).ToNot(HaveOccurred()) @@ -139,7 +139,7 @@ var _ = Describe("Sharing Cross-User Isolation", Ordered, func() { userA = createUser("share-user-a", "share-user-a", "Share User A", false) userB = createUser("share-user-b", "share-user-b", "Share User B", false) - albums, err := ds.Album(ctx).GetAll(model.QueryOptions{ + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"album.name": "Abbey Road"}, }) Expect(err).ToNot(HaveOccurred()) @@ -221,7 +221,7 @@ var _ = Describe("Sharing Downloadable Default", func() { resp := doReq("createShare", append([]string{"id", albumID}, params...)...) Expect(resp.Status).To(Equal(responses.StatusOK)) Expect(resp.Shares.Share).To(HaveLen(1)) - share, err := ds.Share(ctx).Get(resp.Shares.Share[0].ID) + share, err := ds.Share().Get(ctx, resp.Shares.Share[0].ID) Expect(err).ToNot(HaveOccurred()) return share } @@ -248,7 +248,7 @@ var _ = Describe("Sharing Downloadable Default", func() { resp := doReq("updateShare", "id", share.ID, "description", "Updated") Expect(resp.Status).To(Equal(responses.StatusOK)) - updated, err := ds.Share(ctx).Get(share.ID) + updated, err := ds.Share().Get(ctx, share.ID) Expect(err).ToNot(HaveOccurred()) Expect(updated.Description).To(Equal("Updated")) Expect(updated.Downloadable).To(BeTrue()) @@ -261,7 +261,7 @@ var _ = Describe("Sharing Downloadable Default", func() { resp := doReq("updateShare", "id", share.ID, "downloadable", "false") Expect(resp.Status).To(Equal(responses.StatusOK)) - updated, err := ds.Share(ctx).Get(share.ID) + updated, err := ds.Share().Get(ctx, share.ID) Expect(err).ToNot(HaveOccurred()) Expect(updated.Downloadable).To(BeFalse()) Expect(updated.Description).To(Equal("Keep me")) @@ -273,7 +273,7 @@ var _ = Describe("Sharing Downloadable Default", func() { resp := doReq("updateShare", "id", share.ID, "description", "") Expect(resp.Status).To(Equal(responses.StatusOK)) - updated, err := ds.Share(ctx).Get(share.ID) + updated, err := ds.Share().Get(ctx, share.ID) Expect(err).ToNot(HaveOccurred()) Expect(updated.Description).To(BeEmpty()) }) diff --git a/server/subsonic/e2e/subsonic_sonic_similarity_test.go b/server/subsonic/e2e/subsonic_sonic_similarity_test.go index c0cb1d359..80fa658b7 100644 --- a/server/subsonic/e2e/subsonic_sonic_similarity_test.go +++ b/server/subsonic/e2e/subsonic_sonic_similarity_test.go @@ -100,14 +100,14 @@ var _ = Describe("Sonic Similarity Endpoints", func() { ) BeforeEach(func() { - songs, err := ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Come Together"}, }) Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) comeTogether = songs[0] - songs, err = ds.MediaFile(ctx).GetAll(model.QueryOptions{ + songs, err = ds.MediaFile().GetAll(ctx, model.QueryOptions{ Filters: squirrel.Eq{"title": "Something"}, }) Expect(err).ToNot(HaveOccurred()) diff --git a/server/subsonic/e2e/subsonic_stream_test.go b/server/subsonic/e2e/subsonic_stream_test.go index 281524636..81998760f 100644 --- a/server/subsonic/e2e/subsonic_stream_test.go +++ b/server/subsonic/e2e/subsonic_stream_test.go @@ -21,7 +21,7 @@ var _ = Describe("stream.view (legacy streaming)", Ordered, func() { BeforeAll(func() { setupTestDB() - songs, err := ds.MediaFile(ctx).GetAll() + songs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) byTitle := map[string]string{} for _, s := range songs { diff --git a/server/subsonic/e2e/subsonic_transcode_test.go b/server/subsonic/e2e/subsonic_transcode_test.go index afe7d52ca..313066938 100644 --- a/server/subsonic/e2e/subsonic_transcode_test.go +++ b/server/subsonic/e2e/subsonic_transcode_test.go @@ -127,7 +127,7 @@ var _ = Describe("Transcode Endpoints", Ordered, func() { BeforeAll(func() { setupTestDB() - songs, err := ds.MediaFile(ctx).GetAll() + songs, err := ds.MediaFile().GetAll(ctx) Expect(err).ToNot(HaveOccurred()) byTitle := map[string]string{} for _, s := range songs { @@ -153,29 +153,29 @@ var _ = Describe("Transcode Endpoints", Ordered, func() { // It makes a dummy request to register the player, then updates it via the repository. setPlayerMaxBitRate := func(maxBitRate int) { doReq("ping") - player, err := ds.Player(ctx).FindMatch(adminUser.ID, "test-client", "") + player, err := ds.Player().FindMatch(ctx, adminUser.ID, "test-client", "") Expect(err).ToNot(HaveOccurred()) player.MaxBitRate = maxBitRate - Expect(ds.Player(ctx).Put(player)).To(Succeed()) + Expect(ds.Player().Put(ctx, player)).To(Succeed()) } setPlayerForcedFormat := func(format string) { doReq("ping") - player, err := ds.Player(ctx).FindMatch(adminUser.ID, "test-client", "") + player, err := ds.Player().FindMatch(ctx, adminUser.ID, "test-client", "") Expect(err).ToNot(HaveOccurred()) - trc, err := ds.Transcoding(ctx).FindByFormat(format) + trc, err := ds.Transcoding().FindByFormat(ctx, format) Expect(err).ToNot(HaveOccurred()) player.TranscodingId = trc.ID - Expect(ds.Player(ctx).Put(player)).To(Succeed()) + Expect(ds.Player().Put(ctx, player)).To(Succeed()) } AfterEach(func() { // Reset player MaxBitRate to 0 after each test - player, err := ds.Player(ctx).FindMatch(adminUser.ID, "test-client", "") + player, err := ds.Player().FindMatch(ctx, adminUser.ID, "test-client", "") if err == nil { player.MaxBitRate = 0 player.TranscodingId = "" - _ = ds.Player(ctx).Put(player) + _ = ds.Player().Put(ctx, player) } }) @@ -595,13 +595,13 @@ var _ = Describe("Transcode Endpoints", Ordered, func() { Expect(token).ToNot(BeEmpty()) // Save original UpdatedAt and restore after test - mf, err := ds.MediaFile(ctx).Get(mp3TrackID) + mf, err := ds.MediaFile().Get(ctx, mp3TrackID) Expect(err).ToNot(HaveOccurred()) originalUpdatedAt := mf.UpdatedAt // Update the media file's UpdatedAt to simulate a change after token issuance mf.UpdatedAt = time.Now().Add(time.Hour) - Expect(ds.MediaFile(ctx).Put(mf)).To(Succeed()) + Expect(ds.MediaFile().Put(ctx, mf)).To(Succeed()) // Attempt to stream with the now-stale token w := doRawReq("getTranscodeStream", "mediaId", mp3TrackID, "mediaType", "song", "transcodeParams", token) @@ -609,7 +609,7 @@ var _ = Describe("Transcode Endpoints", Ordered, func() { // Restore original UpdatedAt mf.UpdatedAt = originalUpdatedAt - Expect(ds.MediaFile(ctx).Put(mf)).To(Succeed()) + Expect(ds.MediaFile().Put(ctx, mf)).To(Succeed()) }) It("returns 500 when stream creation fails", func() { diff --git a/server/subsonic/library_scanning.go b/server/subsonic/library_scanning.go index 9630425d2..89e5cbc9b 100644 --- a/server/subsonic/library_scanning.go +++ b/server/subsonic/library_scanning.go @@ -53,7 +53,7 @@ func (api *Router) StartScan(r *http.Request) (*responses.Subsonic, error) { } // Validate all libraries in targets exist and user has access to them - userLibraries, err := api.ds.User(ctx).GetUserLibraries(loggedUser.ID) + userLibraries, err := api.ds.User().GetUserLibraries(ctx, loggedUser.ID) if err != nil { return nil, newError(responses.ErrorGeneric, "Internal error") } @@ -67,7 +67,7 @@ func (api *Router) StartScan(r *http.Request) (*responses.Subsonic, error) { // Special case: if single library with empty path and it's the only library in DB, call ScanAll if len(targets) == 1 && targets[0].FolderPath == "" { - allLibs, err := api.ds.Library(ctx).GetAll() + allLibs, err := api.ds.Library().GetAll(ctx) if err != nil { return nil, newError(responses.ErrorGeneric, "Internal error") } diff --git a/server/subsonic/library_scanning_test.go b/server/subsonic/library_scanning_test.go index 771fc3352..e2b8827f4 100644 --- a/server/subsonic/library_scanning_test.go +++ b/server/subsonic/library_scanning_test.go @@ -80,7 +80,7 @@ var _ = Describe("LibraryScanning", func() { It("triggers a selective scan with single target parameter", func() { // Setup mocks mockUserRepo := tests.CreateMockUserRepo() - _ = mockUserRepo.SetUserLibraries("admin-id", []int{1, 2}) + _ = mockUserRepo.SetUserLibraries(GinkgoT().Context(), "admin-id", []int{1, 2}) mockDS := &tests.MockDataStore{MockedUser: mockUserRepo} api.ds = mockDS @@ -116,7 +116,7 @@ var _ = Describe("LibraryScanning", func() { It("triggers a selective scan with multiple target parameters", func() { // Setup mocks mockUserRepo := tests.CreateMockUserRepo() - _ = mockUserRepo.SetUserLibraries("admin-id", []int{1, 2}) + _ = mockUserRepo.SetUserLibraries(GinkgoT().Context(), "admin-id", []int{1, 2}) mockDS := &tests.MockDataStore{MockedUser: mockUserRepo} api.ds = mockDS @@ -154,7 +154,7 @@ var _ = Describe("LibraryScanning", func() { It("triggers a selective full scan with target and fullScan parameters", func() { // Setup mocks mockUserRepo := tests.CreateMockUserRepo() - _ = mockUserRepo.SetUserLibraries("admin-id", []int{1}) + _ = mockUserRepo.SetUserLibraries(GinkgoT().Context(), "admin-id", []int{1}) mockDS := &tests.MockDataStore{MockedUser: mockUserRepo} api.ds = mockDS @@ -235,7 +235,7 @@ var _ = Describe("LibraryScanning", func() { It("returns error when library does not exist", func() { // Setup mocks - user has access to library 1 and 2 only mockUserRepo := tests.CreateMockUserRepo() - _ = mockUserRepo.SetUserLibraries("admin-id", []int{1, 2}) + _ = mockUserRepo.SetUserLibraries(GinkgoT().Context(), "admin-id", []int{1, 2}) mockDS := &tests.MockDataStore{MockedUser: mockUserRepo} api.ds = mockDS @@ -264,7 +264,7 @@ var _ = Describe("LibraryScanning", func() { It("calls ScanAll when single library with empty path and only one library exists", func() { // Setup mocks - single library in DB mockUserRepo := tests.CreateMockUserRepo() - _ = mockUserRepo.SetUserLibraries("admin-id", []int{1}) + _ = mockUserRepo.SetUserLibraries(GinkgoT().Context(), "admin-id", []int{1}) mockLibraryRepo := &tests.MockLibraryRepo{} mockLibraryRepo.SetData(model.Libraries{ {ID: 1, Name: "Music Library", Path: "/music"}, @@ -302,7 +302,7 @@ var _ = Describe("LibraryScanning", func() { It("calls ScanFolders when single library with empty path but multiple libraries exist", func() { // Setup mocks - multiple libraries in DB mockUserRepo := tests.CreateMockUserRepo() - _ = mockUserRepo.SetUserLibraries("admin-id", []int{1, 2}) + _ = mockUserRepo.SetUserLibraries(GinkgoT().Context(), "admin-id", []int{1, 2}) mockLibraryRepo := &tests.MockLibraryRepo{} mockLibraryRepo.SetData(model.Libraries{ {ID: 1, Name: "Music Library", Path: "/music"}, diff --git a/server/subsonic/media_annotation.go b/server/subsonic/media_annotation.go index a43162cf4..d6d160570 100644 --- a/server/subsonic/media_annotation.go +++ b/server/subsonic/media_annotation.go @@ -48,19 +48,19 @@ func (api *Router) setRating(ctx context.Context, id string, rating int) error { } switch entity.(type) { case *model.Artist: - repo = api.ds.Artist(ctx) + repo = api.ds.Artist() resource = "artist" case *model.Album: - repo = api.ds.Album(ctx) + repo = api.ds.Album() resource = "album" case *model.Playlist: - repo = api.ds.Playlist(ctx) + repo = api.ds.Playlist() resource = "playlist" default: - repo = api.ds.MediaFile(ctx) + repo = api.ds.MediaFile() resource = "song" } - err = repo.SetRating(rating, id) + err = repo.SetRating(ctx, rating, id) if err != nil { return err } @@ -129,19 +129,19 @@ func (api *Router) setStar(ctx context.Context, star bool, ids ...string) error } switch entity.(type) { case *model.Artist: - repo = tx.Artist(ctx) + repo = tx.Artist() resource = "artist" case *model.Album: - repo = tx.Album(ctx) + repo = tx.Album() resource = "album" case *model.Playlist: - repo = tx.Playlist(ctx) + repo = tx.Playlist() resource = "playlist" default: - repo = tx.MediaFile(ctx) + repo = tx.MediaFile() resource = "song" } - if err := repo.SetStar(star, id); err != nil { + if err := repo.SetStar(ctx, star, id); err != nil { return err } event = event.With(resource, id) @@ -210,7 +210,7 @@ func (api *Router) scrobblerSubmit(ctx context.Context, ids []string, times []ti } func (api *Router) scrobblerNowPlaying(ctx context.Context, trackId string, position int) error { - mf, err := api.ds.MediaFile(ctx).Get(trackId) + mf, err := api.ds.MediaFile().Get(ctx, trackId) if err != nil { return err } diff --git a/server/subsonic/media_annotation_test.go b/server/subsonic/media_annotation_test.go index 1b16dfc68..9948929e3 100644 --- a/server/subsonic/media_annotation_test.go +++ b/server/subsonic/media_annotation_test.go @@ -77,7 +77,7 @@ var _ = Describe("MediaAnnotationController", func() { Context("submission=false", func() { var req *http.Request BeforeEach(func() { - _ = ds.MediaFile(ctx).Put(&model.MediaFile{ID: "12"}) + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "12"}) ctx = request.WithPlayer(ctx, model.Player{ID: "player-1"}) req = newGetRequest("id=12", "submission=false") req = req.WithContext(ctx) diff --git a/server/subsonic/media_retrieval.go b/server/subsonic/media_retrieval.go index 8a5152a9d..f5e2ceb25 100644 --- a/server/subsonic/media_retrieval.go +++ b/server/subsonic/media_retrieval.go @@ -29,7 +29,7 @@ func (api *Router) GetAvatar(w http.ResponseWriter, r *http.Request) (*responses return nil, err } ctx := r.Context() - u, err := api.ds.User(ctx).FindByUsername(username) + u, err := api.ds.User().FindByUsername(ctx, username) if err != nil { return nil, err } @@ -128,7 +128,7 @@ func (api *Router) GetLyricsBySongId(r *http.Request) (*responses.Subsonic, erro return nil, err } - mediaFile, err := api.ds.MediaFile(r.Context()).Get(id) + mediaFile, err := api.ds.MediaFile().Get(r.Context(), id) if err != nil { return nil, err } diff --git a/server/subsonic/media_retrieval_test.go b/server/subsonic/media_retrieval_test.go index 7610c866a..ad193518d 100644 --- a/server/subsonic/media_retrieval_test.go +++ b/server/subsonic/media_retrieval_test.go @@ -34,7 +34,7 @@ var _ = Describe("MediaRetrievalController", func() { albumRepo := &tests.MockAlbumRepo{} albumRepo.SetData(model.Albums{{ID: "34"}}) // the id the specs request, made accessible radioRepo := tests.CreateMockedRadioRepo() - Expect(radioRepo.Put(&model.Radio{ID: "rd1", Name: "Radio"})).To(Succeed()) + Expect(radioRepo.Put(GinkgoT().Context(), &model.Radio{ID: "rd1", Name: "Radio"})).To(Succeed()) ds = &tests.MockDataStore{ MockedMediaFile: mockRepo, MockedAlbum: albumRepo, @@ -293,8 +293,8 @@ type mockedMediaFile struct { tests.MockMediaFileRepo } -func (m *mockedMediaFile) GetAll(opts ...model.QueryOptions) (model.MediaFiles, error) { - data, err := m.MockMediaFileRepo.GetAll(opts...) +func (m *mockedMediaFile) GetAll(ctx context.Context, opts ...model.QueryOptions) (model.MediaFiles, error) { + data, err := m.MockMediaFileRepo.GetAll(ctx, opts...) if err != nil { return nil, err } diff --git a/server/subsonic/middlewares.go b/server/subsonic/middlewares.go index 35e13eaa5..6617661a9 100644 --- a/server/subsonic/middlewares.go +++ b/server/subsonic/middlewares.go @@ -109,7 +109,7 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { username, isInternalAuth := fromInternalOrProxyAuth(r) if username != "" { authType := If(isInternalAuth, "internal", "reverse-proxy") - usr, err = ds.User(ctx).FindByUsername(username) + usr, err = ds.User().FindByUsername(ctx, username) if errors.Is(err, context.Canceled) { log.Debug(ctx, "API: Request canceled when authenticating", "auth", authType, "username", username, "remoteAddr", r.RemoteAddr, err) return @@ -139,7 +139,7 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { return } - usr, err = ds.User(ctx).FindByUsernameWithPassword(username) + usr, err = ds.User().FindByUsernameWithPassword(ctx, username) if err == nil { err = validateCredentials(usr, pass, token, salt, jwt) } diff --git a/server/subsonic/middlewares_test.go b/server/subsonic/middlewares_test.go index 62de09e76..0879ee540 100644 --- a/server/subsonic/middlewares_test.go +++ b/server/subsonic/middlewares_test.go @@ -42,11 +42,13 @@ func newPostRequest(queryParam string, formFields ...string) *http.Request { } var _ = Describe("Middlewares", func() { + var ctx context.Context var next *mockHandler var w *httptest.ResponseRecorder var ds model.DataStore BeforeEach(func() { + ctx = GinkgoT().Context() next = &mockHandler{} w = httptest.NewRecorder() ds = &tests.MockDataStore{} @@ -147,8 +149,8 @@ var _ = Describe("Middlewares", func() { Describe("Authenticate", func() { BeforeEach(func() { - ur := ds.User(context.TODO()) - _ = ur.Put(&model.User{ + ur := ds.User() + _ = ur.Put(ctx, &model.User{ UserName: "admin", NewPassword: "wordpass", }) @@ -344,7 +346,7 @@ var _ = Describe("Middlewares", func() { It("counts attempts against unknown usernames", func() { failTimes(3, "u=newuser", "p=secret") - _ = ds.User(context.TODO()).Put(&model.User{UserName: "newuser", NewPassword: "secret"}) + _ = ds.User().Put(ctx, &model.User{UserName: "newuser", NewPassword: "secret"}) serve(newGetRequest("u=newuser", "p=secret")) Expect(next.called).To(BeFalse()) @@ -365,7 +367,7 @@ var _ = Describe("Middlewares", func() { }) It("does not count server errors", func() { - userRepo := ds.User(context.TODO()).(*tests.MockedUserRepo) + userRepo := ds.User().(*tests.MockedUserRepo) userRepo.Error = errors.New("db down") failTimes(5, "u=admin", "p=wordpass") userRepo.Error = nil @@ -375,7 +377,7 @@ var _ = Describe("Middlewares", func() { }) It("does not block other usernames from the same IP", func() { - _ = ds.User(context.TODO()).Put(&model.User{UserName: "other", NewPassword: "otherpass"}) + _ = ds.User().Put(ctx, &model.User{UserName: "other", NewPassword: "otherpass"}) failTimes(3, "u=admin", "p=WRONG") serve(newGetRequest("u=other", "p=otherpass")) @@ -422,7 +424,7 @@ var _ = Describe("Middlewares", func() { conf.Server.AuthRequestLimit = 5 conf.Server.AuthWindowLength = time.Minute gate = &gatedUserRepo{ - UserRepository: ds.User(context.TODO()), + UserRepository: ds.User(), entered: make(chan struct{}, 64), proceed: make(chan struct{}), } @@ -583,14 +585,14 @@ var _ = Describe("Middlewares", func() { var usr *model.User BeforeEach(func() { - ur := ds.User(context.TODO()) - _ = ur.Put(&model.User{ + ur := ds.User() + _ = ur.Put(ctx, &model.User{ UserName: "admin", NewPassword: "wordpass", }) var err error - usr, err = ur.FindByUsernameWithPassword("admin") + usr, err = ur.FindByUsernameWithPassword(ctx, "admin") if err != nil { panic(err) } @@ -728,7 +730,7 @@ type gatedDataStore struct { users model.UserRepository } -func (g *gatedDataStore) User(context.Context) model.UserRepository { return g.users } +func (g *gatedDataStore) User() model.UserRepository { return g.users } type gatedUserRepo struct { model.UserRepository @@ -737,11 +739,11 @@ type gatedUserRepo struct { lookups atomic.Int32 } -func (g *gatedUserRepo) FindByUsernameWithPassword(username string) (*model.User, error) { +func (g *gatedUserRepo) FindByUsernameWithPassword(ctx context.Context, username string) (*model.User, error) { g.lookups.Add(1) g.entered <- struct{}{} <-g.proceed - return g.UserRepository.FindByUsernameWithPassword(username) + return g.UserRepository.FindByUsernameWithPassword(ctx, username) } type countingHandler struct{ calls atomic.Int32 } diff --git a/server/subsonic/radio.go b/server/subsonic/radio.go index 1fb266f1d..2c4033467 100644 --- a/server/subsonic/radio.go +++ b/server/subsonic/radio.go @@ -32,7 +32,7 @@ func (api *Router) CreateInternetRadio(r *http.Request) (*responses.Subsonic, er Name: name, } - err = api.ds.Radio(ctx).Put(radio) + err = api.ds.Radio().Put(ctx, radio) if err != nil { return nil, err } @@ -47,7 +47,7 @@ func (api *Router) DeleteInternetRadio(r *http.Request) (*responses.Subsonic, er return nil, err } - err = api.ds.Radio(r.Context()).Delete(id) + err = api.ds.Radio().Delete(r.Context(), id) if err != nil { return nil, err } @@ -56,7 +56,7 @@ func (api *Router) DeleteInternetRadio(r *http.Request) (*responses.Subsonic, er func (api *Router) GetInternetRadios(r *http.Request) (*responses.Subsonic, error) { ctx := r.Context() - radios, err := api.ds.Radio(ctx).GetAll(model.QueryOptions{Sort: "name"}) + radios, err := api.ds.Radio().GetAll(ctx, model.QueryOptions{Sort: "name"}) if err != nil { return nil, err } @@ -119,7 +119,7 @@ func (api *Router) UpdateInternetRadio(r *http.Request) (*responses.Subsonic, er Name: name, } - err = api.ds.Radio(ctx).Put(radio, "StreamUrl", "HomePageUrl", "Name") + err = api.ds.Radio().Put(ctx, radio, "StreamUrl", "HomePageUrl", "Name") if err != nil { return nil, err } diff --git a/server/subsonic/searching.go b/server/subsonic/searching.go index 35233a98f..fb370fc9d 100644 --- a/server/subsonic/searching.go +++ b/server/subsonic/searching.go @@ -42,7 +42,7 @@ func (api *Router) getSearchParams(r *http.Request) (*searchParams, error) { return sp, nil } -type searchFunc[T any] func(q string, options ...model.QueryOptions) (T, error) +type searchFunc[T any] func(ctx context.Context, q string, options ...model.QueryOptions) (T, error) func callSearch[T any](ctx context.Context, s searchFunc[T], q string, options model.QueryOptions, result *T) func() error { return func() error { @@ -52,7 +52,7 @@ func callSearch[T any](ctx context.Context, s searchFunc[T], q string, options m typ := strings.TrimPrefix(reflect.TypeOf(*result).String(), "model.") var err error start := time.Now() - *result, err = s(q, options) + *result, err = s(ctx, q, options) if err != nil { log.Error(ctx, "Error searching "+typ, "query", q, "elapsed", time.Since(start), err) } else { @@ -79,9 +79,9 @@ func (api *Router) searchAll(ctx context.Context, sp *searchParams, musicFolderI // Run searches in parallel g, ctx := errgroup.WithContext(ctx) - g.Go(callSearch(ctx, api.ds.MediaFile(ctx).Search, q, songOpts, &mediaFiles)) - g.Go(callSearch(ctx, api.ds.Album(ctx).Search, q, albumOpts, &albums)) - g.Go(callSearch(ctx, api.ds.Artist(ctx).Search, q, artistOpts, &artists)) + g.Go(callSearch(ctx, api.ds.MediaFile().Search, q, songOpts, &mediaFiles)) + g.Go(callSearch(ctx, api.ds.Album().Search, q, albumOpts, &albums)) + g.Go(callSearch(ctx, api.ds.Artist().Search, q, artistOpts, &artists)) err := g.Wait() if err == nil { log.Debug(ctx, fmt.Sprintf("Search resulted in %d songs, %d albums and %d artists", diff --git a/server/subsonic/searching_test.go b/server/subsonic/searching_test.go index d31a50cfa..177fa2133 100644 --- a/server/subsonic/searching_test.go +++ b/server/subsonic/searching_test.go @@ -26,9 +26,9 @@ var _ = Describe("Search", func() { router = New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) // Get references to the mock repositories so we can inspect their Options - mockAlbumRepo = ds.Album(nil).(*tests.MockAlbumRepo) - mockArtistRepo = ds.Artist(nil).(*tests.MockArtistRepo) - mockMediaFileRepo = ds.MediaFile(nil).(*tests.MockMediaFileRepo) + mockAlbumRepo = ds.Album().(*tests.MockAlbumRepo) + mockArtistRepo = ds.Artist().(*tests.MockArtistRepo) + mockMediaFileRepo = ds.MediaFile().(*tests.MockMediaFileRepo) }) Context("musicFolderId parameter", func() { diff --git a/server/subsonic/sharing.go b/server/subsonic/sharing.go index c4b735832..e5f4f2a54 100644 --- a/server/subsonic/sharing.go +++ b/server/subsonic/sharing.go @@ -6,7 +6,6 @@ import ( "strings" "time" - "github.com/deluan/rest" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/server/public" @@ -16,8 +15,8 @@ import ( ) func (api *Router) GetShares(r *http.Request) (*responses.Subsonic, error) { - repo := api.share.NewRepository(r.Context()).(model.ShareRepository) - shares, err := repo.GetAll(model.QueryOptions{Sort: "created_at desc"}) + repo := api.share.Repository() + shares, err := repo.GetAll(r.Context(), model.QueryOptions{Sort: "created_at desc"}) if err != nil { return nil, err } @@ -60,7 +59,7 @@ func (api *Router) CreateShare(r *http.Request) (*responses.Subsonic, error) { } description, _ := p.String("description") - repo := api.share.NewRepository(r.Context()) + repo := api.share.Repository() share := &model.Share{ Description: description, Downloadable: p.BoolOr("downloadable", conf.Server.DefaultDownloadableShare && conf.Server.EnableDownloads), @@ -68,12 +67,12 @@ func (api *Router) CreateShare(r *http.Request) (*responses.Subsonic, error) { ResourceIDs: strings.Join(ids, ","), } - id, err := repo.(rest.Persistable).Save(share) + id, err := repo.Save(r.Context(), share) if err != nil { return nil, err } - share, err = repo.(model.ShareRepository).Get(id) + share, err = repo.Get(r.Context(), id) if err != nil { return nil, err } @@ -90,18 +89,17 @@ func (api *Router) UpdateShare(r *http.Request) (*responses.Subsonic, error) { return nil, err } - repo := api.share.NewRepository(r.Context()) + repo := api.share.Repository() // The update always writes description and downloadable, so read back the // stored value for whichever one the client omitted. description := p.StringPtr("description") downloadable := p.BoolPtr("downloadable") if description == nil || downloadable == nil { - current, err := repo.Read(id) + cur, err := repo.Read(r.Context(), id) if err != nil { return nil, err } - cur := current.(*model.Share) description = cmp.Or(description, &cur.Description) downloadable = cmp.Or(downloadable, &cur.Downloadable) } @@ -113,7 +111,7 @@ func (api *Router) UpdateShare(r *http.Request) (*responses.Subsonic, error) { ExpiresAt: new(p.TimeOr("expires", time.Time{})), } - err = repo.(rest.Persistable).Update(id, share) + err = repo.Update(r.Context(), id, *share) if err != nil { return nil, err } @@ -128,8 +126,8 @@ func (api *Router) DeleteShare(r *http.Request) (*responses.Subsonic, error) { return nil, err } - repo := api.share.NewRepository(r.Context()) - err = repo.(rest.Persistable).Delete(id) + repo := api.share.Repository() + err = repo.Delete(r.Context(), id) if err != nil { return nil, err } diff --git a/server/subsonic/stream.go b/server/subsonic/stream.go index b4a6b821c..fd26ccc4d 100644 --- a/server/subsonic/stream.go +++ b/server/subsonic/stream.go @@ -27,7 +27,7 @@ func (api *Router) Stream(w http.ResponseWriter, r *http.Request) (*responses.Su format, _ := p.String("format") timeOffset := p.IntOr("timeOffset", 0) - mf, err := api.ds.MediaFile(ctx).Get(id) + mf, err := api.ds.MediaFile().Get(ctx, id) if err != nil { return nil, err } diff --git a/server/subsonic/transcode.go b/server/subsonic/transcode.go index d64bce605..7a011a616 100644 --- a/server/subsonic/transcode.go +++ b/server/subsonic/transcode.go @@ -310,7 +310,7 @@ func (api *Router) GetTranscodeDecision(w http.ResponseWriter, r *http.Request) } // Get media file - mf, err := api.ds.MediaFile(ctx).Get(mediaID) + mf, err := api.ds.MediaFile().Get(ctx, mediaID) if err != nil { if errors.Is(err, model.ErrNotFound) { return nil, newError(responses.ErrorDataNotFound, "media file not found: %s", mediaID) @@ -399,7 +399,7 @@ func (api *Router) GetTranscodeStream(w http.ResponseWriter, r *http.Request) (* } // Fetch the media file - mf, err := api.ds.MediaFile(ctx).Get(mediaID) + mf, err := api.ds.MediaFile().Get(ctx, mediaID) if err != nil { if errors.Is(err, model.ErrNotFound) { http.Error(w, "Not Found", http.StatusNotFound) diff --git a/tests/harness/harness.go b/tests/harness/harness.go index 5949c4fae..92196ef26 100644 --- a/tests/harness/harness.go +++ b/tests/harness/harness.go @@ -61,14 +61,14 @@ func SetupDB(ctx context.Context, users ...*model.User) *DB { auth.Init(ds) h.Library = model.Library{ID: 1, Name: "Music Library", Path: "fake:///music"} - Expect(ds.Library(ctx).Put(&h.Library)).To(Succeed()) + Expect(ds.Library().Put(ctx, &h.Library)).To(Succeed()) for _, u := range users { seeded := *u seeded.NewPassword = "password" - Expect(ds.User(ctx).Put(&seeded)).To(Succeed()) - Expect(ds.User(ctx).SetUserLibraries(u.ID, []int{h.Library.ID})).To(Succeed()) - loaded, err := ds.User(ctx).FindByUsername(u.UserName) + Expect(ds.User().Put(ctx, &seeded)).To(Succeed()) + Expect(ds.User().SetUserLibraries(ctx, u.ID, []int{h.Library.ID})).To(Succeed()) + loaded, err := ds.User().FindByUsername(ctx, u.UserName) Expect(err).ToNot(HaveOccurred()) u.Libraries = loaded.Libraries } diff --git a/tests/mock_album_repo.go b/tests/mock_album_repo.go index 1b14f225b..f8100c189 100644 --- a/tests/mock_album_repo.go +++ b/tests/mock_album_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "errors" "sync" "time" @@ -39,7 +40,7 @@ func (m *MockAlbumRepo) SetData(albums model.Albums) { } } -func (m *MockAlbumRepo) Exists(id string) (bool, error) { +func (m *MockAlbumRepo) Exists(_ context.Context, id string) (bool, error) { if m.Err { return false, errors.New("unexpected error") } @@ -47,7 +48,7 @@ func (m *MockAlbumRepo) Exists(id string) (bool, error) { return found, nil } -func (m *MockAlbumRepo) Get(id string) (*model.Album, error) { +func (m *MockAlbumRepo) Get(_ context.Context, id string) (*model.Album, error) { if m.Err { return nil, errors.New("unexpected error") } @@ -57,7 +58,7 @@ func (m *MockAlbumRepo) Get(id string) (*model.Album, error) { return nil, model.ErrNotFound } -func (m *MockAlbumRepo) Put(al *model.Album) error { +func (m *MockAlbumRepo) Put(_ context.Context, al *model.Album) error { if m.Err { return errors.New("unexpected error") } @@ -71,7 +72,7 @@ func (m *MockAlbumRepo) Put(al *model.Album) error { return nil } -func (m *MockAlbumRepo) GetAll(qo ...model.QueryOptions) (model.Albums, error) { +func (m *MockAlbumRepo) GetAll(_ context.Context, qo ...model.QueryOptions) (model.Albums, error) { if len(qo) > 0 { // Recording the last options is a read-path write, and callers resolve concurrently. m.optionsMu.Lock() @@ -84,8 +85,8 @@ func (m *MockAlbumRepo) GetAll(qo ...model.QueryOptions) (model.Albums, error) { return m.All, nil } -func (m *MockAlbumRepo) GetCursor(qo ...model.QueryOptions) (model.AlbumCursor, error) { - res, err := m.GetAll(qo...) +func (m *MockAlbumRepo) GetCursor(ctx context.Context, qo ...model.QueryOptions) (model.AlbumCursor, error) { + res, err := m.GetAll(ctx, qo...) if err != nil { return nil, err } @@ -98,7 +99,7 @@ func (m *MockAlbumRepo) GetCursor(qo ...model.QueryOptions) (model.AlbumCursor, }, nil } -func (m *MockAlbumRepo) IncPlayCount(id string, timestamp time.Time) error { +func (m *MockAlbumRepo) IncPlayCount(_ context.Context, id string, timestamp time.Time) error { if m.Err { return errors.New("unexpected error") } @@ -109,11 +110,11 @@ func (m *MockAlbumRepo) IncPlayCount(id string, timestamp time.Time) error { } return model.ErrNotFound } -func (m *MockAlbumRepo) CountAll(...model.QueryOptions) (int64, error) { +func (m *MockAlbumRepo) CountAll(_ context.Context, _ ...model.QueryOptions) (int64, error) { return int64(len(m.All)), nil } -func (m *MockAlbumRepo) GetTouchedAlbums(libID int) (model.AlbumCursor, error) { +func (m *MockAlbumRepo) GetTouchedAlbums(_ context.Context, libID int) (model.AlbumCursor, error) { if m.Err { return nil, errors.New("unexpected error") } @@ -135,11 +136,11 @@ func (m *MockAlbumRepo) GetTouchedAlbums(libID int) (model.AlbumCursor, error) { }, nil } -func (m *MockAlbumRepo) UpdateExternalInfo(album *model.Album) error { - return m.Put(album) +func (m *MockAlbumRepo) UpdateExternalInfo(ctx context.Context, album *model.Album) error { + return m.Put(ctx, album) } -func (m *MockAlbumRepo) Search(q string, options ...model.QueryOptions) (model.Albums, error) { +func (m *MockAlbumRepo) Search(_ context.Context, q string, options ...model.QueryOptions) (model.Albums, error) { m.SearchQuery = q if len(options) > 0 { m.Options = options[0] @@ -152,7 +153,7 @@ func (m *MockAlbumRepo) Search(q string, options ...model.QueryOptions) (model.A } // ReassignAnnotation reassigns annotations from one album to another -func (m *MockAlbumRepo) ReassignAnnotation(prevID string, newID string) error { +func (m *MockAlbumRepo) ReassignAnnotation(_ context.Context, prevID string, newID string) error { if m.Err { return errors.New("unexpected error") } @@ -165,7 +166,7 @@ func (m *MockAlbumRepo) ReassignAnnotation(prevID string, newID string) error { } // CopyAttributes copies attributes from one album to another -func (m *MockAlbumRepo) CopyAttributes(fromID, toID string, columns ...string) error { +func (m *MockAlbumRepo) CopyAttributes(_ context.Context, fromID, toID string, columns ...string) error { if m.Err { return errors.New("unexpected error") } @@ -191,7 +192,7 @@ func (m *MockAlbumRepo) CopyAttributes(fromID, toID string, columns ...string) e } // SetRating sets the rating for an album -func (m *MockAlbumRepo) SetRating(rating int, itemID string) error { +func (m *MockAlbumRepo) SetRating(_ context.Context, rating int, itemID string) error { if m.Err { return errors.New("unexpected error") } @@ -202,7 +203,7 @@ func (m *MockAlbumRepo) SetRating(rating int, itemID string) error { } // SetStar sets the starred status for albums -func (m *MockAlbumRepo) SetStar(starred bool, itemIDs ...string) error { +func (m *MockAlbumRepo) SetStar(_ context.Context, starred bool, itemIDs ...string) error { if m.Err { return errors.New("unexpected error") } @@ -214,7 +215,7 @@ func (m *MockAlbumRepo) SetStar(starred bool, itemIDs ...string) error { return nil } -func (m *MockAlbumRepo) GetYears(libraryIDs ...int) ([]int, error) { +func (m *MockAlbumRepo) GetYears(_ context.Context, libraryIDs ...int) ([]int, error) { if m.Err { return nil, errors.New("error") } diff --git a/tests/mock_artist_repo.go b/tests/mock_artist_repo.go index db7d54d5d..fe63edd94 100644 --- a/tests/mock_artist_repo.go +++ b/tests/mock_artist_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "errors" "time" @@ -32,7 +33,7 @@ func (m *MockArtistRepo) SetData(artists model.Artists) { } } -func (m *MockArtistRepo) Exists(id string) (bool, error) { +func (m *MockArtistRepo) Exists(_ context.Context, id string) (bool, error) { if m.Err { return false, errors.New("Error!") } @@ -40,7 +41,7 @@ func (m *MockArtistRepo) Exists(id string) (bool, error) { return found, nil } -func (m *MockArtistRepo) Get(id string) (*model.Artist, error) { +func (m *MockArtistRepo) Get(_ context.Context, id string) (*model.Artist, error) { if m.Err { return nil, errors.New("Error!") } @@ -50,7 +51,7 @@ func (m *MockArtistRepo) Get(id string) (*model.Artist, error) { return nil, model.ErrNotFound } -func (m *MockArtistRepo) Put(ar *model.Artist, columsToUpdate ...string) error { +func (m *MockArtistRepo) Put(_ context.Context, ar *model.Artist, columsToUpdate ...string) error { if m.Err { return errors.New("error") } @@ -64,7 +65,7 @@ func (m *MockArtistRepo) Put(ar *model.Artist, columsToUpdate ...string) error { return nil } -func (m *MockArtistRepo) IncPlayCount(id string, timestamp time.Time) error { +func (m *MockArtistRepo) IncPlayCount(_ context.Context, id string, timestamp time.Time) error { if m.Err { return errors.New("error") } @@ -76,7 +77,7 @@ func (m *MockArtistRepo) IncPlayCount(id string, timestamp time.Time) error { return model.ErrNotFound } -func (m *MockArtistRepo) SetStar(starred bool, itemIDs ...string) error { +func (m *MockArtistRepo) SetStar(_ context.Context, starred bool, itemIDs ...string) error { if m.Err { return errors.New("error") } @@ -88,7 +89,7 @@ func (m *MockArtistRepo) SetStar(starred bool, itemIDs ...string) error { return nil } -func (m *MockArtistRepo) SetRating(rating int, itemID string) error { +func (m *MockArtistRepo) SetRating(_ context.Context, rating int, itemID string) error { if m.Err { return errors.New("error") } @@ -98,7 +99,7 @@ func (m *MockArtistRepo) SetRating(rating int, itemID string) error { return nil } -func (m *MockArtistRepo) GetAll(options ...model.QueryOptions) (model.Artists, error) { +func (m *MockArtistRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Artists, error) { if len(options) > 0 { m.Options = options[0] } @@ -116,8 +117,8 @@ func (m *MockArtistRepo) GetAll(options ...model.QueryOptions) (model.Artists, e return allArtists, nil } -func (m *MockArtistRepo) GetCursor(options ...model.QueryOptions) (model.ArtistCursor, error) { - res, err := m.GetAll(options...) +func (m *MockArtistRepo) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.ArtistCursor, error) { + res, err := m.GetAll(ctx, options...) if err != nil { return nil, err } @@ -130,30 +131,30 @@ func (m *MockArtistRepo) GetCursor(options ...model.QueryOptions) (model.ArtistC }, nil } -func (m *MockArtistRepo) UpdateExternalInfo(artist *model.Artist) error { - return m.Put(artist) +func (m *MockArtistRepo) UpdateExternalInfo(ctx context.Context, artist *model.Artist) error { + return m.Put(ctx, artist) } -func (m *MockArtistRepo) RefreshStats(allArtists bool) (int64, error) { +func (m *MockArtistRepo) RefreshStats(_ context.Context, allArtists bool) (int64, error) { if m.Err { return 0, errors.New("mock repo error") } return int64(len(m.Data)), nil } -func (m *MockArtistRepo) RefreshPlayCounts() (int64, error) { +func (m *MockArtistRepo) RefreshPlayCounts(_ context.Context) (int64, error) { if m.Err { return 0, errors.New("mock repo error") } return int64(len(m.Data)), nil } -func (m *MockArtistRepo) GetIndex(includeMissing bool, libraryIds []int, roles ...model.Role) (model.ArtistIndexes, error) { +func (m *MockArtistRepo) GetIndex(ctx context.Context, includeMissing bool, libraryIds []int, roles ...model.Role) (model.ArtistIndexes, error) { if m.Err { return nil, errors.New("mock repo error") } - artists, err := m.GetAll() + artists, err := m.GetAll(ctx) if err != nil { return nil, err } @@ -181,14 +182,14 @@ func (m *MockArtistRepo) GetIndex(includeMissing bool, libraryIds []int, roles . return result, nil } -func (m *MockArtistRepo) CountAll(...model.QueryOptions) (int64, error) { +func (m *MockArtistRepo) CountAll(context.Context, ...model.QueryOptions) (int64, error) { if m.Err { return 0, errors.New("mock repo error") } return int64(len(m.Data)), nil } -func (m *MockArtistRepo) Search(q string, options ...model.QueryOptions) (model.Artists, error) { +func (m *MockArtistRepo) Search(ctx context.Context, q string, options ...model.QueryOptions) (model.Artists, error) { if len(options) > 0 { m.Options = options[0] } @@ -196,7 +197,7 @@ func (m *MockArtistRepo) Search(q string, options ...model.QueryOptions) (model. return nil, errors.New("unexpected error") } // Simple mock implementation - just return all artists for testing - return m.GetAll() + return m.GetAll(ctx) } var _ model.ArtistRepository = (*MockArtistRepo)(nil) diff --git a/tests/mock_artwork_queue_repo.go b/tests/mock_artwork_queue_repo.go index c482e2150..c6a7917f0 100644 --- a/tests/mock_artwork_queue_repo.go +++ b/tests/mock_artwork_queue_repo.go @@ -2,6 +2,7 @@ package tests import ( "cmp" + "context" "slices" "sync" "time" @@ -26,7 +27,7 @@ func CreateMockArtworkQueueRepo() *MockArtworkQueueRepo { return &MockArtworkQueueRepo{Data: map[string]model.ArtworkQueueItem{}} } -func (m *MockArtworkQueueRepo) Get(kind model.Kind, id, imageType string) (*model.ArtworkQueueItem, error) { +func (m *MockArtworkQueueRepo) Get(_ context.Context, kind model.Kind, id, imageType string) (*model.ArtworkQueueItem, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -39,7 +40,7 @@ func (m *MockArtworkQueueRepo) Get(kind model.Kind, id, imageType string) (*mode return &it, nil } -func (m *MockArtworkQueueRepo) Enqueue(items ...model.ArtworkQueueItem) error { +func (m *MockArtworkQueueRepo) Enqueue(_ context.Context, items ...model.ArtworkQueueItem) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -73,7 +74,7 @@ func (m *MockArtworkQueueRepo) enqueueLocked(items []model.ArtworkQueueItem) { } // EnqueueIfMissing mirrors the SQL anti-join: skip anything that already has an item_artwork row. -func (m *MockArtworkQueueRepo) EnqueueIfMissing(items ...model.ArtworkQueueItem) error { +func (m *MockArtworkQueueRepo) EnqueueIfMissing(_ context.Context, items ...model.ArtworkQueueItem) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -96,7 +97,7 @@ func (m *MockArtworkQueueRepo) EnqueueIfMissing(items ...model.ArtworkQueueItem) return nil } -func (m *MockArtworkQueueRepo) DequeueBatch(n int, kinds ...string) ([]model.ArtworkQueueItem, error) { +func (m *MockArtworkQueueRepo) DequeueBatch(_ context.Context, n int, kinds ...string) ([]model.ArtworkQueueItem, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -118,7 +119,7 @@ func (m *MockArtworkQueueRepo) DequeueBatch(n int, kinds ...string) ([]model.Art return res, nil } -func (m *MockArtworkQueueRepo) MarkFailedIfUnchanged(kind, id, imageType string, seenRetryAt, retryAt time.Time, trace string) error { +func (m *MockArtworkQueueRepo) MarkFailedIfUnchanged(_ context.Context, kind, id, imageType string, seenRetryAt, retryAt time.Time, trace string) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -134,7 +135,7 @@ func (m *MockArtworkQueueRepo) MarkFailedIfUnchanged(kind, id, imageType string, return nil } -func (m *MockArtworkQueueRepo) DeleteIfUnchanged(kind, id, imageType string, retryAt time.Time) error { +func (m *MockArtworkQueueRepo) DeleteIfUnchanged(_ context.Context, kind, id, imageType string, retryAt time.Time) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -147,7 +148,7 @@ func (m *MockArtworkQueueRepo) DeleteIfUnchanged(kind, id, imageType string, ret return nil } -func (m *MockArtworkQueueRepo) PurgeDangling() (int64, error) { +func (m *MockArtworkQueueRepo) PurgeDangling(context.Context) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -174,7 +175,7 @@ func queueFilterMatches(it model.ArtworkQueueItem, kinds []model.Kind, prioritie (len(priorities) == 0 || slices.Contains(priorities, it.Priority)) } -func (m *MockArtworkQueueRepo) PurgeQueued(kinds []model.Kind, priorities []int) (int64, error) { +func (m *MockArtworkQueueRepo) PurgeQueued(_ context.Context, kinds []model.Kind, priorities []int) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -191,7 +192,7 @@ func (m *MockArtworkQueueRepo) PurgeQueued(kinds []model.Kind, priorities []int) return purged, nil } -func (m *MockArtworkQueueRepo) Count() (int64, error) { +func (m *MockArtworkQueueRepo) Count(context.Context) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -200,7 +201,7 @@ func (m *MockArtworkQueueRepo) Count() (int64, error) { return int64(len(m.Data)), nil } -func (m *MockArtworkQueueRepo) CountQueued(kinds []model.Kind, priorities []int) ([]model.ArtworkQueueStat, error) { +func (m *MockArtworkQueueRepo) CountQueued(_ context.Context, kinds []model.Kind, priorities []int) ([]model.ArtworkQueueStat, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -226,7 +227,7 @@ func (m *MockArtworkQueueRepo) CountQueued(kinds []model.Kind, priorities []int) return res, nil } -func (m *MockArtworkQueueRepo) EnqueuePreservingBackoff(items ...model.ArtworkQueueItem) error { +func (m *MockArtworkQueueRepo) EnqueuePreservingBackoff(_ context.Context, items ...model.ArtworkQueueItem) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -276,7 +277,7 @@ func (m *MockArtworkQueueRepo) matchingSource(kind model.Kind, sources []string) return res } -func (m *MockArtworkQueueRepo) CountBySource(kind model.Kind, sources []string) (int64, error) { +func (m *MockArtworkQueueRepo) CountBySource(_ context.Context, kind model.Kind, sources []string) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -285,7 +286,7 @@ func (m *MockArtworkQueueRepo) CountBySource(kind model.Kind, sources []string) return int64(len(m.matchingSource(kind, sources))), nil } -func (m *MockArtworkQueueRepo) SourcesInUse(kind model.Kind) ([]string, error) { +func (m *MockArtworkQueueRepo) SourcesInUse(_ context.Context, kind model.Kind) ([]string, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -295,7 +296,7 @@ func (m *MockArtworkQueueRepo) SourcesInUse(kind model.Kind) ([]string, error) { return slice.Unique(sources), nil } -func (m *MockArtworkQueueRepo) EnqueueBySource(kind model.Kind, sources []string, priority int) (int64, error) { +func (m *MockArtworkQueueRepo) EnqueueBySource(_ context.Context, kind model.Kind, sources []string, priority int) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -322,7 +323,7 @@ func (m *MockArtworkQueueRepo) EnqueueBySource(kind model.Kind, sources []string } // EnqueueAllMissing mirrors the SQL set-difference insert: ExistingIDs[kind] minus ItemArtworkSource. -func (m *MockArtworkQueueRepo) EnqueueAllMissing(kind model.Kind, priority int) (int64, error) { +func (m *MockArtworkQueueRepo) EnqueueAllMissing(_ context.Context, kind model.Kind, priority int) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { diff --git a/tests/mock_artwork_repo.go b/tests/mock_artwork_repo.go index 5d76a0169..350c0e9f6 100644 --- a/tests/mock_artwork_repo.go +++ b/tests/mock_artwork_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "maps" "sync" "time" @@ -26,7 +27,7 @@ func CreateMockArtworkRepo() *MockArtworkRepo { func iaKey(kind, id, imageType string) string { return kind + "|" + id + "|" + imageType } -func (m *MockArtworkRepo) GetImage(hash string) (*model.Artwork, error) { +func (m *MockArtworkRepo) GetImage(_ context.Context, hash string) (*model.Artwork, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -38,7 +39,7 @@ func (m *MockArtworkRepo) GetImage(hash string) (*model.Artwork, error) { return nil, model.ErrNotFound } -func (m *MockArtworkRepo) PutImage(a *model.Artwork) error { +func (m *MockArtworkRepo) PutImage(_ context.Context, a *model.Artwork) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -50,7 +51,7 @@ func (m *MockArtworkRepo) PutImage(a *model.Artwork) error { return nil } -func (m *MockArtworkRepo) GetMimeByHash() (map[string]string, error) { +func (m *MockArtworkRepo) GetMimeByHash(context.Context) (map[string]string, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -63,7 +64,7 @@ func (m *MockArtworkRepo) GetMimeByHash() (map[string]string, error) { return mimes, nil } -func (m *MockArtworkRepo) PurgeDanglingItems() (int64, error) { +func (m *MockArtworkRepo) PurgeDanglingItems(context.Context) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -84,7 +85,7 @@ func (m *MockArtworkRepo) PurgeDanglingItems() (int64, error) { return purged, nil } -func (m *MockArtworkRepo) PurgeOrphans(createdBefore time.Time) (int64, error) { +func (m *MockArtworkRepo) PurgeOrphans(_ context.Context, createdBefore time.Time) (int64, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -110,7 +111,7 @@ func (m *MockArtworkRepo) referenced(hash string) bool { return false } -func (m *MockArtworkRepo) GetItemArtwork(kind model.Kind, id, imageType string) (*model.ItemArtwork, error) { +func (m *MockArtworkRepo) GetItemArtwork(_ context.Context, kind model.Kind, id, imageType string) (*model.ItemArtwork, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -122,7 +123,7 @@ func (m *MockArtworkRepo) GetItemArtwork(kind model.Kind, id, imageType string) return nil, model.ErrNotFound } -func (m *MockArtworkRepo) PutLastFailure(kind model.Kind, id, imageType, trace string) error { +func (m *MockArtworkRepo) PutLastFailure(_ context.Context, kind model.Kind, id, imageType, trace string) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -136,7 +137,7 @@ func (m *MockArtworkRepo) PutLastFailure(kind model.Kind, id, imageType, trace s return nil } -func (m *MockArtworkRepo) PutItemArtwork(ia *model.ItemArtwork) error { +func (m *MockArtworkRepo) PutItemArtwork(_ context.Context, ia *model.ItemArtwork) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -153,7 +154,7 @@ func (m *MockArtworkRepo) PutItemArtwork(ia *model.ItemArtwork) error { return nil } -func (m *MockArtworkRepo) DeleteForItems(kind model.Kind, ids []string) error { +func (m *MockArtworkRepo) DeleteForItems(_ context.Context, kind model.Kind, ids []string) error { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { @@ -167,7 +168,7 @@ func (m *MockArtworkRepo) DeleteForItems(kind model.Kind, ids []string) error { return nil } -func (m *MockArtworkRepo) GetInfoForItems(kind model.Kind, ids []string) (map[string]model.ItemArtworkInfo, error) { +func (m *MockArtworkRepo) GetInfoForItems(_ context.Context, kind model.Kind, ids []string) (map[string]model.ItemArtworkInfo, error) { m.mu.Lock() defer m.mu.Unlock() if m.Err != nil { diff --git a/tests/mock_data_store.go b/tests/mock_data_store.go index 6a0ebbb31..db798eece 100644 --- a/tests/mock_data_store.go +++ b/tests/mock_data_store.go @@ -38,76 +38,76 @@ type MockDataStore struct { GCError error } -func (db *MockDataStore) Library(ctx context.Context) model.LibraryRepository { +func (db *MockDataStore) Library() model.LibraryRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedLibrary != nil { return db.MockedLibrary } if db.RealDS != nil { - return db.RealDS.Library(ctx) + return db.RealDS.Library() } db.MockedLibrary = &MockLibraryRepo{} return db.MockedLibrary } -func (db *MockDataStore) Folder(ctx context.Context) model.FolderRepository { +func (db *MockDataStore) Folder() model.FolderRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedFolder != nil { return db.MockedFolder } if db.RealDS != nil { - return db.RealDS.Folder(ctx) + return db.RealDS.Folder() } db.MockedFolder = struct{ model.FolderRepository }{} return db.MockedFolder } -func (db *MockDataStore) Tag(ctx context.Context) model.TagRepository { +func (db *MockDataStore) Tag() model.TagRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedTag != nil { return db.MockedTag } if db.RealDS != nil { - return db.RealDS.Tag(ctx) + return db.RealDS.Tag() } db.MockedTag = &MockTagRepo{} return db.MockedTag } -func (db *MockDataStore) Album(ctx context.Context) model.AlbumRepository { +func (db *MockDataStore) Album() model.AlbumRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedAlbum != nil { return db.MockedAlbum } if db.RealDS != nil { - return db.RealDS.Album(ctx) + return db.RealDS.Album() } db.MockedAlbum = CreateMockAlbumRepo() return db.MockedAlbum } -func (db *MockDataStore) Artist(ctx context.Context) model.ArtistRepository { +func (db *MockDataStore) Artist() model.ArtistRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedArtist != nil { return db.MockedArtist } if db.RealDS != nil { - return db.RealDS.Artist(ctx) + return db.RealDS.Artist() } db.MockedArtist = CreateMockArtistRepo() return db.MockedArtist } -func (db *MockDataStore) MediaFile(ctx context.Context) model.MediaFileRepository { +func (db *MockDataStore) MediaFile() model.MediaFileRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.RealDS != nil && db.MockedMediaFile == nil { - return db.RealDS.MediaFile(ctx) + return db.RealDS.MediaFile() } if db.MockedMediaFile == nil { db.MockedMediaFile = CreateMockMediaFileRepo() @@ -115,128 +115,128 @@ func (db *MockDataStore) MediaFile(ctx context.Context) model.MediaFileRepositor return db.MockedMediaFile } -func (db *MockDataStore) Genre(ctx context.Context) model.GenreRepository { +func (db *MockDataStore) Genre() model.GenreRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedGenre != nil { return db.MockedGenre } if db.RealDS != nil { - return db.RealDS.Genre(ctx) + return db.RealDS.Genre() } db.MockedGenre = &MockedGenreRepo{} return db.MockedGenre } -func (db *MockDataStore) Playlist(ctx context.Context) model.PlaylistRepository { +func (db *MockDataStore) Playlist() model.PlaylistRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedPlaylist != nil { return db.MockedPlaylist } if db.RealDS != nil { - return db.RealDS.Playlist(ctx) + return db.RealDS.Playlist() } db.MockedPlaylist = CreateMockPlaylistRepo() return db.MockedPlaylist } -func (db *MockDataStore) PlayQueue(ctx context.Context) model.PlayQueueRepository { +func (db *MockDataStore) PlayQueue() model.PlayQueueRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedPlayQueue != nil { return db.MockedPlayQueue } if db.RealDS != nil { - return db.RealDS.PlayQueue(ctx) + return db.RealDS.PlayQueue() } db.MockedPlayQueue = &MockPlayQueueRepo{} return db.MockedPlayQueue } -func (db *MockDataStore) UserProps(ctx context.Context) model.UserPropsRepository { +func (db *MockDataStore) UserProps() model.UserPropsRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedUserProps != nil { return db.MockedUserProps } if db.RealDS != nil { - return db.RealDS.UserProps(ctx) + return db.RealDS.UserProps() } db.MockedUserProps = &MockedUserPropsRepo{} return db.MockedUserProps } -func (db *MockDataStore) Property(ctx context.Context) model.PropertyRepository { +func (db *MockDataStore) Property() model.PropertyRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedProperty != nil { return db.MockedProperty } if db.RealDS != nil { - return db.RealDS.Property(ctx) + return db.RealDS.Property() } db.MockedProperty = &MockedPropertyRepo{} return db.MockedProperty } -func (db *MockDataStore) Share(ctx context.Context) model.ShareRepository { +func (db *MockDataStore) Share() model.ShareRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedShare != nil { return db.MockedShare } if db.RealDS != nil { - return db.RealDS.Share(ctx) + return db.RealDS.Share() } db.MockedShare = &MockShareRepo{} return db.MockedShare } -func (db *MockDataStore) User(ctx context.Context) model.UserRepository { +func (db *MockDataStore) User() model.UserRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedUser != nil { return db.MockedUser } if db.RealDS != nil { - return db.RealDS.User(ctx) + return db.RealDS.User() } db.MockedUser = CreateMockUserRepo() return db.MockedUser } -func (db *MockDataStore) Transcoding(ctx context.Context) model.TranscodingRepository { +func (db *MockDataStore) Transcoding() model.TranscodingRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedTranscoding != nil { return db.MockedTranscoding } if db.RealDS != nil { - return db.RealDS.Transcoding(ctx) + return db.RealDS.Transcoding() } db.MockedTranscoding = struct{ model.TranscodingRepository }{} return db.MockedTranscoding } -func (db *MockDataStore) Player(ctx context.Context) model.PlayerRepository { +func (db *MockDataStore) Player() model.PlayerRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedPlayer != nil { return db.MockedPlayer } if db.RealDS != nil { - return db.RealDS.Player(ctx) + return db.RealDS.Player() } db.MockedPlayer = struct{ model.PlayerRepository }{} return db.MockedPlayer } -func (db *MockDataStore) ScrobbleBuffer(ctx context.Context) model.ScrobbleBufferRepository { +func (db *MockDataStore) ScrobbleBuffer() model.ScrobbleBufferRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.RealDS != nil && db.MockedScrobbleBuffer == nil { - return db.RealDS.ScrobbleBuffer(ctx) + return db.RealDS.ScrobbleBuffer() } db.scrobbleBufferMu.Lock() defer db.scrobbleBufferMu.Unlock() @@ -246,75 +246,75 @@ func (db *MockDataStore) ScrobbleBuffer(ctx context.Context) model.ScrobbleBuffe return db.MockedScrobbleBuffer } -func (db *MockDataStore) Scrobble(ctx context.Context) model.ScrobbleRepository { +func (db *MockDataStore) Scrobble() model.ScrobbleRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedScrobble != nil { return db.MockedScrobble } if db.RealDS != nil { - return db.RealDS.Scrobble(ctx) + return db.RealDS.Scrobble() } - db.MockedScrobble = &MockScrobbleRepo{ctx: ctx} + db.MockedScrobble = &MockScrobbleRepo{} return db.MockedScrobble } -func (db *MockDataStore) Radio(ctx context.Context) model.RadioRepository { +func (db *MockDataStore) Radio() model.RadioRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedRadio != nil { return db.MockedRadio } if db.RealDS != nil { - return db.RealDS.Radio(ctx) + return db.RealDS.Radio() } db.MockedRadio = CreateMockedRadioRepo() return db.MockedRadio } -func (db *MockDataStore) Plugin(ctx context.Context) model.PluginRepository { +func (db *MockDataStore) Plugin() model.PluginRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedPlugin != nil { return db.MockedPlugin } if db.RealDS != nil { - return db.RealDS.Plugin(ctx) + return db.RealDS.Plugin() } db.MockedPlugin = CreateMockPluginRepo() return db.MockedPlugin } -func (db *MockDataStore) Artwork(ctx context.Context) model.ArtworkRepository { +func (db *MockDataStore) Artwork() model.ArtworkRepository { db.repoMu.Lock() defer db.repoMu.Unlock() - return db.artworkLocked(ctx) + return db.artworkLocked() } // artworkLocked is the body of Artwork for callers already holding repoMu; repoMu is a plain // Mutex, so re-entering through the exported method would deadlock. -func (db *MockDataStore) artworkLocked(ctx context.Context) model.ArtworkRepository { +func (db *MockDataStore) artworkLocked() model.ArtworkRepository { if db.MockedArtwork != nil { return db.MockedArtwork } if db.RealDS != nil { - return db.RealDS.Artwork(ctx) + return db.RealDS.Artwork() } db.MockedArtwork = CreateMockArtworkRepo() return db.MockedArtwork } -func (db *MockDataStore) ArtworkQueue(ctx context.Context) model.ArtworkQueueRepository { +func (db *MockDataStore) ArtworkQueue() model.ArtworkQueueRepository { db.repoMu.Lock() defer db.repoMu.Unlock() if db.MockedArtworkQueue != nil { return db.MockedArtworkQueue } if db.RealDS != nil { - return db.RealDS.ArtworkQueue(ctx) + return db.RealDS.ArtworkQueue() } q := CreateMockArtworkQueueRepo() - if aw, ok := db.artworkLocked(ctx).(*MockArtworkRepo); ok { + if aw, ok := db.artworkLocked().(*MockArtworkRepo); ok { q.ItemArtworkSource = aw } db.MockedArtworkQueue = q @@ -333,37 +333,6 @@ func (db *MockDataStore) WithTxRetry(ctx context.Context, block func(ctx context return block(ctx, db) } -func (db *MockDataStore) Resource(ctx context.Context, m any) model.ResourceRepository { - switch m.(type) { - case model.MediaFile, *model.MediaFile: - return db.MediaFile(ctx).(model.ResourceRepository) - case model.Album, *model.Album: - return db.Album(ctx).(model.ResourceRepository) - case model.Artist, *model.Artist: - return db.Artist(ctx).(model.ResourceRepository) - case model.User, *model.User: - return db.User(ctx).(model.ResourceRepository) - case model.Playlist, *model.Playlist: - return db.Playlist(ctx).(model.ResourceRepository) - case model.Radio, *model.Radio: - return db.Radio(ctx).(model.ResourceRepository) - case model.Share, *model.Share: - return db.Share(ctx).(model.ResourceRepository) - case model.Genre, *model.Genre: - return db.Genre(ctx).(model.ResourceRepository) - case model.Tag, *model.Tag: - return db.Tag(ctx).(model.ResourceRepository) - case model.Transcoding, *model.Transcoding: - return db.Transcoding(ctx).(model.ResourceRepository) - case model.Player, *model.Player: - return db.Player(ctx).(model.ResourceRepository) - case model.Plugin, *model.Plugin: - return db.Plugin(ctx).(model.ResourceRepository) - default: - return struct{ model.ResourceRepository }{} - } -} - func (db *MockDataStore) GC(context.Context, ...int) error { db.GCCalled = true if db.GCError != nil { diff --git a/tests/mock_genre_repo.go b/tests/mock_genre_repo.go index ad3ee1a6a..917e2c626 100644 --- a/tests/mock_genre_repo.go +++ b/tests/mock_genre_repo.go @@ -1,6 +1,9 @@ package tests import ( + "context" + + "github.com/deluan/rest" "github.com/navidrome/navidrome/model" ) @@ -16,7 +19,7 @@ func (r *MockedGenreRepo) init() { } } -func (r *MockedGenreRepo) GetAll(options ...model.QueryOptions) (model.Genres, error) { +func (r *MockedGenreRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Genres, error) { if len(options) > 0 { r.Options = options[0] } @@ -32,7 +35,7 @@ func (r *MockedGenreRepo) GetAll(options ...model.QueryOptions) (model.Genres, e return all, nil } -func (r *MockedGenreRepo) Get(id string) (*model.Genre, error) { +func (r *MockedGenreRepo) Get(_ context.Context, id string) (*model.Genre, error) { if r.Error != nil { return nil, r.Error } @@ -51,3 +54,21 @@ func (r *MockedGenreRepo) Put(g *model.Genre) error { r.Data[g.ID] = *g return nil } + +func (r *MockedGenreRepo) Count(context.Context, ...rest.QueryOptions) (int64, error) { + if r.Error != nil { + return 0, r.Error + } + r.init() + return int64(len(r.Data)), nil +} + +func (r *MockedGenreRepo) Read(ctx context.Context, id string) (*model.Genre, error) { + return r.Get(ctx, id) +} + +func (r *MockedGenreRepo) ReadAll(ctx context.Context, _ ...rest.QueryOptions) ([]model.Genre, error) { + return r.GetAll(ctx) +} + +var _ model.GenreRepository = (*MockedGenreRepo)(nil) diff --git a/tests/mock_library_repo.go b/tests/mock_library_repo.go index 0f7af2aab..e21dcccce 100644 --- a/tests/mock_library_repo.go +++ b/tests/mock_library_repo.go @@ -27,7 +27,7 @@ func (m *MockLibraryRepo) SetData(data model.Libraries) { } } -func (m *MockLibraryRepo) GetAll(...model.QueryOptions) (model.Libraries, error) { +func (m *MockLibraryRepo) GetAll(_ context.Context, _ ...model.QueryOptions) (model.Libraries, error) { if m.Err != nil { return nil, m.Err } @@ -42,7 +42,7 @@ func (m *MockLibraryRepo) GetAll(...model.QueryOptions) (model.Libraries, error) return libraries, nil } -func (m *MockLibraryRepo) CountAll(qo ...model.QueryOptions) (int64, error) { +func (m *MockLibraryRepo) CountAll(_ context.Context, qo ...model.QueryOptions) (int64, error) { if m.Err != nil { return 0, m.Err } @@ -71,7 +71,7 @@ func (m *MockLibraryRepo) CountAll(qo ...model.QueryOptions) (int64, error) { return int64(len(m.Data)), nil } -func (m *MockLibraryRepo) Get(id int) (*model.Library, error) { +func (m *MockLibraryRepo) Get(_ context.Context, id int) (*model.Library, error) { if m.Err != nil { return nil, m.Err } @@ -81,7 +81,7 @@ func (m *MockLibraryRepo) Get(id int) (*model.Library, error) { return nil, model.ErrNotFound } -func (m *MockLibraryRepo) GetPath(id int) (string, error) { +func (m *MockLibraryRepo) GetPath(_ context.Context, id int) (string, error) { if m.Err != nil { return "", m.Err } @@ -91,7 +91,7 @@ func (m *MockLibraryRepo) GetPath(id int) (string, error) { return "", model.ErrNotFound } -func (m *MockLibraryRepo) Put(library *model.Library, colsToUpdate ...string) error { +func (m *MockLibraryRepo) Put(_ context.Context, library *model.Library, colsToUpdate ...string) error { m.PutCols = colsToUpdate if m.PutFn != nil { return m.PutFn(library) @@ -106,7 +106,7 @@ func (m *MockLibraryRepo) Put(library *model.Library, colsToUpdate ...string) er return nil } -func (m *MockLibraryRepo) Delete(id int) error { +func (m *MockLibraryRepo) Delete(_ context.Context, id int) error { if m.Err != nil { return m.Err } @@ -117,48 +117,48 @@ func (m *MockLibraryRepo) Delete(id int) error { return nil } -func (m *MockLibraryRepo) StoreMusicFolder() error { +func (m *MockLibraryRepo) StoreMusicFolder(_ context.Context) error { if m.Err != nil { return m.Err } return nil } -func (m *MockLibraryRepo) AddArtist(id int, artistID string) error { +func (m *MockLibraryRepo) AddArtist(_ context.Context, id int, artistID string) error { if m.Err != nil { return m.Err } return nil } -func (m *MockLibraryRepo) ScanBegin(id int, fullScan bool) error { +func (m *MockLibraryRepo) ScanBegin(_ context.Context, id int, fullScan bool) error { if m.Err != nil { return m.Err } return nil } -func (m *MockLibraryRepo) ScanEnd(id int) error { +func (m *MockLibraryRepo) ScanEnd(_ context.Context, id int) error { if m.Err != nil { return m.Err } return nil } -func (m *MockLibraryRepo) ScanInProgress() (bool, error) { +func (m *MockLibraryRepo) ScanInProgress(_ context.Context) (bool, error) { if m.Err != nil { return false, m.Err } return false, nil } -func (m *MockLibraryRepo) RefreshStats(id int) error { +func (m *MockLibraryRepo) RefreshStats(_ context.Context, id int) error { return nil } // User-library association methods - mock implementations -func (m *MockLibraryRepo) GetUsersWithLibraryAccess(libraryID int) (model.Users, error) { +func (m *MockLibraryRepo) GetUsersWithLibraryAccess(_ context.Context, libraryID int) (model.Users, error) { if m.Err != nil { return nil, m.Err } @@ -166,31 +166,22 @@ func (m *MockLibraryRepo) GetUsersWithLibraryAccess(libraryID int) (model.Users, return model.Users{}, nil } -func (m *MockLibraryRepo) Count(options ...rest.QueryOptions) (int64, error) { - return m.CountAll() +func (m *MockLibraryRepo) Count(ctx context.Context, _ ...rest.QueryOptions) (int64, error) { + return m.CountAll(ctx) } -func (m *MockLibraryRepo) Read(id string) (any, error) { +func (m *MockLibraryRepo) Read(ctx context.Context, id string) (*model.Library, error) { idInt, _ := strconv.Atoi(id) - return m.Get(idInt) + return m.Get(ctx, idInt) } -func (m *MockLibraryRepo) ReadAll(options ...rest.QueryOptions) (any, error) { - return m.GetAll() -} - -func (m *MockLibraryRepo) EntityName() string { - return "library" -} - -func (m *MockLibraryRepo) NewInstance() any { - return &model.Library{} +func (m *MockLibraryRepo) ReadAll(ctx context.Context, _ ...rest.QueryOptions) ([]model.Library, error) { + return m.GetAll(ctx) } // REST Repository methods (string-based IDs) -func (m *MockLibraryRepo) Save(entity any) (string, error) { - lib := entity.(*model.Library) +func (m *MockLibraryRepo) Save(_ context.Context, lib *model.Library) (string, error) { if m.Err != nil { return "", m.Err } @@ -214,8 +205,8 @@ func (m *MockLibraryRepo) Save(entity any) (string, error) { return strconv.Itoa(lib.ID), nil } -func (m *MockLibraryRepo) Update(id string, entity any, cols ...string) error { - lib := entity.(*model.Library) +func (m *MockLibraryRepo) Update(_ context.Context, id string, entity model.Library, _ ...string) error { + lib := &entity if m.Err != nil { return m.Err } @@ -307,4 +298,4 @@ func (m *MockLibraryRepo) ValidateLibraryAccess(ctx context.Context, userID stri } var _ model.LibraryRepository = (*MockLibraryRepo)(nil) -var _ model.ResourceRepository = (*MockLibraryRepo)(nil) +var _ rest.Repository[model.Library] = (*MockLibraryRepo)(nil) diff --git a/tests/mock_library_service.go b/tests/mock_library_service.go index 78693197d..f5e1f0387 100644 --- a/tests/mock_library_service.go +++ b/tests/mock_library_service.go @@ -14,7 +14,7 @@ type MockLibraryService struct { *MockLibraryRepo } -// MockLibraryRestAdapter adapts MockLibraryRepo to rest.Repository interface +// MockLibraryRestAdapter adapts MockLibraryRepo to the REST repository interface type MockLibraryRestAdapter struct { *MockLibraryRepo } @@ -33,12 +33,15 @@ func NewMockLibraryService() *MockLibraryService { return &MockLibraryService{MockLibraryRepo: repo} } -func (m *MockLibraryService) NewRepository(ctx context.Context) rest.Repository { +func (m *MockLibraryService) Repository() rest.Repository[model.Library] { return &MockLibraryRestAdapter{MockLibraryRepo: m.MockLibraryRepo} } -// rest.Repository interface implementation - -func (a *MockLibraryRestAdapter) Delete(id string) error { - return a.DeleteByStringID(id) +func (a *MockLibraryRestAdapter) Delete(_ context.Context, ids ...string) error { + for _, id := range ids { + if err := a.DeleteByStringID(id); err != nil { + return err + } + } + return nil } diff --git a/tests/mock_mediafile_repo.go b/tests/mock_mediafile_repo.go index 2093a007d..392a4d8e0 100644 --- a/tests/mock_mediafile_repo.go +++ b/tests/mock_mediafile_repo.go @@ -2,6 +2,7 @@ package tests import ( "cmp" + "context" "errors" "maps" "slices" @@ -55,7 +56,7 @@ func (m *MockMediaFileRepo) SetData(mfs model.MediaFiles) { } } -func (m *MockMediaFileRepo) Exists(id string) (bool, error) { +func (m *MockMediaFileRepo) Exists(_ context.Context, id string) (bool, error) { if m.Err { return false, errors.New("error") } @@ -63,7 +64,7 @@ func (m *MockMediaFileRepo) Exists(id string) (bool, error) { return found, nil } -func (m *MockMediaFileRepo) Get(id string) (*model.MediaFile, error) { +func (m *MockMediaFileRepo) Get(_ context.Context, id string) (*model.MediaFile, error) { if m.Err { return nil, errors.New("error") } @@ -77,7 +78,7 @@ func (m *MockMediaFileRepo) Get(id string) (*model.MediaFile, error) { return nil, model.ErrNotFound } -func (m *MockMediaFileRepo) AddBookmark(id, _ string, _ int64) error { +func (m *MockMediaFileRepo) AddBookmark(_ context.Context, id, _ string, _ int64) error { if m.Err { return errors.New("error") } @@ -85,7 +86,7 @@ func (m *MockMediaFileRepo) AddBookmark(id, _ string, _ int64) error { return nil } -func (m *MockMediaFileRepo) GetWithParticipants(id string) (*model.MediaFile, error) { +func (m *MockMediaFileRepo) GetWithParticipants(_ context.Context, id string) (*model.MediaFile, error) { if m.Err { return nil, errors.New("error") } @@ -95,11 +96,11 @@ func (m *MockMediaFileRepo) GetWithParticipants(id string) (*model.MediaFile, er return nil, model.ErrNotFound } -func (m *MockMediaFileRepo) GetAllByTags(_ model.TagName, _ []string, options ...model.QueryOptions) (model.MediaFiles, error) { - return m.GetAll(options...) +func (m *MockMediaFileRepo) GetAllByTags(ctx context.Context, _ model.TagName, _ []string, options ...model.QueryOptions) (model.MediaFiles, error) { + return m.GetAll(ctx, options...) } -func (m *MockMediaFileRepo) GetAll(qo ...model.QueryOptions) (model.MediaFiles, error) { +func (m *MockMediaFileRepo) GetAll(_ context.Context, qo ...model.QueryOptions) (model.MediaFiles, error) { if len(qo) > 0 { m.Options = qo[0] } @@ -117,8 +118,8 @@ func (m *MockMediaFileRepo) GetAll(qo ...model.QueryOptions) (model.MediaFiles, return result, nil } -func (m *MockMediaFileRepo) GetRandom(qo ...model.QueryOptions) (model.MediaFiles, error) { - res, err := m.GetAll(qo...) +func (m *MockMediaFileRepo) GetRandom(ctx context.Context, qo ...model.QueryOptions) (model.MediaFiles, error) { + res, err := m.GetAll(ctx, qo...) if err != nil { return nil, err } @@ -128,8 +129,8 @@ func (m *MockMediaFileRepo) GetRandom(qo ...model.QueryOptions) (model.MediaFile return res, nil } -func (m *MockMediaFileRepo) GetCursor(qo ...model.QueryOptions) (model.MediaFileCursor, error) { - res, err := m.GetAll(qo...) +func (m *MockMediaFileRepo) GetCursor(ctx context.Context, qo ...model.QueryOptions) (model.MediaFileCursor, error) { + res, err := m.GetAll(ctx, qo...) if err != nil { return nil, err } @@ -142,11 +143,11 @@ func (m *MockMediaFileRepo) GetCursor(qo ...model.QueryOptions) (model.MediaFile }, nil } -func (m *MockMediaFileRepo) GetCursorWithArtwork(qo ...model.QueryOptions) (model.MediaFileCursor, error) { - return m.GetCursor(qo...) +func (m *MockMediaFileRepo) GetCursorWithArtwork(ctx context.Context, qo ...model.QueryOptions) (model.MediaFileCursor, error) { + return m.GetCursor(ctx, qo...) } -func (m *MockMediaFileRepo) Put(mf *model.MediaFile) error { +func (m *MockMediaFileRepo) Put(_ context.Context, mf *model.MediaFile) error { if m.Err { return errors.New("error") } @@ -157,7 +158,7 @@ func (m *MockMediaFileRepo) Put(mf *model.MediaFile) error { return nil } -func (m *MockMediaFileRepo) UpdateProbeData(id string, data string) error { +func (m *MockMediaFileRepo) UpdateProbeData(_ context.Context, id string, data string) error { if m.Err { return errors.New("error") } @@ -168,7 +169,7 @@ func (m *MockMediaFileRepo) UpdateProbeData(id string, data string) error { return model.ErrNotFound } -func (m *MockMediaFileRepo) Delete(id string) error { +func (m *MockMediaFileRepo) Delete(_ context.Context, id string) error { if m.Err { return errors.New("error") } @@ -179,7 +180,7 @@ func (m *MockMediaFileRepo) Delete(id string) error { return nil } -func (m *MockMediaFileRepo) ReassignReferences(prevID, newID string) error { +func (m *MockMediaFileRepo) ReassignReferences(_ context.Context, prevID, newID string) error { if m.Err { return errors.New("error") } @@ -190,7 +191,7 @@ func (m *MockMediaFileRepo) ReassignReferences(prevID, newID string) error { return nil } -func (m *MockMediaFileRepo) IncPlayCount(id string, timestamp time.Time) error { +func (m *MockMediaFileRepo) IncPlayCount(_ context.Context, id string, timestamp time.Time) error { if m.Err { return errors.New("error") } @@ -202,7 +203,7 @@ func (m *MockMediaFileRepo) IncPlayCount(id string, timestamp time.Time) error { return model.ErrNotFound } -func (m *MockMediaFileRepo) SetStar(starred bool, itemIDs ...string) error { +func (m *MockMediaFileRepo) SetStar(_ context.Context, starred bool, itemIDs ...string) error { if m.Err { return errors.New("error") } @@ -214,7 +215,7 @@ func (m *MockMediaFileRepo) SetStar(starred bool, itemIDs ...string) error { return nil } -func (m *MockMediaFileRepo) SetRating(rating int, itemID string) error { +func (m *MockMediaFileRepo) SetRating(_ context.Context, rating int, itemID string) error { if m.Err { return errors.New("error") } @@ -240,7 +241,7 @@ func (m *MockMediaFileRepo) FindByAlbum(artistId string) (model.MediaFiles, erro return res, nil } -func (m *MockMediaFileRepo) GetMissingAndMatching(libId int) (model.MediaFileCursor, error) { +func (m *MockMediaFileRepo) GetMissingAndMatching(_ context.Context, libId int) (model.MediaFileCursor, error) { if m.Err { return nil, errors.New("error") } @@ -274,7 +275,7 @@ func (m *MockMediaFileRepo) GetMissingAndMatching(libId int) (model.MediaFileCur }, nil } -func (m *MockMediaFileRepo) CountAll(opts ...model.QueryOptions) (int64, error) { +func (m *MockMediaFileRepo) CountAll(_ context.Context, opts ...model.QueryOptions) (int64, error) { if m.Err { return 0, errors.New("error") } @@ -287,7 +288,7 @@ func (m *MockMediaFileRepo) CountAll(opts ...model.QueryOptions) (int64, error) return int64(len(m.Data)), nil } -func (m *MockMediaFileRepo) DeleteAllMissing() (int64, error) { +func (m *MockMediaFileRepo) DeleteAllMissing(_ context.Context) (int64, error) { if m.Err { return 0, errors.New("error") } @@ -305,28 +306,20 @@ func (m *MockMediaFileRepo) DeleteAllMissing() (int64, error) { return count, nil } -// ResourceRepository methods -func (m *MockMediaFileRepo) Count(...rest.QueryOptions) (int64, error) { - return m.CountAll() +// REST repository methods +func (m *MockMediaFileRepo) Count(ctx context.Context, _ ...rest.QueryOptions) (int64, error) { + return m.CountAll(ctx) } -func (m *MockMediaFileRepo) Read(id string) (any, error) { - return m.Get(id) +func (m *MockMediaFileRepo) Read(ctx context.Context, id string) (*model.MediaFile, error) { + return m.Get(ctx, id) } -func (m *MockMediaFileRepo) ReadAll(...rest.QueryOptions) (any, error) { - return m.GetAll() +func (m *MockMediaFileRepo) ReadAll(ctx context.Context, _ ...rest.QueryOptions) ([]model.MediaFile, error) { + return m.GetAll(ctx) } -func (m *MockMediaFileRepo) EntityName() string { - return "mediafile" -} - -func (m *MockMediaFileRepo) NewInstance() any { - return &model.MediaFile{} -} - -func (m *MockMediaFileRepo) Search(q string, options ...model.QueryOptions) (model.MediaFiles, error) { +func (m *MockMediaFileRepo) Search(ctx context.Context, q string, options ...model.QueryOptions) (model.MediaFiles, error) { if len(options) > 0 { m.Options = options[0] } @@ -334,11 +327,11 @@ func (m *MockMediaFileRepo) Search(q string, options ...model.QueryOptions) (mod return nil, errors.New("unexpected error") } // Simple mock implementation - just return all media files for testing - return m.GetAll() + return m.GetAll(ctx) } // Cross-library move detection mock methods -func (m *MockMediaFileRepo) FindRecentFilesByMBZTrackID(missing model.MediaFile, since time.Time) (model.MediaFiles, error) { +func (m *MockMediaFileRepo) FindRecentFilesByMBZTrackID(_ context.Context, missing model.MediaFile, since time.Time) (model.MediaFiles, error) { if m.Err { return nil, errors.New("error") } @@ -360,7 +353,7 @@ func (m *MockMediaFileRepo) FindRecentFilesByMBZTrackID(missing model.MediaFile, return result, nil } -func (m *MockMediaFileRepo) FindRecentFilesByProperties(missing model.MediaFile, since time.Time) (model.MediaFiles, error) { +func (m *MockMediaFileRepo) FindRecentFilesByProperties(_ context.Context, missing model.MediaFile, since time.Time) (model.MediaFiles, error) { if m.Err { return nil, errors.New("error") } @@ -386,7 +379,7 @@ func (m *MockMediaFileRepo) FindRecentFilesByProperties(missing model.MediaFile, return result, nil } -func (m *MockMediaFileRepo) MatchesCriteria(string, criteria.Criteria) (bool, error) { +func (m *MockMediaFileRepo) MatchesCriteria(context.Context, string, criteria.Criteria) (bool, error) { if m.MatchesCriteriaErr != nil { return false, m.MatchesCriteriaErr } @@ -394,4 +387,4 @@ func (m *MockMediaFileRepo) MatchesCriteria(string, criteria.Criteria) (bool, er } var _ model.MediaFileRepository = (*MockMediaFileRepo)(nil) -var _ model.ResourceRepository = (*MockMediaFileRepo)(nil) +var _ rest.Repository[model.MediaFile] = (*MockMediaFileRepo)(nil) diff --git a/tests/mock_playlist_repo.go b/tests/mock_playlist_repo.go index 0fa9618ae..824e701f6 100644 --- a/tests/mock_playlist_repo.go +++ b/tests/mock_playlist_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "errors" "time" @@ -43,7 +44,7 @@ func (m *MockPlaylistRepo) SetData(playlists model.Playlists) { } } -func (m *MockPlaylistRepo) GetAll(options ...model.QueryOptions) (model.Playlists, error) { +func (m *MockPlaylistRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Playlists, error) { if len(options) > 0 { m.Options = options[0] } @@ -53,8 +54,8 @@ func (m *MockPlaylistRepo) GetAll(options ...model.QueryOptions) (model.Playlist return m.All, nil } -func (m *MockPlaylistRepo) GetCursor(options ...model.QueryOptions) (model.PlaylistCursor, error) { - res, err := m.GetAll(options...) +func (m *MockPlaylistRepo) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.PlaylistCursor, error) { + res, err := m.GetAll(ctx, options...) if err != nil { return nil, err } @@ -67,7 +68,7 @@ func (m *MockPlaylistRepo) GetCursor(options ...model.QueryOptions) (model.Playl }, nil } -func (m *MockPlaylistRepo) Get(id string) (*model.Playlist, error) { +func (m *MockPlaylistRepo) Get(_ context.Context, id string) (*model.Playlist, error) { if m.Err { return nil, errors.New("error") } @@ -79,11 +80,11 @@ func (m *MockPlaylistRepo) Get(id string) (*model.Playlist, error) { return nil, model.ErrNotFound } -func (m *MockPlaylistRepo) GetWithTracks(id string, _, _ bool) (*model.Playlist, error) { - return m.Get(id) +func (m *MockPlaylistRepo) GetWithTracks(ctx context.Context, id string, _, _ bool) (*model.Playlist, error) { + return m.Get(ctx, id) } -func (m *MockPlaylistRepo) Put(pls *model.Playlist, _ ...string) error { +func (m *MockPlaylistRepo) Put(_ context.Context, pls *model.Playlist, _ ...string) error { if m.Err { return errors.New("error") } @@ -97,7 +98,7 @@ func (m *MockPlaylistRepo) Put(pls *model.Playlist, _ ...string) error { return nil } -func (m *MockPlaylistRepo) FindByPath(path string) (*model.Playlist, error) { +func (m *MockPlaylistRepo) FindByPath(_ context.Context, path string) (*model.Playlist, error) { if m.Err { return nil, errors.New("error") } @@ -109,15 +110,15 @@ func (m *MockPlaylistRepo) FindByPath(path string) (*model.Playlist, error) { return nil, model.ErrNotFound } -func (m *MockPlaylistRepo) Delete(id string) error { +func (m *MockPlaylistRepo) Delete(_ context.Context, ids ...string) error { if m.Err { return errors.New("error") } - m.Deleted = append(m.Deleted, id) + m.Deleted = append(m.Deleted, ids...) return nil } -func (m *MockPlaylistRepo) SetStar(starred bool, ids ...string) error { +func (m *MockPlaylistRepo) SetStar(_ context.Context, starred bool, ids ...string) error { if m.Err { return errors.New("error") } @@ -130,7 +131,7 @@ func (m *MockPlaylistRepo) SetStar(starred bool, ids ...string) error { return nil } -func (m *MockPlaylistRepo) SetRating(rating int, id string) error { +func (m *MockPlaylistRepo) SetRating(_ context.Context, rating int, id string) error { if m.Err { return errors.New("error") } @@ -141,26 +142,26 @@ func (m *MockPlaylistRepo) SetRating(rating int, id string) error { return nil } -func (m *MockPlaylistRepo) IncPlayCount(string, time.Time) error { +func (m *MockPlaylistRepo) IncPlayCount(context.Context, string, time.Time) error { if m.Err { return errors.New("error") } return nil } -func (m *MockPlaylistRepo) ReassignAnnotation(string, string) error { +func (m *MockPlaylistRepo) ReassignAnnotation(context.Context, string, string) error { if m.Err { return errors.New("error") } return nil } -func (m *MockPlaylistRepo) Tracks(_ string, refreshSmartPlaylist bool) model.PlaylistTrackRepository { +func (m *MockPlaylistRepo) Tracks(_ context.Context, _ string, refreshSmartPlaylist bool) model.PlaylistTrackRepository { m.TracksRefreshed = refreshSmartPlaylist return m.TracksRepo } -func (m *MockPlaylistRepo) Exists(id string) (bool, error) { +func (m *MockPlaylistRepo) Exists(_ context.Context, id string) (bool, error) { if m.Err { return false, errors.New("error") } @@ -171,14 +172,14 @@ func (m *MockPlaylistRepo) Exists(id string) (bool, error) { return false, nil } -func (m *MockPlaylistRepo) Count(_ ...rest.QueryOptions) (int64, error) { +func (m *MockPlaylistRepo) Count(_ context.Context, _ ...rest.QueryOptions) (int64, error) { if m.Err { return 0, errors.New("error") } return int64(len(m.Data)), nil } -func (m *MockPlaylistRepo) CountAll(_ ...model.QueryOptions) (int64, error) { +func (m *MockPlaylistRepo) CountAll(_ context.Context, _ ...model.QueryOptions) (int64, error) { if m.Err { return 0, errors.New("error") } diff --git a/tests/mock_playlist_track_repo.go b/tests/mock_playlist_track_repo.go index 5666e7cc7..794393371 100644 --- a/tests/mock_playlist_track_repo.go +++ b/tests/mock_playlist_track_repo.go @@ -1,6 +1,8 @@ package tests import ( + "context" + "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils/slice" ) @@ -40,21 +42,21 @@ func (m *MockPlaylistTrackRepo) page(options ...model.QueryOptions) model.Playli return tracks } -func (m *MockPlaylistTrackRepo) CountAll(_ ...model.QueryOptions) (int64, error) { +func (m *MockPlaylistTrackRepo) CountAll(_ context.Context, _ ...model.QueryOptions) (int64, error) { if m.Err != nil { return 0, m.Err } return int64(len(m.Data)), nil } -func (m *MockPlaylistTrackRepo) GetAll(options ...model.QueryOptions) (model.PlaylistTracks, error) { +func (m *MockPlaylistTrackRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.PlaylistTracks, error) { if m.Err != nil { return nil, m.Err } return m.page(options...), nil } -func (m *MockPlaylistTrackRepo) GetCursor(options ...model.QueryOptions) (model.PlaylistTrackCursor, error) { +func (m *MockPlaylistTrackRepo) GetCursor(_ context.Context, options ...model.QueryOptions) (model.PlaylistTrackCursor, error) { if m.Err != nil { return nil, m.Err } @@ -68,21 +70,21 @@ func (m *MockPlaylistTrackRepo) GetCursor(options ...model.QueryOptions) (model. }, nil } -func (m *MockPlaylistTrackRepo) GetAlbumIDs(...model.QueryOptions) ([]string, error) { +func (m *MockPlaylistTrackRepo) GetAlbumIDs(context.Context, ...model.QueryOptions) ([]string, error) { if m.Err != nil { return nil, m.Err } return m.AlbumIDs, nil } -func (m *MockPlaylistTrackRepo) GetMediaFileIDs(options ...model.QueryOptions) ([]string, error) { +func (m *MockPlaylistTrackRepo) GetMediaFileIDs(_ context.Context, options ...model.QueryOptions) ([]string, error) { if m.Err != nil { return nil, m.Err } return slice.Map(m.page(options...), func(t model.PlaylistTrack) string { return t.MediaFileID }), nil } -func (m *MockPlaylistTrackRepo) Add(ids []string) (int, error) { +func (m *MockPlaylistTrackRepo) Add(_ context.Context, ids []string) (int, error) { m.AddedIds = append(m.AddedIds, ids...) if m.Err != nil { return 0, m.Err @@ -90,38 +92,38 @@ func (m *MockPlaylistTrackRepo) Add(ids []string) (int, error) { return m.AddCount, nil } -func (m *MockPlaylistTrackRepo) Insert(ids []string, pos int) (int, error) { +func (m *MockPlaylistTrackRepo) Insert(ctx context.Context, ids []string, pos int) (int, error) { m.InsertPos = pos - return m.Add(ids) + return m.Add(ctx, ids) } -func (m *MockPlaylistTrackRepo) AddAlbums(_ []string) (int, error) { +func (m *MockPlaylistTrackRepo) AddAlbums(_ context.Context, _ []string) (int, error) { if m.Err != nil { return 0, m.Err } return m.AddCount, nil } -func (m *MockPlaylistTrackRepo) AddArtists(_ []string) (int, error) { +func (m *MockPlaylistTrackRepo) AddArtists(_ context.Context, _ []string) (int, error) { if m.Err != nil { return 0, m.Err } return m.AddCount, nil } -func (m *MockPlaylistTrackRepo) AddDiscs(_ []model.DiscID) (int, error) { +func (m *MockPlaylistTrackRepo) AddDiscs(_ context.Context, _ []model.DiscID) (int, error) { if m.Err != nil { return 0, m.Err } return m.AddCount, nil } -func (m *MockPlaylistTrackRepo) Delete(ids ...string) error { +func (m *MockPlaylistTrackRepo) Delete(_ context.Context, ids ...string) error { m.DeletedIds = append(m.DeletedIds, ids...) return m.Err } -func (m *MockPlaylistTrackRepo) Reorder(_, _ int) error { +func (m *MockPlaylistTrackRepo) Reorder(_ context.Context, _, _ int) error { m.Reordered = true return m.Err } diff --git a/tests/mock_playqueue_repo.go b/tests/mock_playqueue_repo.go index 19976db57..b445a5085 100644 --- a/tests/mock_playqueue_repo.go +++ b/tests/mock_playqueue_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "errors" "github.com/navidrome/navidrome/model" @@ -13,7 +14,7 @@ type MockPlayQueueRepo struct { LastCols []string } -func (m *MockPlayQueueRepo) Store(q *model.PlayQueue, cols ...string) error { +func (m *MockPlayQueueRepo) Store(_ context.Context, q *model.PlayQueue, cols ...string) error { if m.Err { return errors.New("error") } @@ -26,7 +27,7 @@ func (m *MockPlayQueueRepo) Store(q *model.PlayQueue, cols ...string) error { return nil } -func (m *MockPlayQueueRepo) RetrieveWithMediaFiles(userId string) (*model.PlayQueue, error) { +func (m *MockPlayQueueRepo) RetrieveWithMediaFiles(_ context.Context, userId string) (*model.PlayQueue, error) { if m.Err { return nil, errors.New("error") } @@ -40,7 +41,7 @@ func (m *MockPlayQueueRepo) RetrieveWithMediaFiles(userId string) (*model.PlayQu return &qCopy, nil } -func (m *MockPlayQueueRepo) Retrieve(userId string) (*model.PlayQueue, error) { +func (m *MockPlayQueueRepo) Retrieve(_ context.Context, userId string) (*model.PlayQueue, error) { if m.Err { return nil, errors.New("error") } @@ -56,7 +57,7 @@ func (m *MockPlayQueueRepo) Retrieve(userId string) (*model.PlayQueue, error) { return &qCopy, nil } -func (m *MockPlayQueueRepo) Clear(userId string) error { +func (m *MockPlayQueueRepo) Clear(_ context.Context, userId string) error { if m.Err { return errors.New("error") } diff --git a/tests/mock_plugin_repo.go b/tests/mock_plugin_repo.go index e65d56def..feaf9b212 100644 --- a/tests/mock_plugin_repo.go +++ b/tests/mock_plugin_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "errors" "time" @@ -29,7 +30,7 @@ func (m *MockPluginRepo) SetError(err bool) { m.Err = err } -func (m *MockPluginRepo) ClearErrors() error { +func (m *MockPluginRepo) ClearErrors(context.Context) error { if m.Err { return errors.New("unexpected error") } @@ -55,7 +56,7 @@ func (m *MockPluginRepo) SetPermitted(permitted bool) { m.Permitted = permitted } -func (m *MockPluginRepo) Get(id string) (*model.Plugin, error) { +func (m *MockPluginRepo) Get(_ context.Context, id string) (*model.Plugin, error) { if !m.Permitted { return nil, rest.ErrPermissionDenied } @@ -68,11 +69,11 @@ func (m *MockPluginRepo) Get(id string) (*model.Plugin, error) { return nil, model.ErrNotFound } -func (m *MockPluginRepo) Read(id string) (any, error) { - return m.Get(id) +func (m *MockPluginRepo) Read(ctx context.Context, id string) (*model.Plugin, error) { + return m.Get(ctx, id) } -func (m *MockPluginRepo) Put(p *model.Plugin) error { +func (m *MockPluginRepo) Put(_ context.Context, p *model.Plugin) error { if !m.Permitted { return rest.ErrPermissionDenied } @@ -105,7 +106,7 @@ func (m *MockPluginRepo) Put(p *model.Plugin) error { return nil } -func (m *MockPluginRepo) Delete(id string) error { +func (m *MockPluginRepo) Delete(_ context.Context, id string) error { if !m.Permitted { return rest.ErrPermissionDenied } @@ -123,7 +124,7 @@ func (m *MockPluginRepo) Delete(id string) error { return nil } -func (m *MockPluginRepo) GetAll(qo ...model.QueryOptions) (model.Plugins, error) { +func (m *MockPluginRepo) GetAll(_ context.Context, qo ...model.QueryOptions) (model.Plugins, error) { if len(qo) > 0 { m.Options = qo[0] } @@ -136,7 +137,7 @@ func (m *MockPluginRepo) GetAll(qo ...model.QueryOptions) (model.Plugins, error) return m.All, nil } -func (m *MockPluginRepo) CountAll(qo ...model.QueryOptions) (int64, error) { +func (m *MockPluginRepo) CountAll(_ context.Context, qo ...model.QueryOptions) (int64, error) { if len(qo) > 0 { m.Options = qo[0] } @@ -149,36 +150,16 @@ func (m *MockPluginRepo) CountAll(qo ...model.QueryOptions) (int64, error) { return int64(len(m.All)), nil } -// rest.Repository interface methods -func (m *MockPluginRepo) Count(options ...rest.QueryOptions) (int64, error) { +// REST repository methods +func (m *MockPluginRepo) Count(_ context.Context, _ ...rest.QueryOptions) (int64, error) { if !m.Permitted { return 0, rest.ErrPermissionDenied } return int64(len(m.All)), nil } -func (m *MockPluginRepo) EntityName() string { - return "plugin" -} - -func (m *MockPluginRepo) NewInstance() any { - return &model.Plugin{} -} - -func (m *MockPluginRepo) ReadAll(options ...rest.QueryOptions) (any, error) { - return m.GetAll() -} - -func (m *MockPluginRepo) Save(entity any) (string, error) { - p := entity.(*model.Plugin) - err := m.Put(p) - return p.ID, err -} - -func (m *MockPluginRepo) Update(id string, entity any, cols ...string) error { - p := entity.(*model.Plugin) - p.ID = id - return m.Put(p) +func (m *MockPluginRepo) ReadAll(ctx context.Context, _ ...rest.QueryOptions) ([]model.Plugin, error) { + return m.GetAll(ctx) } var _ model.PluginRepository = (*MockPluginRepo)(nil) diff --git a/tests/mock_property_repo.go b/tests/mock_property_repo.go index 9adc66e6d..949f894c1 100644 --- a/tests/mock_property_repo.go +++ b/tests/mock_property_repo.go @@ -1,6 +1,10 @@ package tests -import "github.com/navidrome/navidrome/model" +import ( + "context" + + "github.com/navidrome/navidrome/model" +) type MockedPropertyRepo struct { model.PropertyRepository @@ -14,7 +18,7 @@ func (p *MockedPropertyRepo) init() { } } -func (p *MockedPropertyRepo) Put(id string, value string) error { +func (p *MockedPropertyRepo) Put(_ context.Context, id string, value string) error { if p.Error != nil { return p.Error } @@ -23,7 +27,7 @@ func (p *MockedPropertyRepo) Put(id string, value string) error { return nil } -func (p *MockedPropertyRepo) Get(id string) (string, error) { +func (p *MockedPropertyRepo) Get(_ context.Context, id string) (string, error) { if p.Error != nil { return "", p.Error } @@ -34,7 +38,7 @@ func (p *MockedPropertyRepo) Get(id string) (string, error) { return "", model.ErrNotFound } -func (p *MockedPropertyRepo) Delete(id string) error { +func (p *MockedPropertyRepo) Delete(_ context.Context, id string) error { if p.Error != nil { return p.Error } @@ -46,12 +50,12 @@ func (p *MockedPropertyRepo) Delete(id string) error { return model.ErrNotFound } -func (p *MockedPropertyRepo) DefaultGet(id string, defaultValue string) (string, error) { +func (p *MockedPropertyRepo) DefaultGet(ctx context.Context, id string, defaultValue string) (string, error) { if p.Error != nil { return "", p.Error } p.init() - v, err := p.Get(id) + v, err := p.Get(ctx, id) if err != nil { return defaultValue, nil //nolint:nilerr } diff --git a/tests/mock_radio_repository.go b/tests/mock_radio_repository.go index 20f81ec45..21898ea53 100644 --- a/tests/mock_radio_repository.go +++ b/tests/mock_radio_repository.go @@ -1,6 +1,7 @@ package tests import ( + "context" "errors" "github.com/navidrome/navidrome/model" @@ -23,29 +24,27 @@ func (m *MockedRadioRepo) SetError(err bool) { m.Err = err } -func (m *MockedRadioRepo) CountAll(options ...model.QueryOptions) (int64, error) { +func (m *MockedRadioRepo) CountAll(_ context.Context, options ...model.QueryOptions) (int64, error) { if m.Err { return 0, errors.New("error") } return int64(len(m.Data)), nil } -func (m *MockedRadioRepo) Delete(id string) error { +func (m *MockedRadioRepo) Delete(_ context.Context, ids ...string) error { if m.Err { return errors.New("Error!") } - - _, found := m.Data[id] - - if !found { - return errors.New("not found") + for _, id := range ids { + if _, found := m.Data[id]; !found { + return errors.New("not found") + } + delete(m.Data, id) } - - delete(m.Data, id) return nil } -func (m *MockedRadioRepo) Exists(id string) (bool, error) { +func (m *MockedRadioRepo) Exists(_ context.Context, id string) (bool, error) { if m.Err { return false, errors.New("Error!") } @@ -53,7 +52,7 @@ func (m *MockedRadioRepo) Exists(id string) (bool, error) { return found, nil } -func (m *MockedRadioRepo) Get(id string) (*model.Radio, error) { +func (m *MockedRadioRepo) Get(_ context.Context, id string) (*model.Radio, error) { if m.Err { return nil, errors.New("Error!") } @@ -63,7 +62,7 @@ func (m *MockedRadioRepo) Get(id string) (*model.Radio, error) { return nil, model.ErrNotFound } -func (m *MockedRadioRepo) GetAll(qo ...model.QueryOptions) (model.Radios, error) { +func (m *MockedRadioRepo) GetAll(_ context.Context, qo ...model.QueryOptions) (model.Radios, error) { if len(qo) > 0 { m.Options = qo[0] } @@ -73,7 +72,7 @@ func (m *MockedRadioRepo) GetAll(qo ...model.QueryOptions) (model.Radios, error) return m.All, nil } -func (m *MockedRadioRepo) Put(radio *model.Radio, _ ...string) error { +func (m *MockedRadioRepo) Put(_ context.Context, radio *model.Radio, _ ...string) error { if m.Err { return errors.New("error") } diff --git a/tests/mock_scrobble_buffer_repo.go b/tests/mock_scrobble_buffer_repo.go index 2eb5e8a93..91177eeef 100644 --- a/tests/mock_scrobble_buffer_repo.go +++ b/tests/mock_scrobble_buffer_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "sync" "time" @@ -17,7 +18,7 @@ func CreateMockedScrobbleBufferRepo() *MockedScrobbleBufferRepo { return &MockedScrobbleBufferRepo{} } -func (m *MockedScrobbleBufferRepo) UserIDs(service string) ([]string, error) { +func (m *MockedScrobbleBufferRepo) UserIDs(_ context.Context, service string) ([]string, error) { if m.Error != nil { return nil, m.Error } @@ -36,7 +37,7 @@ func (m *MockedScrobbleBufferRepo) UserIDs(service string) ([]string, error) { return result, nil } -func (m *MockedScrobbleBufferRepo) Enqueue(service, userId, mediaFileId string, playTime time.Time) error { +func (m *MockedScrobbleBufferRepo) Enqueue(_ context.Context, service, userId, mediaFileId string, playTime time.Time) error { if m.Error != nil { return m.Error } @@ -52,7 +53,7 @@ func (m *MockedScrobbleBufferRepo) Enqueue(service, userId, mediaFileId string, return nil } -func (m *MockedScrobbleBufferRepo) Next(service, userId string) (*model.ScrobbleEntry, error) { +func (m *MockedScrobbleBufferRepo) Next(_ context.Context, service, userId string) (*model.ScrobbleEntry, error) { if m.Error != nil { return nil, m.Error } @@ -66,7 +67,7 @@ func (m *MockedScrobbleBufferRepo) Next(service, userId string) (*model.Scrobble return nil, nil } -func (m *MockedScrobbleBufferRepo) Dequeue(entry *model.ScrobbleEntry) error { +func (m *MockedScrobbleBufferRepo) Dequeue(_ context.Context, entry *model.ScrobbleEntry) error { if m.Error != nil { return m.Error } @@ -83,7 +84,7 @@ func (m *MockedScrobbleBufferRepo) Dequeue(entry *model.ScrobbleEntry) error { return nil } -func (m *MockedScrobbleBufferRepo) Discard(service string) error { +func (m *MockedScrobbleBufferRepo) Discard(_ context.Context, service string) error { if m.Error != nil { return m.Error } @@ -99,7 +100,7 @@ func (m *MockedScrobbleBufferRepo) Discard(service string) error { return nil } -func (m *MockedScrobbleBufferRepo) Length() (int64, error) { +func (m *MockedScrobbleBufferRepo) Length(context.Context) (int64, error) { if m.Error != nil { return 0, m.Error } @@ -107,3 +108,5 @@ func (m *MockedScrobbleBufferRepo) Length() (int64, error) { defer m.mu.RUnlock() return int64(len(m.Data)), nil } + +var _ model.ScrobbleBufferRepository = (*MockedScrobbleBufferRepo)(nil) diff --git a/tests/mock_scrobble_repo.go b/tests/mock_scrobble_repo.go index d6d88d221..a76bc5d59 100644 --- a/tests/mock_scrobble_repo.go +++ b/tests/mock_scrobble_repo.go @@ -5,16 +5,16 @@ import ( "strconv" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" ) type MockScrobbleRepo struct { RecordedScrobbles []model.Scrobble - ctx context.Context } -func (m *MockScrobbleRepo) Get(id string) (*model.Scrobble, error) { +func (m *MockScrobbleRepo) Get(_ context.Context, id string) (*model.Scrobble, error) { for idx := range m.RecordedScrobbles { if strconv.FormatInt(m.RecordedScrobbles[idx].ID, 10) == id { return &m.RecordedScrobbles[idx], nil @@ -24,16 +24,16 @@ func (m *MockScrobbleRepo) Get(id string) (*model.Scrobble, error) { return nil, model.ErrNotFound } -func (m *MockScrobbleRepo) GetAll(options ...model.QueryOptions) (model.Scrobbles, error) { +func (m *MockScrobbleRepo) GetAll(_ context.Context, _ ...model.QueryOptions) (model.Scrobbles, error) { return m.RecordedScrobbles, nil } -func (m *MockScrobbleRepo) CountAll(options ...model.QueryOptions) (int64, error) { +func (m *MockScrobbleRepo) CountAll(_ context.Context, _ ...model.QueryOptions) (int64, error) { return int64(len(m.RecordedScrobbles)), nil } -func (m *MockScrobbleRepo) RecordScrobble(fileID string, submissionTime time.Time) error { - user, _ := request.UserFrom(m.ctx) +func (m *MockScrobbleRepo) RecordScrobble(ctx context.Context, fileID string, submissionTime time.Time) error { + user, _ := request.UserFrom(ctx) m.RecordedScrobbles = append(m.RecordedScrobbles, model.Scrobble{ MediaFileID: fileID, UserID: user.ID, @@ -42,4 +42,16 @@ func (m *MockScrobbleRepo) RecordScrobble(fileID string, submissionTime time.Tim return nil } +func (m *MockScrobbleRepo) Count(ctx context.Context, _ ...rest.QueryOptions) (int64, error) { + return m.CountAll(ctx) +} + +func (m *MockScrobbleRepo) Read(ctx context.Context, id string) (*model.Scrobble, error) { + return m.Get(ctx, id) +} + +func (m *MockScrobbleRepo) ReadAll(ctx context.Context, _ ...rest.QueryOptions) ([]model.Scrobble, error) { + return m.GetAll(ctx) +} + var _ model.ScrobbleRepository = (*MockScrobbleRepo)(nil) diff --git a/tests/mock_share_repo.go b/tests/mock_share_repo.go index 9fbf0057e..0c872dc0d 100644 --- a/tests/mock_share_repo.go +++ b/tests/mock_share_repo.go @@ -1,14 +1,13 @@ package tests import ( - "github.com/deluan/rest" + "context" + "github.com/navidrome/navidrome/model" ) type MockShareRepo struct { model.ShareRepository - rest.Repository - rest.Persistable Entity any ID string @@ -16,11 +15,10 @@ type MockShareRepo struct { Error error } -func (m *MockShareRepo) Save(entity any) (string, error) { +func (m *MockShareRepo) Save(_ context.Context, s *model.Share) (string, error) { if m.Error != nil { return "", m.Error } - s := entity.(*model.Share) if s.ID == "" { s.ID = "id" } @@ -28,24 +26,24 @@ func (m *MockShareRepo) Save(entity any) (string, error) { return s.ID, nil } -func (m *MockShareRepo) Update(id string, entity any, cols ...string) error { +func (m *MockShareRepo) Update(_ context.Context, id string, entity model.Share, cols ...string) error { if m.Error != nil { return m.Error } m.ID = id - m.Entity = entity + m.Entity = &entity m.Cols = cols return nil } -func (m *MockShareRepo) Exists(id string) (bool, error) { +func (m *MockShareRepo) Exists(_ context.Context, id string) (bool, error) { if m.Error != nil { return false, m.Error } return id == m.ID, nil } -func (m *MockShareRepo) Get(id string) (*model.Share, error) { +func (m *MockShareRepo) Get(_ context.Context, id string) (*model.Share, error) { if m.Error != nil { return nil, m.Error } diff --git a/tests/mock_tag_repo.go b/tests/mock_tag_repo.go index a59035ea6..f1f252efe 100644 --- a/tests/mock_tag_repo.go +++ b/tests/mock_tag_repo.go @@ -1,6 +1,8 @@ package tests import ( + "context" + "github.com/navidrome/navidrome/model" ) @@ -13,7 +15,7 @@ type MockTagRepo struct { Err error } -func (r *MockTagRepo) GetAll(_ model.TagName, options ...model.QueryOptions) (model.TagList, error) { +func (r *MockTagRepo) GetAll(_ context.Context, _ model.TagName, options ...model.QueryOptions) (model.TagList, error) { if len(options) > 0 { r.Options = options[0] } diff --git a/tests/mock_transcoding_repo.go b/tests/mock_transcoding_repo.go index 641daca8a..52eb16eed 100644 --- a/tests/mock_transcoding_repo.go +++ b/tests/mock_transcoding_repo.go @@ -1,16 +1,20 @@ package tests -import "github.com/navidrome/navidrome/model" +import ( + "context" + + "github.com/navidrome/navidrome/model" +) type MockTranscodingRepo struct { model.TranscodingRepository } -func (m *MockTranscodingRepo) Get(id string) (*model.Transcoding, error) { +func (m *MockTranscodingRepo) Get(_ context.Context, id string) (*model.Transcoding, error) { return &model.Transcoding{ID: id, TargetFormat: "mp3", DefaultBitRate: 160}, nil } -func (m *MockTranscodingRepo) FindByFormat(format string) (*model.Transcoding, error) { +func (m *MockTranscodingRepo) FindByFormat(_ context.Context, format string) (*model.Transcoding, error) { switch format { case "mp3": return &model.Transcoding{ID: "mp31", TargetFormat: "mp3", DefaultBitRate: 160}, nil diff --git a/tests/mock_user_props_repo.go b/tests/mock_user_props_repo.go index 1b1e17650..278a68431 100644 --- a/tests/mock_user_props_repo.go +++ b/tests/mock_user_props_repo.go @@ -1,6 +1,10 @@ package tests -import "github.com/navidrome/navidrome/model" +import ( + "context" + + "github.com/navidrome/navidrome/model" +) type MockedUserPropsRepo struct { model.UserPropsRepository @@ -14,7 +18,7 @@ func (p *MockedUserPropsRepo) init() { } } -func (p *MockedUserPropsRepo) Put(userId, key string, value string) error { +func (p *MockedUserPropsRepo) Put(_ context.Context, userId, key string, value string) error { if p.Error != nil { return p.Error } @@ -23,7 +27,7 @@ func (p *MockedUserPropsRepo) Put(userId, key string, value string) error { return nil } -func (p *MockedUserPropsRepo) Get(userId, key string) (string, error) { +func (p *MockedUserPropsRepo) Get(_ context.Context, userId, key string) (string, error) { if p.Error != nil { return "", p.Error } @@ -34,7 +38,7 @@ func (p *MockedUserPropsRepo) Get(userId, key string) (string, error) { return "", model.ErrNotFound } -func (p *MockedUserPropsRepo) Delete(userId, key string) error { +func (p *MockedUserPropsRepo) Delete(_ context.Context, userId, key string) error { if p.Error != nil { return p.Error } @@ -46,12 +50,12 @@ func (p *MockedUserPropsRepo) Delete(userId, key string) error { return model.ErrNotFound } -func (p *MockedUserPropsRepo) DefaultGet(userId, key string, defaultValue string) (string, error) { +func (p *MockedUserPropsRepo) DefaultGet(ctx context.Context, userId, key string, defaultValue string) (string, error) { if p.Error != nil { return "", p.Error } p.init() - v, err := p.Get(userId, key) + v, err := p.Get(ctx, userId, key) if err != nil { return defaultValue, nil //nolint:nilerr } diff --git a/tests/mock_user_repo.go b/tests/mock_user_repo.go index 2d6ff3c02..58c985157 100644 --- a/tests/mock_user_repo.go +++ b/tests/mock_user_repo.go @@ -1,6 +1,7 @@ package tests import ( + "context" "encoding/base64" "fmt" "strings" @@ -23,14 +24,14 @@ type MockedUserRepo struct { UserLibraries map[string][]int // userID -> libraryIDs } -func (u *MockedUserRepo) CountAll(qo ...model.QueryOptions) (int64, error) { +func (u *MockedUserRepo) CountAll(_ context.Context, qo ...model.QueryOptions) (int64, error) { if u.Error != nil { return 0, u.Error } return int64(len(u.Data)), nil } -func (u *MockedUserRepo) Put(usr *model.User) error { +func (u *MockedUserRepo) Put(_ context.Context, usr *model.User) error { if u.Error != nil { return u.Error } @@ -42,7 +43,7 @@ func (u *MockedUserRepo) Put(usr *model.User) error { return nil } -func (u *MockedUserRepo) FindByUsername(username string) (*model.User, error) { +func (u *MockedUserRepo) FindByUsername(_ context.Context, username string) (*model.User, error) { if u.Error != nil { return nil, u.Error } @@ -53,11 +54,11 @@ func (u *MockedUserRepo) FindByUsername(username string) (*model.User, error) { return usr, nil } -func (u *MockedUserRepo) FindByUsernameWithPassword(username string) (*model.User, error) { - return u.FindByUsername(username) +func (u *MockedUserRepo) FindByUsernameWithPassword(ctx context.Context, username string) (*model.User, error) { + return u.FindByUsername(ctx, username) } -func (u *MockedUserRepo) FindFirstAdmin() (*model.User, error) { +func (u *MockedUserRepo) FindFirstAdmin(_ context.Context) (*model.User, error) { if u.Error != nil { return nil, u.Error } @@ -69,7 +70,7 @@ func (u *MockedUserRepo) FindFirstAdmin() (*model.User, error) { return nil, model.ErrNotFound } -func (u *MockedUserRepo) Get(id string) (*model.User, error) { +func (u *MockedUserRepo) Get(_ context.Context, id string) (*model.User, error) { if u.Error != nil { return nil, u.Error } @@ -81,7 +82,7 @@ func (u *MockedUserRepo) Get(id string) (*model.User, error) { return nil, model.ErrNotFound } -func (u *MockedUserRepo) GetAll(options ...model.QueryOptions) (model.Users, error) { +func (u *MockedUserRepo) GetAll(_ context.Context, options ...model.QueryOptions) (model.Users, error) { if u.Error != nil { return nil, u.Error } @@ -92,7 +93,7 @@ func (u *MockedUserRepo) GetAll(options ...model.QueryOptions) (model.Users, err return users, nil } -func (u *MockedUserRepo) UpdateLastLoginAt(id string) error { +func (u *MockedUserRepo) UpdateLastLoginAt(_ context.Context, id string) error { for _, usr := range u.Data { if usr.ID == id { usr.LastLoginAt = new(time.Now()) @@ -102,7 +103,7 @@ func (u *MockedUserRepo) UpdateLastLoginAt(id string) error { return u.Error } -func (u *MockedUserRepo) UpdateLastAccessAt(id string) error { +func (u *MockedUserRepo) UpdateLastAccessAt(_ context.Context, id string) error { for _, usr := range u.Data { if usr.ID == id { usr.LastAccessAt = new(time.Now()) @@ -114,7 +115,7 @@ func (u *MockedUserRepo) UpdateLastAccessAt(id string) error { // Library association methods - mock implementations -func (u *MockedUserRepo) GetUserLibraries(userID string) (model.Libraries, error) { +func (u *MockedUserRepo) GetUserLibraries(_ context.Context, userID string) (model.Libraries, error) { if u.Error != nil { return nil, u.Error } @@ -135,7 +136,7 @@ func (u *MockedUserRepo) GetUserLibraries(userID string) (model.Libraries, error return libraries, nil } -func (u *MockedUserRepo) SetUserLibraries(userID string, libraryIDs []int) error { +func (u *MockedUserRepo) SetUserLibraries(_ context.Context, userID string, libraryIDs []int) error { if u.Error != nil { return u.Error } @@ -146,10 +147,19 @@ func (u *MockedUserRepo) SetUserLibraries(userID string, libraryIDs []int) error return nil } -func (u *MockedUserRepo) Delete(id string) error { +func (u *MockedUserRepo) Delete(_ context.Context, ids ...string) error { if u.Error != nil { return u.Error } + for _, id := range ids { + if err := u.deleteOne(id); err != nil { + return err + } + } + return nil +} + +func (u *MockedUserRepo) deleteOne(id string) error { for key, usr := range u.Data { if usr.ID == id { delete(u.Data, key) @@ -160,19 +170,17 @@ func (u *MockedUserRepo) Delete(id string) error { return model.ErrNotFound } -func (u *MockedUserRepo) Save(entity any) (string, error) { - usr := entity.(*model.User) - if err := u.Put(usr); err != nil { +func (u *MockedUserRepo) Save(ctx context.Context, usr *model.User) (string, error) { + if err := u.Put(ctx, usr); err != nil { return "", err } return usr.ID, nil } -func (u *MockedUserRepo) Update(id string, entity any, cols ...string) error { +func (u *MockedUserRepo) Update(ctx context.Context, id string, entity model.User, _ ...string) error { if u.Error != nil { return u.Error } - usr := entity.(*model.User) - usr.ID = id - return u.Put(usr) + entity.ID = id + return u.Put(ctx, &entity) } diff --git a/tests/mock_user_service.go b/tests/mock_user_service.go index f2700de45..bde843d1a 100644 --- a/tests/mock_user_service.go +++ b/tests/mock_user_service.go @@ -1,9 +1,8 @@ package tests import ( - "context" - "github.com/deluan/rest" + "github.com/navidrome/navidrome/model" ) // MockUserService provides a simple wrapper around MockedUserRepo @@ -13,7 +12,7 @@ type MockUserService struct { *MockedUserRepo } -// MockUserRestAdapter adapts MockedUserRepo to rest.Repository interface +// MockUserRestAdapter adapts MockedUserRepo to the REST repository interface type MockUserRestAdapter struct { *MockedUserRepo } @@ -25,6 +24,6 @@ func NewMockUserService() *MockUserService { return &MockUserService{MockedUserRepo: repo} } -func (m *MockUserService) NewRepository(ctx context.Context) rest.Repository { +func (m *MockUserService) Repository() rest.Repository[model.User] { return &MockUserRestAdapter{MockedUserRepo: m.MockedUserRepo} }