From 961ee8c4136fbd18b78de1929ee5f97a1ae37f7c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Tue, 22 Sep 2026 19:51:32 -0400 Subject: [PATCH 01/26] fix(scanner): keep tag numbers within the int32 range (#6202) * fix(scanner): keep tag numbers within the int32 range A track number of 4294967295 (-1 stored as an unsigned 32-bit tag) was saved as-is by 64-bit builds. 32-bit builds (armv5/6/7, 386) cannot read that value back into an int, so every scan failed with "converting driver.Value type int64 to a int: value out of range" when loading the folder's media files. Track and disc numbers (and their totals) are now parsed as int32 and fall back to 0 when out of range, matching how unparseable values are handled. BPM values outside the int32 range are dropped. A migration resets existing out-of-range track_number, disc_number and bpm values, and removes out-of-range keys from album.discs, so databases written by 64-bit builds are readable again by 32-bit ones. Persistent IDs are unaffected because they use the raw tag text. Fixes #6200 * fix(scanner): accept the int32 minimum as a BPM value The BPM range check compared the absolute value against MaxInt32, which rejected -2147483648 even though it fits in an int32. Compare against MinInt32 and MaxInt32 separately, matching atoi32 and the migration. * fix(scanner): treat negative track, disc and BPM values as missing Track numbers, disc numbers and BPM can never be negative, so negative tag values now map to 0 (track/disc, including totals) or nil (BPM), the same as unparseable ones. The migration resets existing negative values as well as the ones above the int32 range, and keeps only album disc keys from 0 to MaxInt32. --- ...428_clamp_media_file_int32_tag_numbers.sql | 24 +++++++++++++++++++ model/metadata/map_mediafile.go | 4 ++-- model/metadata/map_mediafile_test.go | 9 +++++++ model/metadata/metadata.go | 19 +++++++++------ model/metadata/metadata_test.go | 4 ++++ 5 files changed, 51 insertions(+), 9 deletions(-) create mode 100644 db/migrations/20260922230428_clamp_media_file_int32_tag_numbers.sql diff --git a/db/migrations/20260922230428_clamp_media_file_int32_tag_numbers.sql b/db/migrations/20260922230428_clamp_media_file_int32_tag_numbers.sql new file mode 100644 index 000000000..dced5d744 --- /dev/null +++ b/db/migrations/20260922230428_clamp_media_file_int32_tag_numbers.sql @@ -0,0 +1,24 @@ +-- +goose Up +-- +goose StatementBegin +-- 32-bit builds cannot read values above the int32 range written by 64-bit builds. +update media_file set track_number = 0 +where track_number < 0 or track_number > 2147483647; + +update media_file set disc_number = 0 +where disc_number < 0 or disc_number > 2147483647; + +update media_file set bpm = null +where bpm < 0 or bpm > 2147483647; + +update album set discs = ( + select json_group_object(key, value) from json_each(album.discs) + where cast(key as integer) between 0 and 2147483647 +) +where json_valid(discs) and exists ( + select 1 from json_each(album.discs) + where cast(key as integer) not between 0 and 2147483647 +); +-- +goose StatementEnd + +-- +goose Down +SELECT 1; diff --git a/model/metadata/map_mediafile.go b/model/metadata/map_mediafile.go index b3ce4ef02..6d12feba9 100644 --- a/model/metadata/map_mediafile.go +++ b/model/metadata/map_mediafile.go @@ -37,8 +37,8 @@ func (md Metadata) ToMediaFile(libID int, folderID string) model.MediaFile { mf.CatalogNum = md.String(model.TagCatalogNumber) mf.Comment = md.String(model.TagComment) if f := md.NullableFloat(model.TagBPM); f != nil { - if v := int(math.Round(*f)); v != 0 { - mf.BPM = new(v) + if r := math.Round(*f); r > 0 && r <= math.MaxInt32 { + mf.BPM = new(int(r)) } } mf.Lyrics = md.mapLyrics() diff --git a/model/metadata/map_mediafile_test.go b/model/metadata/map_mediafile_test.go index 75a7ed358..baaf8fab5 100644 --- a/model/metadata/map_mediafile_test.go +++ b/model/metadata/map_mediafile_test.go @@ -131,6 +131,15 @@ var _ = Describe("ToMediaFile", func() { Expect(toMediaFile(model.RawTags{"BPM": {"0"}}).BPM).To(BeNil()) Expect(toMediaFile(model.RawTags{"BPM": {"fast"}}).BPM).To(BeNil()) }) + It("leaves BPM nil when the tag does not fit in 32 bits", func() { + Expect(toMediaFile(model.RawTags{"BPM": {"4294967295"}}).BPM).To(BeNil()) + }) + It("leaves BPM nil when the tag is negative", func() { + Expect(toMediaFile(model.RawTags{"BPM": {"-120"}}).BPM).To(BeNil()) + }) + It("keeps the largest 32-bit BPM value", func() { + Expect(toMediaFile(model.RawTags{"BPM": {"2147483647"}}).BPM).To(Equal(new(2147483647))) + }) }) Describe("BitDepth", func() { diff --git a/model/metadata/metadata.go b/model/metadata/metadata.go index 0efbe94ec..7843e7010 100644 --- a/model/metadata/metadata.go +++ b/model/metadata/metadata.go @@ -148,15 +148,20 @@ func (md Metadata) tuple(key model.TagName) (int, int) { return 0, 0 } tuple := strings.Split(tag, "/") - t1, t2 := 0, 0 - t1, _ = strconv.Atoi(tuple[0]) + total := md.first(key + "total") if len(tuple) > 1 { - t2, _ = strconv.Atoi(tuple[1]) - } else { - t2tag := md.first(key + "total") - t2, _ = strconv.Atoi(t2tag) + total = tuple[1] } - return t1, t2 + return tagNumber(tuple[0]), tagNumber(total) +} + +// tagNumber rejects negatives and values above int32, so the DB stays readable by 32-bit builds. +func tagNumber(s string) int { + v, err := strconv.ParseInt(s, 10, 32) + if err != nil || v < 0 { + return 0 + } + return int(v) } var dateRegex = regexp.MustCompile(`([12]\d\d\d)`) diff --git a/model/metadata/metadata_test.go b/model/metadata/metadata_test.go index 09a2dfde0..c84d93981 100644 --- a/model/metadata/metadata_test.go +++ b/model/metadata/metadata_test.go @@ -226,6 +226,10 @@ var _ = Describe("Metadata", func() { Entry(nil, "2/10", "", 2, 10), Entry(nil, "", "", 0, 0), Entry(nil, "A", "", 0, 0), + Entry("ignores values that do not fit in 32 bits", "4294967295", "4294967296", 0, 0), + Entry("ignores a total that does not fit in 32 bits", "2/4294967295", "", 2, 0), + Entry("keeps the largest 32-bit value", "2147483647", "", 2147483647, 0), + Entry("ignores negative values", "-1", "-2", 0, 0), ) Describe("Performers", func() { From a3f41fb422dee23036fd45ead569f0c278a036a8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20S=C3=A1nchez=20Zapico?= Date: Wed, 23 Sep 2026 15:47:16 +0200 Subject: [PATCH 02/26] sec(server): sanitize user-controlled filenames in Content-Disposition (#5895) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * sec(server): sanitize user-controlled filenames in Content-Disposition Playlist export, Subsonic download and public share download built the Content-Disposition header by interpolating a user-controlled name into a quoted-string with fmt.Sprintf. A name containing a double quote closes the string early and the rest is parsed as additional parameters, so a playlist named `party"; filename="evil.html` yielded attachment; filename="party"; filename="evil.html.m3u" letting whoever chose the name decide what the browser saves the download as. The names come from playlists, album/artist names and media file tags. Go's net/http already rewrites CR and LF in header values to spaces, so response splitting was not reachable; parameter injection was. Add str.ContentDispositionAttachment, which emits a sanitized ASCII-only quoted `filename` plus an RFC 5987 `filename*` carrying the original UTF-8 name, and use it at all four call sites. The `filename*` parameter also fixes non-ASCII names, which previously went out raw or were mangled by sanitizing. Signed-off-by: zapisanchez * fix(server): keep download names intact and sanitize filename* Rework ContentDispositionAttachment after review. Names with no ASCII letters now fall back to download. instead of a bare extension (東京.mp3 gave filename="mp3"). filename* is built from the same sanitized name as the ASCII fallback, so path separators, reserved characters, control and bidi characters, and invalid UTF-8 no longer reach it. The ASCII fallback transliterates accents and typographic punctuation (Legião -> Legiao, She’s -> She's) through the existing sanitize.Accents and str.Clear helpers, keeps leading dots, and only trims trailing ones. Names are capped at 255 bytes, keeping the extension. Pure ASCII names now get only the quoted filename parameter, so the header for them matches the previous output byte for byte. filename* is encoded with mime.FormatMediaType instead of a hand-written RFC 5987 encoder. Adds tests for the M3U export and Subsonic download headers. * fix(server): handle dot-only names and long fake extensions A name made only of dots trimmed down to an empty filename. It now falls back to download, like an empty stem does. path.Ext treats anything after the last dot as the extension, so a long suffix with no real extension was kept whole and replaced the stem with download, going past the 255-byte cap. Suffixes longer than 16 bytes are now treated as part of the stem and truncated with it. Neither case is reachable from the current call sites, which always append a short extension. --------- Signed-off-by: zapisanchez Co-authored-by: Deluan --- server/nativeapi/playlists.go | 4 +- server/nativeapi/playlists_test.go | 40 ++++++ server/public/handle_downloads.go | 7 +- server/public/handle_downloads_test.go | 21 +++ .../e2e/subsonic_media_retrieval_test.go | 4 +- server/subsonic/stream.go | 10 +- utils/str/content_disposition.go | 80 ++++++++++++ utils/str/content_disposition_test.go | 120 ++++++++++++++++++ 8 files changed, 271 insertions(+), 15 deletions(-) create mode 100644 utils/str/content_disposition.go create mode 100644 utils/str/content_disposition_test.go diff --git a/server/nativeapi/playlists.go b/server/nativeapi/playlists.go index d215eb9dd..00c1575a5 100644 --- a/server/nativeapi/playlists.go +++ b/server/nativeapi/playlists.go @@ -16,6 +16,7 @@ import ( "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils/req" + "github.com/navidrome/navidrome/utils/str" ) type restHandler = func(rest.RepositoryConstructor, ...rest.Logger) http.HandlerFunc @@ -101,8 +102,7 @@ func handleExportPlaylist(pls playlists.Playlists) http.HandlerFunc { log.Debug(ctx, "Exporting playlist as M3U", "playlistId", plsId, "name", playlist.Name) w.Header().Set("Content-Type", "audio/x-mpegurl") - disposition := fmt.Sprintf("attachment; filename=\"%s.m3u\"", playlist.Name) - w.Header().Set("Content-Disposition", disposition) + w.Header().Set("Content-Disposition", str.ContentDispositionAttachment(playlist.Name+".m3u")) _, err = w.Write([]byte(playlist.ToM3U8())) //nolint:gosec if err != nil { diff --git a/server/nativeapi/playlists_test.go b/server/nativeapi/playlists_test.go index 2f6c578f9..349b4a662 100644 --- a/server/nativeapi/playlists_test.go +++ b/server/nativeapi/playlists_test.go @@ -9,6 +9,7 @@ import ( "time" "github.com/deluan/rest" + "github.com/go-chi/chi/v5" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" @@ -183,6 +184,37 @@ var _ = Describe("Playlist Tracks Endpoint", func() { }) }) +var _ = Describe("handleExportPlaylist", func() { + export := func(name string) *httptest.ResponseRecorder { + r := chi.NewRouter() + r.Get("/playlist/{playlistId}", handleExportPlaylist(&mockPlaylistsService{ + playlist: &model.Playlist{ID: "pls-1", Name: name}, + })) + w := httptest.NewRecorder() + r.ServeHTTP(w, httptest.NewRequest("GET", "/playlist/pls-1", nil)) + return w + } + + It("names the download after the playlist", func() { + w := export("Road Trip") + + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Disposition")).To(Equal(`attachment; filename="Road Trip.m3u"`)) + }) + + It("does not let the playlist name inject a second filename parameter", func() { + w := export(`party"; filename="evil.html`) + + Expect(w.Header().Get("Content-Disposition")).To(Equal(`attachment; filename="party_; filename=_evil.html.m3u"`)) + }) + + It("keeps non-ASCII names in filename*", func() { + w := export("Кино") + + Expect(w.Header().Get("Content-Disposition")).To(Equal(`attachment; filename="download.m3u"; filename*=utf-8''%D0%9A%D0%B8%D0%BD%D0%BE.m3u`)) + }) +}) + var _ = Describe("writePlaylistError", func() { DescribeTable("maps a service error to an HTTP status", func(err error, expected int) { @@ -231,6 +263,7 @@ func (m *mockPlaylistTrackRepo) Read(id string) (any, error) { type mockPlaylistsService struct { playlists.Playlists tracksRepo rest.Repository + playlist *model.Playlist removeImageFn func(ctx context.Context, id string) error setImageFn func(ctx context.Context, id string, reader io.Reader, ext string) error } @@ -249,6 +282,13 @@ func (m *mockPlaylistsService) SetImage(ctx context.Context, id string, reader i return model.ErrNotFound } +func (m *mockPlaylistsService) GetWithTracks(_ context.Context, _ string) (*model.Playlist, error) { + if m.playlist == nil { + return nil, model.ErrNotFound + } + return m.playlist, nil +} + func (m *mockPlaylistsService) TracksRepository(_ context.Context, _ string, _ bool) rest.Repository { return m.tracksRepo } diff --git a/server/public/handle_downloads.go b/server/public/handle_downloads.go index 0012c4b35..e50d7609a 100644 --- a/server/public/handle_downloads.go +++ b/server/public/handle_downloads.go @@ -2,9 +2,7 @@ package public import ( "cmp" - "fmt" "net/http" - "strings" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils/req" @@ -31,9 +29,8 @@ func (pub *Router) handleDownloads(w http.ResponseWriter, r *http.Request) { return } - name := str.SanitizeFilename(cmp.Or(s.Description, s.ID)) - name = strings.ReplaceAll(name, ",", "_") - w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=%q", name+".zip")) + name := cmp.Or(s.Description, s.ID) + w.Header().Set("Content-Disposition", str.ContentDispositionAttachment(name+".zip")) w.Header().Set("Content-Type", "application/zip") err = pub.archiver.ZipShare(ctx, s, w) diff --git a/server/public/handle_downloads_test.go b/server/public/handle_downloads_test.go index 1a97f4379..d2118f47e 100644 --- a/server/public/handle_downloads_test.go +++ b/server/public/handle_downloads_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "io" + "mime" "net/http" "net/http/httptest" "time" @@ -95,6 +96,26 @@ var _ = Describe("handleDownloads", func() { Expect(w.Header().Get("Content-Disposition")).To(Equal(`attachment; filename="AC_DC_ Live_ 1979.zip"`)) }) + It("sanitizes the UTF-8 filename* as well", func() { + shareIs(&model.Share{ID: "abc123", Description: "Sigur Rós/Live", Downloadable: true}) + + w := makeRequest("abc123") + + Expect(w.Header().Get("Content-Disposition")).To(Equal(`attachment; filename="Sigur Ros_Live.zip"; filename*=utf-8''Sigur%20R%C3%B3s_Live.zip`)) + }) + + It("does not let the share description inject a second filename parameter", func() { + shareIs(&model.Share{ID: "abc123", Description: `mix"; filename="evil.html`, Downloadable: true}) + + w := makeRequest("abc123") + + disposition := w.Header().Get("Content-Disposition") + Expect(disposition).ToNot(ContainSubstring(`filename="evil.html`)) + _, params, err := mime.ParseMediaType(disposition) + Expect(err).ToNot(HaveOccurred()) + Expect(params["filename"]).To(Equal(`mix_; filename=_evil.html.zip`)) + }) + It("returns 403 without invoking the archiver when the share is not downloadable", func() { shareIs(&model.Share{ID: "abc123", Description: "No Download", Downloadable: false}) diff --git a/server/subsonic/e2e/subsonic_media_retrieval_test.go b/server/subsonic/e2e/subsonic_media_retrieval_test.go index 079b131ff..268d93b82 100644 --- a/server/subsonic/e2e/subsonic_media_retrieval_test.go +++ b/server/subsonic/e2e/subsonic_media_retrieval_test.go @@ -106,7 +106,7 @@ var _ = Describe("Media Retrieval Endpoints", Ordered, func() { }) Describe("Download", func() { - var trackID string + var trackID, trackTitle string BeforeAll(func() { // All test tracks are mp3 at 320kbps @@ -114,6 +114,7 @@ var _ = Describe("Media Retrieval Endpoints", Ordered, func() { Expect(err).ToNot(HaveOccurred()) Expect(songs).ToNot(BeEmpty()) trackID = songs[0].ID + trackTitle = songs[0].Title }) It("returns error when id parameter is missing", func() { @@ -143,6 +144,7 @@ var _ = Describe("Media Retrieval Endpoints", Ordered, func() { Expect(w.Code).To(Equal(http.StatusOK)) Expect(streamerSpy.LastRequest.Format).To(Equal("opus")) Expect(streamerSpy.LastRequest.BitRate).To(Equal(128)) + Expect(w.Header().Get("Content-Disposition")).To(Equal(`attachment; filename="` + trackTitle + `.opus"`)) }) It("returns error when downloads are disabled", func() { diff --git a/server/subsonic/stream.go b/server/subsonic/stream.go index 28b4585f0..b4a6b821c 100644 --- a/server/subsonic/stream.go +++ b/server/subsonic/stream.go @@ -3,10 +3,8 @@ package subsonic import ( "context" "errors" - "fmt" "net/http" "strconv" - "strings" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/core/stream" @@ -15,6 +13,7 @@ import ( "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server/subsonic/responses" "github.com/navidrome/navidrome/utils/req" + "github.com/navidrome/navidrome/utils/str" ) func (api *Router) Stream(w http.ResponseWriter, r *http.Request) (*responses.Subsonic, error) { @@ -94,9 +93,7 @@ func (api *Router) Download(w http.ResponseWriter, r *http.Request) (*responses. } setHeaders := func(name string) { - name = strings.ReplaceAll(name, ",", "_") - disposition := fmt.Sprintf("attachment; filename=\"%s.zip\"", name) - w.Header().Set("Content-Disposition", disposition) + w.Header().Set("Content-Disposition", str.ContentDispositionAttachment(name+".zip")) w.Header().Set("Content-Type", "application/zip") } @@ -115,8 +112,7 @@ func (api *Router) Download(w http.ResponseWriter, r *http.Request) (*responses. } }() - disposition := fmt.Sprintf("attachment; filename=\"%s\"", stream.Name()) - w.Header().Set("Content-Disposition", disposition) + w.Header().Set("Content-Disposition", str.ContentDispositionAttachment(stream.Name())) _, err = stream.Serve(ctx, w, r) return nil, err diff --git a/utils/str/content_disposition.go b/utils/str/content_disposition.go new file mode 100644 index 000000000..34ae8a5c6 --- /dev/null +++ b/utils/str/content_disposition.go @@ -0,0 +1,80 @@ +package str + +import ( + "cmp" + "fmt" + "mime" + "path" + "strings" + "unicode" + "unicode/utf8" + + "github.com/deluan/sanitize" +) + +const ( + maxFilenameBytes = 255 + maxExtensionBytes = 16 + fallbackFilename = "download" +) + +// ContentDispositionAttachment builds an RFC 6266 attachment header value for a user-controlled filename. +// Non-ASCII names also get a filename* parameter, which clients prefer over the ASCII fallback. +func ContentDispositionAttachment(filename string) string { + stem, ext := splitFilename(filename) + header := fmt.Sprintf("attachment; filename=%q", joinFilename(toASCII(stem), toASCII(ext))) + name := joinFilename(stem, ext) + if isASCII(name) { + return header + } + // FormatMediaType emits a percent-encoded filename* for non-ASCII values + extended := mime.FormatMediaType("attachment", map[string]string{"filename": name}) + return header + strings.TrimPrefix(extended, "attachment") +} + +func splitFilename(filename string) (stem, ext string) { + name := strings.Map(func(r rune) rune { + if unicode.IsControl(r) || unicode.Is(unicode.Bidi_Control, r) { + return -1 + } + return r + }, SanitizeFilename(strings.ToValidUTF8(filename, "_"))) + ext = path.Ext(name) + if len(ext) > maxExtensionBytes { + ext = "" + } + stem = strings.TrimSuffix(name, ext) + if limit := maxFilenameBytes - len(ext); len(stem) > limit { + for limit > 0 && !utf8.RuneStart(stem[limit]) { + limit-- + } + stem = stem[:limit] + } + return stem, ext +} + +func joinFilename(stem, ext string) string { + stem = cmp.Or(strings.TrimSpace(stem), fallbackFilename) + return cmp.Or(strings.TrimRight(stem+ext, " ."), fallbackFilename) +} + +func toASCII(s string) string { + return strings.Map(func(r rune) rune { + switch { + case r == ',': + return '_' + case r >= ' ' && r <= '~': + return r + } + return -1 + }, sanitize.Accents(SanitizeFilename(Clear(s)))) +} + +func isASCII(s string) bool { + for i := 0; i < len(s); i++ { + if s[i] >= utf8.RuneSelf { + return false + } + } + return true +} diff --git a/utils/str/content_disposition_test.go b/utils/str/content_disposition_test.go new file mode 100644 index 000000000..944392a4e --- /dev/null +++ b/utils/str/content_disposition_test.go @@ -0,0 +1,120 @@ +package str_test + +import ( + "mime" + "regexp" + "strings" + + "github.com/navidrome/navidrome/utils/str" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var asciiFilenameRe = regexp.MustCompile(`(?:^|; )filename="([^"]*)"`) + +// asciiName returns the quoted filename parameter, the one used by clients that ignore filename*. +func asciiName(filename string) string { + m := asciiFilenameRe.FindAllStringSubmatch(str.ContentDispositionAttachment(filename), -1) + ExpectWithOffset(1, m).To(HaveLen(1), "expected exactly one quoted filename parameter") + return m[0][1] +} + +// decodedName returns the name an RFC 6266 client picks, preferring filename* when present. +func decodedName(filename string) string { + disp, params, err := mime.ParseMediaType(str.ContentDispositionAttachment(filename)) + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + ExpectWithOffset(1, disp).To(Equal("attachment")) + return params["filename"] +} + +var _ = Describe("ContentDispositionAttachment", func() { + Describe("parameter injection", func() { + const attack = `party"; filename="evil.html.m3u` + + It("does not let a quote in the name open a second parameter", func() { + Expect(str.ContentDispositionAttachment(attack)).To(Equal(`attachment; filename="party_; filename=_evil.html.m3u"`)) + }) + + It("keeps the header parseable", func() { + Expect(decodedName(attack)).To(Equal("party_; filename=_evil.html.m3u")) + }) + + It("does not let a quote in a non-ASCII name inject either", func() { + const utf8Attack = `東京"; filename="evil.html.m3u` + header := str.ContentDispositionAttachment(utf8Attack) + Expect(strings.Count(header, `"`)).To(Equal(2)) + Expect(decodedName(utf8Attack)).To(Equal("東京_; filename=_evil.html.m3u")) + }) + }) + + Describe("ASCII names", func() { + It("sends only the quoted filename, unchanged", func() { + Expect(str.ContentDispositionAttachment("My Playlist.m3u")).To(Equal(`attachment; filename="My Playlist.m3u"`)) + }) + + It("keeps leading dots", func() { + Expect(asciiName("...And Justice for All.zip")).To(Equal("...And Justice for All.zip")) + }) + + It("trims surrounding spaces and trailing dots", func() { + Expect(asciiName(" Greatest Hits .zip")).To(Equal("Greatest Hits.zip")) + Expect(asciiName("Loose End. ")).To(Equal("Loose End")) + }) + + It("falls back to a placeholder when only dots are left", func() { + Expect(str.ContentDispositionAttachment("...")).To(Equal(`attachment; filename="download"`)) + }) + + It("caps a long name whose last dot is not an extension", func() { + name := asciiName("a." + strings.Repeat("x", 300)) + Expect(len(name)).To(BeNumerically("<=", 255)) + Expect(name).To(HavePrefix("a.xxx")) + }) + + It("replaces path separators and reserved characters", func() { + Expect(asciiName("AC/DC: Live, 1979?.zip")).To(Equal("AC_DC_ Live_ 1979_.zip")) + }) + + It("drops control characters", func() { + Expect(asciiName("line\r\nbreak\x00\t.mp3")).To(Equal("linebreak.mp3")) + }) + }) + + Describe("non-ASCII names", func() { + It("transliterates accents in the ASCII fallback", func() { + Expect(asciiName("Legião Urbana.zip")).To(Equal("Legiao Urbana.zip")) + }) + + It("converts typographic punctuation in the ASCII fallback", func() { + Expect(asciiName("She’s a Woman — Live.mp3")).To(Equal("She's a Woman - Live.mp3")) + Expect(asciiName("“Heroes”.mp3")).To(Equal("_Heroes_.mp3")) + Expect(decodedName("She’s a Woman.mp3")).To(Equal("She’s a Woman.mp3")) + }) + + It("keeps the extension when no ASCII letters survive", func() { + Expect(asciiName("東京.mp3")).To(Equal("download.mp3")) + Expect(asciiName("Кино.m3u")).To(Equal("download.m3u")) + Expect(asciiName("東京")).To(Equal("download")) + }) + + DescribeTable("filename* carries the sanitized UTF-8 name", + func(filename, expected string) { + Expect(decodedName(filename)).To(Equal(expected)) + }, + Entry("accented", "Legião Urbana.zip", "Legião Urbana.zip"), + Entry("CJK", "東京.zip", "東京.zip"), + Entry("emoji", "🎵 mix.zip", "🎵 mix.zip"), + Entry("comma and percent", "Sigur Rós, 100%.zip", "Sigur Rós, 100%.zip"), + Entry("path separators", "Sigur Rós/Live: Heima.zip", "Sigur Rós_Live_ Heima.zip"), + Entry("control characters", "Sigur Rós\r\n\x00.zip", "Sigur Rós.zip"), + Entry("bidi override", "Björk\u202Eexe.mp3", "Björkexe.mp3"), + Entry("invalid UTF-8", "Bj\xf6rk Café.mp3", "Bj_rk Café.mp3"), + ) + + It("caps the name length and keeps the extension", func() { + name := decodedName(strings.Repeat("東", 300) + ".zip") + Expect(len(name)).To(BeNumerically("<=", 255)) + Expect(name).To(HaveSuffix("東.zip")) + }) + }) +}) From 39028f65c80734fcf8a00b2b40ce210cd3e5ea91 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Wed, 23 Sep 2026 09:54:35 -0400 Subject: [PATCH 03/26] fix(jellyfin): honor IsPublic when creating a playlist (#6204) POST /Playlists dropped the client's IsPublic flag, so every playlist was created private. JellyBox Player's create-playlist form defaults its "public" checkbox to true, so JellyBox users could never create a public playlist. Upstream's PlaylistsController passes IsPublic into PlaylistCreationRequest. core/playlists.Create has no visibility parameter and widening it would ripple into the Subsonic and native APIs, so createPlaylist follows the same pattern updatePlaylist already uses: after Create succeeds, a non-nil IsPublic is applied with a follow-up Update. The field is a pointer so an absent one keeps today's default instead of forcing private. If that second write fails the handler surfaces the error through playlistError rather than returning the id: answering 200 for a playlist that is not as visible as the client asked is the same silent drop this fixes. --- server/jellyfin/e2e/e2e_suite_test.go | 8 +++++- server/jellyfin/e2e/playlists_test.go | 22 +++++++++++++---- server/jellyfin/playlists.go | 8 ++++++ server/jellyfin/playlists_test.go | 35 +++++++++++++++++++++++++++ 4 files changed, 67 insertions(+), 6 deletions(-) diff --git a/server/jellyfin/e2e/e2e_suite_test.go b/server/jellyfin/e2e/e2e_suite_test.go index aa93f7e38..5aa38cba1 100644 --- a/server/jellyfin/e2e/e2e_suite_test.go +++ b/server/jellyfin/e2e/e2e_suite_test.go @@ -222,8 +222,14 @@ func createPlaylistAs(user model.User, name string, encodedIds ...string) string } body, err := json.Marshal(map[string]any{"Name": name, "Ids": encodedIds}) Expect(err).ToNot(HaveOccurred()) + return createPlaylistBodyAs(user, string(body)) +} + +// createPlaylistBodyAs posts a raw create body, for tests that need fields the helpers above don't +// build, and returns the new playlist's decoded id. +func createPlaylistBodyAs(user model.User, body string) string { var res map[string]string - parseInto(postAs(user, "/Playlists", string(body)), &res) + parseInto(postAs(user, "/Playlists", body), &res) Expect(res["Id"]).ToNot(BeEmpty()) id, ok := dto.DecodeID(res["Id"]) Expect(ok).To(BeTrue()) diff --git a/server/jellyfin/e2e/playlists_test.go b/server/jellyfin/e2e/playlists_test.go index 174ba8660..3dd53227f 100644 --- a/server/jellyfin/e2e/playlists_test.go +++ b/server/jellyfin/e2e/playlists_test.go @@ -21,12 +21,19 @@ var _ = Describe("Playlists", func() { } order := func(plID string) []string { return names(playlistItems(plID).Items) } + createWith := func(body string) string { return createPlaylistBodyAs(adminUser, body) } + openAccess := func(plID string) bool { + var info dto.PlaylistInfo + parseInto(get("/Playlists/"+enc(plID)), &info) + return info.OpenAccess + } + Describe("create", func() { It("creates an empty playlist", func() { plID := createPlaylist("Empty", nil) var info dto.PlaylistInfo parseInto(get("/Playlists/"+enc(plID)), &info) - Expect(info.OpenAccess).To(BeFalse()) + Expect(info.OpenAccess).To(BeFalse(), "a playlist created without IsPublic stays private") Expect(info.Shares).To(BeEmpty()) Expect(info.ItemIds).To(BeEmpty()) }) @@ -48,6 +55,14 @@ var _ = Describe("Playlists", func() { Expect(playlistItems(plID).TotalRecordCount).To(Equal(3)) // Abbey Road (2) + Help! (1) }) + It("creates a public playlist when the client sends IsPublic true", func() { + Expect(openAccess(createWith(`{"Name":"Public","Ids":[],"IsPublic":true}`))).To(BeTrue()) + }) + + It("creates a private playlist when the client sends IsPublic false", func() { + Expect(openAccess(createWith(`{"Name":"Private","Ids":[],"IsPublic":false}`))).To(BeFalse()) + }) + // 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() { @@ -355,10 +370,7 @@ var _ = Describe("Playlists", func() { It("makes a playlist public", func() { plID := createPlaylist("Make Public", nil) Expect(post("/Playlists/"+enc(plID), `{"Name":"Make Public","IsPublic":true}`).Code).To(Equal(http.StatusNoContent)) - - var info dto.PlaylistInfo - parseInto(get("/Playlists/"+enc(plID)), &info) - Expect(info.OpenAccess).To(BeTrue()) + Expect(openAccess(plID)).To(BeTrue()) // Now visible to other users. Expect(queryResult(getAs(regularUser, "/Items?IncludeItemTypes=Playlist&Recursive=true")).TotalRecordCount).To(Equal(1)) }) diff --git a/server/jellyfin/playlists.go b/server/jellyfin/playlists.go index 1052a398f..5fd2df8c9 100644 --- a/server/jellyfin/playlists.go +++ b/server/jellyfin/playlists.go @@ -48,6 +48,7 @@ type createPlaylistRequest struct { Name string `json:"Name"` Ids []string `json:"Ids"` MediaType string `json:"MediaType"` + IsPublic *bool `json:"IsPublic"` } // createPlaylist always creates a new playlist (playlistId "" tells core/playlists.Create not to @@ -69,6 +70,13 @@ func (api *Router) createPlaylist(w http.ResponseWriter, r *http.Request) { api.internalError(w, r, err) return } + // Create takes no visibility, so a requested one costs a second write. + if body.IsPublic != nil { + if err := api.playlists.Update(r.Context(), id, nil, nil, body.IsPublic, nil, nil); err != nil { + api.playlistError(w, r, err) + return + } + } api.ok(w, r, map[string]string{"Id": dto.EncodeID(id)}) } diff --git a/server/jellyfin/playlists_test.go b/server/jellyfin/playlists_test.go index ae0a5da3c..edaa63dd1 100644 --- a/server/jellyfin/playlists_test.go +++ b/server/jellyfin/playlists_test.go @@ -55,6 +55,16 @@ type fakePlaylists struct { deletePlaylistID string deleteErr error + + updatePlaylistID string + updatePublic *bool + updateErr error +} + +func (f *fakePlaylists) Update(_ context.Context, playlistID string, _ *string, _ *string, public *bool, _ []string, _ []int) error { + f.updatePlaylistID = playlistID + f.updatePublic = public + return f.updateErr } func (f *fakePlaylists) Delete(_ context.Context, id string) error { @@ -184,6 +194,31 @@ var _ = Describe("Playlists", func() { invoke(api.createPlaylist, w, r) Expect(w.Code).To(Equal(http.StatusInternalServerError)) }) + + createReq := func(body string) *http.Request { + return httptest.NewRequest("POST", "/Playlists", strings.NewReader(body)). + WithContext(GinkgoT().Context()) + } + + DescribeTable("visibility", + func(body string, wantPublic *bool, wantUpdatedID string) { + w := httptest.NewRecorder() + invoke(api.createPlaylist, w, createReq(body)) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(fp.updatePublic).To(Equal(wantPublic)) + Expect(fp.updatePlaylistID).To(Equal(wantUpdatedID)) + }, + Entry("applies IsPublic true", `{"Name":"Mix","IsPublic":true}`, new(true), testID("pl-new")), + Entry("applies an explicit IsPublic false", `{"Name":"Mix","IsPublic":false}`, new(false), testID("pl-new")), + Entry("leaves visibility alone when IsPublic is omitted", `{"Name":"Mix"}`, nil, ""), + ) + + It("returns 500 when the visibility update fails", func() { + fp.updateErr = errors.New("boom") + w := httptest.NewRecorder() + invoke(api.createPlaylist, w, createReq(`{"Name":"Mix","IsPublic":true}`)) + Expect(w.Code).To(Equal(http.StatusInternalServerError)) + }) }) Describe("getPlaylistItems", func() { From 27483a46dcda604b5e494b3b2010b617878e0593 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20S=C3=A1nchez=20Zapico?= Date: Wed, 23 Sep 2026 18:10:25 +0200 Subject: [PATCH 04/26] fix(server): fail startup on initial setup errors and fix JSON/M3U response headers (#5897) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(server): stop swallowing errors and correct two response bugs Four independent bugs found while reviewing the HTTP layer: initial_setup.go: createInitialAdminUser assigned the users.Put error to a shadowed err, so the outer err (always nil by then, since a CountAll failure panics) was returned instead. A failure to create the admin user was reported as success, and initialSetup went on to commit the "setup complete" property in the same transaction — so no admin user existed and initial setup was skipped on every later boot. auth.go: createAdminUser logged the Put error but returned nil, so createAdmin fell through to doLogin and answered 401 "Invalid username or password" instead of surfacing the real failure. It also logged the whole model.User, which puts the new admin's password in the log in clear text; every other call site logs user.UserName. native_api.go: writeDeleteManyResponse did not return after http.Error when marshaling failed, then wrote a nil body over the 500. It also built the single-id body by hand with html.EscapeString, which does not escape backslashes, so an id ending in one produced `{"id":"a\"}` — invalid JSON. Both shapes now go through json.Marshal. A failed Write is now logged rather than answered with http.Error, which could not work once the body had started. handle_shares.go: handleM3U set Content-Type after WriteHeader, so it was never sent and shared playlists were served with a sniffed type. Signed-off-by: zapisanchez * fix(server): address review feedback - writeDeleteManyResponse uses rest.RespondWithJSON, so the response now has Content-Type: application/json. This also removes a marshal error branch that could never run. - createInitialAdminUser returns the CountAll error instead of panicking, and wraps its errors. initialSetup now stops the server with log.Fatal when setup fails. Before, the error was dropped and the server started with a half-done setup. - Trim comments that described PR history. --------- Signed-off-by: zapisanchez Co-authored-by: Deluan --- server/auth.go | 2 +- server/auth_test.go | 12 +++++ server/initial_setup.go | 14 ++--- server/initial_setup_test.go | 23 ++++++++ server/nativeapi/delete_many_response_test.go | 49 +++++++++++++++++ server/nativeapi/native_api.go | 22 +++----- server/public/handle_shares.go | 2 +- server/public/handle_shares_test.go | 52 +++++++++++++++++++ 8 files changed, 154 insertions(+), 22 deletions(-) create mode 100644 server/nativeapi/delete_many_response_test.go create mode 100644 server/public/handle_shares_test.go diff --git a/server/auth.go b/server/auth.go index d9ade7b29..2aaa93e63 100644 --- a/server/auth.go +++ b/server/auth.go @@ -159,7 +159,7 @@ func createAdminUser(ctx context.Context, ds model.DataStore, username, password } err := ds.User(ctx).Put(&initialUser) if err != nil { - log.Error(ctx, "Could not create initial user", "user", initialUser, err) + log.Error(ctx, "Could not create initial user", "user", initialUser.UserName, err) return fmt.Errorf("creating initial user: %w", err) } return nil diff --git a/server/auth_test.go b/server/auth_test.go index a4d592c51..abe144a12 100644 --- a/server/auth_test.go +++ b/server/auth_test.go @@ -76,6 +76,18 @@ var _ = Describe("Auth", func() { }) }) + Describe("createAdmin when the user cannot be stored", func() { + It("responds 500 rather than falling through to login", func() { + failing := dsWithFailingPut(errors.New("db is down")) + req = httptest.NewRequest("POST", "/createAdmin", strings.NewReader(`{"username":"johndoe", "password":"secret"}`)) + resp = httptest.NewRecorder() + + createAdmin(failing)(resp, req) + + Expect(resp.Code).To(Equal(http.StatusInternalServerError)) + }) + }) + Describe("Login from HTTP headers", func() { const ( trustedIpv4 = "192.168.0.42" diff --git a/server/initial_setup.go b/server/initial_setup.go index 7e974dc21..e75220abe 100644 --- a/server/initial_setup.go +++ b/server/initial_setup.go @@ -16,7 +16,7 @@ import ( func initialSetup(ds model.DataStore) { ctx := context.TODO() - _ = ds.WithTx(func(tx model.DataStore) error { + err := ds.WithTx(func(tx model.DataStore) error { if err := tx.Library(ctx).StoreMusicFolder(); err != nil { return err } @@ -36,6 +36,9 @@ func initialSetup(ds model.DataStore) { err = properties.Put(consts.InitialSetupFlagKey, time.Now().String()) return err }, "initial setup") + if err != nil { + log.Fatal("Error running initial setup", err) + } } // If the Dev Admin user is not present, create it @@ -43,7 +46,7 @@ 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}}) if err != nil { - panic(fmt.Sprintf("Could not access User table: %s", err)) + return fmt.Errorf("could not access User table: %w", err) } if c == 0 { newID := id.NewRandom() @@ -57,12 +60,11 @@ func createInitialAdminUser(ds model.DataStore, initialPassword string) error { NewPassword: initialPassword, IsAdmin: true, } - err := users.Put(&initialUser) - if err != nil { - log.Error("Could not create initial admin user", "user", initialUser, err) + if err := users.Put(&initialUser); err != nil { + return fmt.Errorf("could not create initial admin user: %w", err) } } - return err + return nil } func checkFFmpegInstallation() { diff --git a/server/initial_setup_test.go b/server/initial_setup_test.go index 982046f78..0ce8a39fa 100644 --- a/server/initial_setup_test.go +++ b/server/initial_setup_test.go @@ -2,6 +2,7 @@ package server import ( "context" + "errors" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/tests" @@ -9,6 +10,17 @@ import ( . "github.com/onsi/gomega" ) +type failingPutUserRepo struct { + model.UserRepository + err error +} + +func (r *failingPutUserRepo) Put(*model.User) error { return r.err } + +func dsWithFailingPut(err error) model.DataStore { + return &tests.MockDataStore{MockedUser: &failingPutUserRepo{UserRepository: tests.CreateMockUserRepo(), err: err}} +} + var _ = Describe("initial_setup", func() { var ds model.DataStore @@ -32,5 +44,16 @@ var _ = Describe("initial_setup", func() { Expect(createInitialAdminUser(ds, "second")).To(BeNil()) Expect(ur.CountAll()).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)) + }) + + 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)) + }) }) }) diff --git a/server/nativeapi/delete_many_response_test.go b/server/nativeapi/delete_many_response_test.go new file mode 100644 index 000000000..7d911d6cd --- /dev/null +++ b/server/nativeapi/delete_many_response_test.go @@ -0,0 +1,49 @@ +package nativeapi + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("writeDeleteManyResponse", func() { + var w *httptest.ResponseRecorder + + write := func(ids ...string) map[string]any { + w = httptest.NewRecorder() + writeDeleteManyResponse(w, httptest.NewRequest("DELETE", "/missing", nil), ids) + + var body map[string]any + Expect(json.Unmarshal(w.Body.Bytes(), &body)).To(Succeed(), "response body must be valid JSON: %s", w.Body.String()) + return body + } + + It("returns a single id as an object", func() { + Expect(write("abc123")).To(HaveKeyWithValue("id", "abc123")) + }) + + It("returns multiple ids as a list", func() { + Expect(write("a", "b")).To(HaveKeyWithValue("ids", ConsistOf("a", "b"))) + }) + + It("stays valid JSON when the id contains a backslash", func() { + Expect(write(`a\`)).To(HaveKeyWithValue("id", `a\`)) + }) + + It("stays valid JSON when the id contains a quote", func() { + Expect(write(`a"b`)).To(HaveKeyWithValue("id", `a"b`)) + }) + + It("does not HTML-escape the id into entities", func() { + Expect(write("a&b")).To(HaveKeyWithValue("id", "a&b")) + }) + + It("responds 200 with a JSON content type", func() { + write("abc123") + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).To(Equal("application/json")) + }) +}) diff --git a/server/nativeapi/native_api.go b/server/nativeapi/native_api.go index f931834c6..a7c53df09 100644 --- a/server/nativeapi/native_api.go +++ b/server/nativeapi/native_api.go @@ -2,8 +2,6 @@ package nativeapi import ( "context" - "encoding/json" - "html" "net/http" "strconv" "time" @@ -206,22 +204,18 @@ func (api *Router) addMissingFilesRoute(r chi.Router) { } func writeDeleteManyResponse(w http.ResponseWriter, r *http.Request, ids []string) { - var resp []byte - var err error + var payload any if len(ids) == 1 { - resp = []byte(`{"id":"` + html.EscapeString(ids[0]) + `"}`) + payload = struct { + ID string `json:"id"` + }{ID: ids[0]} } else { - resp, err = json.Marshal(&struct { + payload = struct { Ids []string `json:"ids"` - }{Ids: ids}) - if err != nil { - log.Error(r.Context(), "Error marshaling response", "ids", ids, err) - http.Error(w, err.Error(), http.StatusInternalServerError) - } + }{Ids: ids} } - _, err = w.Write(resp) //nolint:gosec - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) + if err := rest.RespondWithJSON(w, http.StatusOK, payload); err != nil { + log.Error(r.Context(), "Error writing response", "ids", ids, err) } } diff --git a/server/public/handle_shares.go b/server/public/handle_shares.go index d67cfe456..367bff501 100644 --- a/server/public/handle_shares.go +++ b/server/public/handle_shares.go @@ -58,8 +58,8 @@ func (pub *Router) handleM3U(w http.ResponseWriter, r *http.Request) { } s = pub.mapShareToM3U(r, *s) - w.WriteHeader(http.StatusOK) w.Header().Set("Content-Type", "audio/x-mpegurl") + w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(s.ToM3U8())) //nolint:gosec } diff --git a/server/public/handle_shares_test.go b/server/public/handle_shares_test.go new file mode 100644 index 000000000..1bf631fd4 --- /dev/null +++ b/server/public/handle_shares_test.go @@ -0,0 +1,52 @@ +package public + +import ( + "net/http" + "net/http/httptest" + + "github.com/navidrome/navidrome/core" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("handleM3U", func() { + var ds *tests.MockDataStore + var shareRepo *tests.MockShareRepo + var pub *Router + + BeforeEach(func() { + ds = &tests.MockDataStore{} + shareRepo = &tests.MockShareRepo{} + ds.MockedShare = shareRepo + pub = &Router{ds: ds, share: core.NewShare(ds)} + }) + + makeRequest := func(id string) *httptest.ResponseRecorder { + r := httptest.NewRequest("GET", "/public/"+id+"/m3u?%3Aid="+id, nil) + w := httptest.NewRecorder() + pub.handleM3U(w, r) + return w + } + + It("sets the M3U content type", func() { + share := &model.Share{ID: "abc123", Tracks: model.MediaFiles{{ID: "t1", Title: "Track 1"}}} + shareRepo.ID = share.ID + shareRepo.Entity = share + + w := makeRequest("abc123") + + Expect(w.Code).To(Equal(http.StatusOK)) + // Result() has the headers sent at WriteHeader time, unlike w.Header() + Expect(w.Result().Header.Get("Content-Type")).To(Equal("audio/x-mpegurl")) + Expect(w.Body.String()).To(HavePrefix("#EXTM3U")) + }) + + It("returns 404 when the share does not exist", func() { + shareRepo.ID = "other" + shareRepo.Entity = &model.Share{ID: "other"} + + Expect(makeRequest("missing").Code).To(Equal(http.StatusNotFound)) + }) +}) From ee6dd1bc031154788988c7b62cac18f2a49c67b5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Wed, 23 Sep 2026 17:04:50 -0400 Subject: [PATCH 05/26] fix(scanner): stop DB lock starvation during scans on slow storage (#6201) * fix(artwork): pause the artwork worker while a scan is running The artwork worker added in 0.64 writes to the database continuously, including while a scan runs. On slow storage the scanner holds the write lock for many seconds per folder, so the two writers keep timing each other out: artwork writes fail with "database is locked", and a single busy timeout on the scanner side aborts the whole scan. The worker now stops dispatching queue items while scanner.IsScanning reports true, including mid-batch, and resumes on the next poll after the scan ends. Artwork requests are unaffected, since they serve local art without the worker. * fix(db): run ANALYZE one index at a time so writers are not starved A full ANALYZE is a single write transaction, so every other write waits for it to finish and fails after the 15s busy timeout. On slow NAS storage it was measured taking over 26 minutes. The analysis now runs ANALYZE per index (per table for unindexed and WITHOUT ROWID tables), which produces the same sqlite_stat1 rows as a full ANALYZE, and pauses briefly between steps (up to 150ms, just above SQLite's longest busy-handler sleep) so waiting writers get the lock. * fix(scanner): ignore Synology @eaDir metadata folders Synology creates an @eaDir folder next to media files, holding one subfolder per file with generated thumbnails. The scanner and watcher treated them as regular folders, which on one reported library added tens of thousands of extra folders to every scan. * fix(db): analyze tables with only partial indexes as a whole A partial index does not record the table's row count, so a table whose only indexes are partial needs a table-level ANALYZE to get the sqlite_stat1 row a full ANALYZE would write. Navidrome's schema has no such table today, but the stepped analysis should match a full ANALYZE for any schema a future migration creates. * fix(scanner): retry busy folder saves and stop phase 1 on a fatal error On slow storage, a single SQLITE_BUSY while saving a folder aborted the whole scan, even when another writer held the lock only briefly. The folder save now runs as a retryable unit: on a busy error it waits (5s, 10s, 15s) and reruns the transaction, up to three times, before failing. Side effects that do not survive a rollback (the album ID map consumed by persistAlbum, the artwork queue items, the image-change record) are rebuilt per attempt or recorded only after a successful commit. When a folder save does fail, phase 1 used to keep walking the library and reading tags for every remaining folder, discarding the results, before reporting the error; a reporter saw 40 silent minutes. The walk now stops as soon as the save fails, and the walker honors cancellation instead of blocking on its channel. Because an early stop leaves folders unvisited, phase 1 no longer marks unvisited folders missing when the phase failed; the resumed scan handles them. * refactor(persistence): move busy retry into DataStore.WithTxRetry The scanner retried its folder save itself, which meant it had to know SQLite error codes. WithTxRetry now owns that policy: it reruns the block in a fresh transaction on SQLITE_BUSY, up to three times with growing delays, and runs it only once when already inside a transaction, since the outer transaction would still hold the lock. The block receives the context to use, and attempts that will be retried carry a marker so a busy statement in them is logged as a warning; only the final attempt logs errors. The scanner's inner error logs are folded into wrapped errors, so a recovered retry no longer prints error-level lines, and the folder path travels in the log context. * fix(persistence): join the enclosing transaction in a nested WithTxRetry Called on a store that is already inside a transaction, WithTxRetry went through WithTx, which opens a second, independent transaction on another connection. That transaction waits on the lock the outer one holds and fails with SQLITE_BUSY, and if it does succeed the outer transaction cannot roll it back. It now runs the block on the enclosing transaction, which owns the lock, the commit and the rollback. Found by a Codex (gpt-6-sol) review. * fix(scanner): retry the remaining scan writes on a busy database Every write step after phase 1 still aborted the whole scan on a single SQLITE_BUSY: phase 1 finalize, phase 2 moves and purge, phase 3 album saves and play count refreshes, the deferred playlist import flag, library ScanBegin, GC, the missing-artwork enqueue, tag counts, and the final library update. They now go through WithTxRetry. The phase 2 move had to be made rerun-safe first: it changed the target track's ID inside the transaction, so a rerun would have deleted the moved track itself, and it marked album annotations as handled even when the transaction rolled back. It now works on a copy per attempt and records the annotation reassignment only after a commit. Artist.RefreshStats is left alone: it updates artists in batches outside a transaction, and one transaction around all of them would hold the write lock for the whole refresh on slow storage. Phase 4 playlist imports go through the playlist service and are left for a follow-up. * fix(scanner): claim the album before moving its annotations The rerun-safe moveMatched checked processedAlbumAnnotations before its transaction and marked the album only after the commit. Phase 2 runs same-library and cross-library moves in separate pipeline stages, so two moves into one album could both pass the check; the second would reassign annotations again and overwrite the album's created_at. The album is now claimed under the lock before the transaction, as the old code effectively did, and the claim is released if the move fails so a later move can still reassign. Found by a Codex (gpt-6-sol) review. * fix(artwork): keep artwork housekeeping from writing during scans The artwork worker already pauses while a scan runs, but its housekeeping jobs did not: the hourly missing-artwork recheck (a bulk INSERT ... SELECT over albums and artists), the startup run of the same recheck, and the daily prune all kept competing with the scanner for the write lock. They now run through LockForMaintenance, like the scheduled DB analysis: they skip while a scan is running and keep a scan from starting until they finish. Skipping the recheck loses nothing, since each scan with changes queues missing artwork at its end. * refactor(scanner): log retried step errors once, from the caller Blocks passed to WithTxRetry still logged their own errors at error level on every attempt, so a busy error that a retry absorbed printed several error lines (GC printed three). They now return wrapped errors and the callers, which already log them, report the final outcome once. Also: drop a leftover variable in phase 1 finalize, check the walk context once, stop repeating the folder field that is already in the log context, stop shadowing finalize's err in phase 3, and format the WithTxRetry scope the same way as WithTx. * test(scanner): make the scanner suite's temp DB cleanup best effort Which DB file the process-wide DB handle opens depends on which spec touches it first. When the Scanner container wins the random order, its temp DB stays open until db.Close after RunSpecs, and on Windows removing the temp dir fails with 'being used by another process'. Ginkgo pins that on the container's last spec, which is now one of the busy-database specs. The sibling suites skip Windows for the same reason; this one now removes its temp dir on a best-effort basis instead, so it keeps running there. --- cmd/root.go | 28 ++- core/artwork/worker.go | 13 ++ core/artwork/worker_test.go | 52 +++++ db/db.go | 6 + db/db_test.go | 16 ++ db/optimize.go | 56 ++++- db/optimize_test.go | 37 ++++ model/datastore.go | 3 + persistence/persistence.go | 60 +++++- persistence/persistence_test.go | 73 +++++++ persistence/sql_base_repository.go | 4 + scanner/phase_1_folders.go | 280 +++++++++++++------------ scanner/phase_2_missing_tracks.go | 97 +++++---- scanner/phase_2_missing_tracks_test.go | 114 ++++++++++ scanner/phase_3_refresh_albums.go | 17 +- scanner/phase_4_playlists.go | 5 +- scanner/scanner.go | 41 ++-- scanner/scanner_test.go | 62 +++++- scanner/walk_dir_tree.go | 13 +- scanner/walk_dir_tree_test.go | 1 + tests/mock_data_store.go | 4 + 21 files changed, 754 insertions(+), 228 deletions(-) diff --git a/cmd/root.go b/cmd/root.go index f632b0425..02cd30240 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -375,10 +375,26 @@ func startPlaybackServer(ctx context.Context) func() error { func startArtworkWorker(ctx context.Context, worker *artwork.Worker) func() error { return func() error { log.Info(ctx, "Starting artwork worker") + // The scanner writes to the DB for its whole run; competing for the write lock makes both fail. + worker.PauseWhile(scanner.IsScanning) return worker.Run(ctx) } } +// outsideScan runs a DB maintenance job unless a scan is running, and keeps a scan from starting +// until it ends; both write to the DB, and competing for the lock can make either fail. +func outsideScan(ctx context.Context, job string, run func(context.Context) error) { + release, ok := scanner.LockForMaintenance() + if !ok { + log.Debug(ctx, "Skipping "+job+" because a scan is in progress") + return + } + defer release() + if err := run(ctx); err != nil { + log.Error(ctx, "Error running "+job, err) + } +} + // scheduleArtworkHousekeeping registers the recurring missing-state and prune jobs, and // reports an artwork config change without acting on it. func scheduleArtworkHousekeeping(ctx context.Context, worker *artwork.Worker) func() error { @@ -386,26 +402,20 @@ func scheduleArtworkHousekeeping(ctx context.Context, worker *artwork.Worker) fu schedulerInstance := scheduler.GetInstance() if _, err := schedulerInstance.Add(consts.ArtworkEnqueueMissingSchedule, func() { - if err := worker.EnqueueMissingAll(ctx); err != nil { - log.Error(ctx, "Error enqueueing missing artwork rechecks", err) - } + outsideScan(ctx, "artwork missing-state recheck", worker.EnqueueMissingAll) }); err != nil { log.Error(ctx, "Error scheduling artwork missing-state recheck", err) } if _, err := schedulerInstance.Add(consts.ArtworkPruneSchedule, func() { - if err := worker.RunPrune(ctx); err != nil { - log.Error(ctx, "Error running artwork prune", err) - } + outsideScan(ctx, "artwork prune", worker.RunPrune) }); err != nil { log.Error(ctx, "Error scheduling artwork prune", err) } // Also run the missing-row recheck once at startup so a never-scanned entity is picked up // immediately, not only on the next hourly tick (e.g. after enabling the feature). - if err := worker.EnqueueMissingAll(ctx); err != nil { - log.Error(ctx, "Error enqueueing missing artwork rechecks", err) - } + outsideScan(ctx, "artwork missing-state recheck", worker.EnqueueMissingAll) if err := worker.ReconcileConfig(ctx); err != nil { log.Error(ctx, "Error checking the artwork config fingerprint", err) diff --git a/core/artwork/worker.go b/core/artwork/worker.go index bb09be55e..e8fb3a11c 100644 --- a/core/artwork/worker.go +++ b/core/artwork/worker.go @@ -46,6 +46,7 @@ type Worker struct { pruneMu sync.RWMutex pools []*drainPool runCtx context.Context + paused func() bool gatesMu sync.Mutex gates map[string]*extGate @@ -59,6 +60,7 @@ func NewWorker(ds model.DataStore, store *ImageStore, ag *agents.Agents, ffmpeg broker: broker, pools: newDrainPools(), runCtx: context.Background(), + paused: func() bool { return false }, gates: map[string]*extGate{}, } w.proc.resolver = newResolver(ds, ag, ffmpeg, w.gate) @@ -90,6 +92,11 @@ var ( } ) +// PauseWhile holds off queue draining whenever paused reports true. Call it before Run. +func (w *Worker) PauseWhile(paused func() bool) { + w.paused = paused +} + // Run blocks draining the queue until ctx is cancelled. func (w *Worker) Run(ctx context.Context) error { w.runCtx = ctx @@ -143,6 +150,9 @@ func (w *Worker) EnqueueMissingAll(ctx context.Context) error { } func (w *Worker) drain(ctx context.Context, concurrency int, kinds ...string) (int, error) { + if w.paused() { + return 0, nil + } // 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...) @@ -170,6 +180,9 @@ func (w *Worker) drain(ctx context.Context, concurrency int, kinds ...string) (i wg.Wait() return len(items), nil //nolint:nilerr // a cancelled drain is a clean stop, not an error } + if w.paused() { + break + } wg.Go(func() { defer func() { <-sem }() out, got := w.process(ctx, item) diff --git a/core/artwork/worker_test.go b/core/artwork/worker_test.go index ebb8de251..d74e3a5ed 100644 --- a/core/artwork/worker_test.go +++ b/core/artwork/worker_test.go @@ -866,6 +866,34 @@ var _ = Describe("Worker", func() { } }) + It("stops dispatching and leaves the rest queued when paused mid-batch", func() { + folderRepo.result = []model.Folder{{ + Path: "tests/fixtures/artist/an-album", + ImageFiles: []string{"cover.jpg"}, + }} + albums := model.Albums{} + 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{ + 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() + return n < 8 + }) + + _, err := w.drain(ctx, 1) + Expect(err).ToNot(HaveOccurred()) + + count, err := queueRepo.Count() + Expect(err).ToNot(HaveOccurred()) + Expect(count).To(Equal(int64(7)), "only the item dispatched before the pause may leave the queue") + }) + 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"}}) @@ -896,6 +924,30 @@ var _ = Describe("Worker", func() { Eventually(done, time.Second).Should(Receive(BeNil())) }) + It("does not drain the queue while paused", func() { + folderRepo.result = []model.Folder{{ + Path: "tests/fixtures/artist/an-album", + ImageFiles: []string{"cover.jpg"}, + }} + ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{ + {ID: "al1", Name: "Album", FolderIDs: []string{"f1"}}, + }) + Expect(queueRepo.Enqueue(model.ArtworkQueueItem{ + ItemKind: "al", ItemID: "al1", Priority: model.ArtworkPriorityScan, + })).To(Succeed()) + w.PauseWhile(func() bool { return true }) + + runCtx, cancel := context.WithCancel(ctx) + done := make(chan error, 1) + go func() { done <- w.Run(runCtx) }() + DeferCleanup(func() { + cancel() + Eventually(done, 2*time.Second).Should(Receive(BeNil())) + }) + + Consistently(func() any { return findQueued(queueRepo, "al", "al1") }, 300*time.Millisecond).ShouldNot(BeNil()) + }) + It("does not leak goroutines after Run exits", func() { DeferCleanup(configtest.SetupConfig()) diff --git a/db/db.go b/db/db.go index 685886edb..a9c6c4a15 100644 --- a/db/db.go +++ b/db/db.go @@ -133,6 +133,12 @@ func ErrorCodes(err error) (code, extended int, ok bool) { return int(se.Code), int(se.ExtendedCode), true } +// IsBusy reports whether err is SQLITE_BUSY, including BUSY_SNAPSHOT, which only a new transaction clears. +func IsBusy(err error) bool { + code, _, ok := ErrorCodes(err) + return ok && code == int(sqlite3.ErrBusy) +} + type statusLogger struct{ numPending int } func (*statusLogger) Fatalf(format string, v ...any) { log.Fatal(fmt.Sprintf(format, v...)) } diff --git a/db/db_test.go b/db/db_test.go index 2ce01dc3d..e3a52e1fb 100644 --- a/db/db_test.go +++ b/db/db_test.go @@ -3,8 +3,12 @@ package db_test import ( "context" "database/sql" + "errors" + "fmt" "testing" + "github.com/mattn/go-sqlite3" + "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/tests" @@ -19,6 +23,18 @@ func TestDB(t *testing.T) { RunSpecs(t, "DB Suite") } +var _ = DescribeTable("IsBusy", + func(err error, expected bool) { + Expect(db.IsBusy(err)).To(Equal(expected)) + }, + Entry("SQLITE_BUSY", sqlite3.Error{Code: sqlite3.ErrBusy}, true), + Entry("SQLITE_BUSY_SNAPSHOT", sqlite3.Error{Code: sqlite3.ErrBusy, ExtendedCode: sqlite3.ErrBusySnapshot}, true), + Entry("a wrapped SQLITE_BUSY", fmt.Errorf("persisting: %w", sqlite3.Error{Code: sqlite3.ErrBusy}), true), + Entry("another SQLite error", sqlite3.Error{Code: sqlite3.ErrConstraint}, false), + Entry("a non-SQLite error", errors.New("database is locked"), false), + Entry("nil", nil, false), +) + var _ = Describe("IsSchemaEmpty", func() { var database *sql.DB var ctx context.Context diff --git a/db/optimize.go b/db/optimize.go index f46906c4e..aff36e8fd 100644 --- a/db/optimize.go +++ b/db/optimize.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "strconv" + "strings" "sync" "time" @@ -134,16 +135,65 @@ func optimizeAt(ctx context.Context, db *sql.DB, now time.Time) error { return recordAnalyzeError(ctx, db, now, fmt.Errorf("marking ANALYZE pending: %w", err)) } log.Debug(ctx, "Refreshing query planner statistics") - _, err := db.ExecContext(ctx, "ANALYZE") - if err != nil { + if err := analyzeInSteps(ctx, db); err != nil { return recordAnalyzeError(ctx, db, now, fmt.Errorf("running ANALYZE: %w", err)) } - if err = recordAnalyzeSuccess(ctx, db, now); err != nil { + if err := recordAnalyzeSuccess(ctx, db, now); err != nil { return recordAnalyzeError(ctx, db, now, err) } return nil } +// One ANALYZE per index (whole table if WITHOUT ROWID or lacking a non-partial index) yields the +// same sqlite_stat1 rows as a full ANALYZE, but frees the write lock between steps. +const analyzeTargetsSQL = ` +SELECT i.name FROM sqlite_schema i JOIN pragma_table_list t ON t.schema = 'main' AND t.name = i.tbl_name +WHERE i.type = 'index' AND t.wr = 0 AND EXISTS (SELECT 1 FROM pragma_index_list(t.name) l WHERE l.partial = 0) +UNION ALL +SELECT t.name FROM pragma_table_list t +WHERE t.schema = 'main' AND t.type IN ('table', 'shadow') AND t.name NOT LIKE 'sqlite_%' + AND (t.wr = 1 OR NOT EXISTS (SELECT 1 FROM pragma_index_list(t.name) l WHERE l.partial = 0))` + +// analyzeMaxYield is just above SQLite's longest busy-handler sleep, so every waiting writer +// retries during the pause. +const analyzeMaxYield = 150 * time.Millisecond + +func analyzeInSteps(ctx context.Context, db *sql.DB) error { + targets, err := analyzeTargets(ctx, db) + if err != nil { + return err + } + for _, target := range targets { + start := time.Now() + if _, err := db.ExecContext(ctx, `ANALYZE "`+strings.ReplaceAll(target, `"`, `""`)+`"`); err != nil { + return fmt.Errorf("analyzing %s: %w", target, err) + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(min(time.Since(start), analyzeMaxYield)): + } + } + return nil +} + +func analyzeTargets(ctx context.Context, db *sql.DB) ([]string, error) { + rows, err := db.QueryContext(ctx, analyzeTargetsSQL) + if err != nil { + return nil, fmt.Errorf("listing ANALYZE targets: %w", err) + } + defer rows.Close() + var targets []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, fmt.Errorf("listing ANALYZE targets: %w", err) + } + targets = append(targets, name) + } + return targets, rows.Err() +} + func recordAnalyzeSuccess(ctx context.Context, db *sql.DB, now time.Time) error { tx, err := db.BeginTx(ctx, nil) if err != nil { diff --git a/db/optimize_test.go b/db/optimize_test.go index da9b3b9c9..6ad9b9cb4 100644 --- a/db/optimize_test.go +++ b/db/optimize_test.go @@ -75,6 +75,43 @@ var _ = Describe("Optimize", func() { Expect(getProperty(consts.DBAnalyzePendingKey)).To(Equal("0")) }) + It("produces the same statistics as a single full ANALYZE", func() { + putProperty(consts.DBAnalyzePendingKey, "1") + for _, stmt := range []string{ + "create table unindexed(id integer primary key, v int)", + "insert into unindexed(v) select flag from analyze_probe", + "create table no_rowid(k text primary key, v int) without rowid", + "insert into no_rowid select 'k' || id, id % 7 from analyze_probe", + "create index no_rowid_v on no_rowid(v)", + "create table partial_only(id integer primary key, v int)", + "insert into partial_only(v) select id % 5 from analyze_probe", + "create index partial_only_v on partial_only(v) where v = 1", + "analyze", + } { + _, err := database.Exec(stmt) + Expect(err).ToNot(HaveOccurred()) + } + statRows := func() []string { + rows, err := database.Query("select tbl || '|' || coalesce(idx, '') || '|' || stat from sqlite_stat1 order by 1") + Expect(err).ToNot(HaveOccurred()) + defer rows.Close() + var res []string + for rows.Next() { + var s string + Expect(rows.Scan(&s)).To(Succeed()) + res = append(res, s) + } + return res + } + fullAnalyze := statRows() + _, err := database.Exec("delete from sqlite_stat1") + Expect(err).ToNot(HaveOccurred()) + + Expect(db.OptimizeDBAt(ctx, database, now)).To(Succeed()) + + Expect(statRows()).To(Equal(fullAnalyze)) + }) + It("runs when no previous analysis was recorded", func() { ran, err := db.OptimizeDBIfNeeded(ctx, database, now) Expect(err).ToNot(HaveOccurred()) diff --git a/model/datastore.go b/model/datastore.go index 273ca714b..26687d5d4 100644 --- a/model/datastore.go +++ b/model/datastore.go @@ -47,5 +47,8 @@ type DataStore interface { WithTx(block func(tx DataStore) error, scope ...string) error WithTxImmediate(block func(tx DataStore) error, scope ...string) error + // WithTxRetry runs block in a transaction, rerunning it while SQLite reports the database busy. + // For background work only (it can take minutes), and block must be safe to rerun after a rollback. + WithTxRetry(ctx context.Context, block func(ctx context.Context, tx DataStore) error, scope ...string) error GC(ctx context.Context, libraryIDs ...int) error } diff --git a/persistence/persistence.go b/persistence/persistence.go index 9d3a33cfc..589812266 100644 --- a/persistence/persistence.go +++ b/persistence/persistence.go @@ -3,6 +3,7 @@ package persistence import ( "context" "database/sql" + "fmt" "reflect" "time" @@ -138,11 +139,15 @@ func (s *SQLStore) Resource(ctx context.Context, m any) model.ResourceRepository return nil } -func (s *SQLStore) WithTx(block func(tx model.DataStore) error, scope ...string) error { - var msg string +func scopeLabel(scope []string) string { if len(scope) > 0 { - msg = scope[0] + return scope[0] } + return "" +} + +func (s *SQLStore) WithTx(block func(tx model.DataStore) error, scope ...string) error { + msg := scopeLabel(scope) start := time.Now() conn, inTx := s.db.(*dbx.DB) if !inTx { @@ -177,6 +182,51 @@ func (s *SQLStore) WithTxImmediate(block func(tx model.DataStore) error, scope . }, scope...) } +// txRetryDelay spaces out reruns of a busy transaction. Each attempt has already waited out the +// busy timeout, so WithTxRetry gives up only after a sustained lock. +var txRetryDelay = 5 * time.Second + +const txMaxRetries = 3 + +func (s *SQLStore) WithTxRetry(ctx context.Context, block func(ctx context.Context, tx model.DataStore) error, scope ...string) error { + // Inside a transaction, join it: the outer one holds the lock and owns commit and rollback + if _, ok := s.db.(*dbx.DB); !ok { + return block(ctx, s) + } + for attempt := 0; ; attempt++ { + attemptCtx := ctx + if attempt < txMaxRetries { + attemptCtx = withBusyRetry(ctx) + } + err := s.WithTx(func(tx model.DataStore) error { return block(attemptCtx, tx) }, scope...) + if attempt == txMaxRetries || !db.IsBusy(err) { + return err + } + log.Warn(ctx, "Database busy, retrying transaction", "scope", scopeLabel(scope), "attempt", attempt+1, err) + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(time.Duration(attempt+1) * txRetryDelay): + } + } +} + +type busyRetryKey struct{} + +// withBusyRetry marks a transaction attempt that WithTxRetry will rerun, so a busy statement in it +// is logged as a warning rather than an error. +func withBusyRetry(ctx context.Context) context.Context { + return context.WithValue(ctx, busyRetryKey{}, true) +} + +func hasBusyRetry(ctx context.Context) bool { + if ctx == nil { + return false + } + retry, _ := ctx.Value(busyRetryKey{}).(bool) + return retry +} + func (s *SQLStore) GC(ctx context.Context, libraryIDs ...int) error { trace := func(ctx context.Context, msg string, f func() error) func() error { return func() error { @@ -207,9 +257,9 @@ func (s *SQLStore) GC(ctx context.Context, libraryIDs ...int) error { trace(ctx, "remove orphan playlist tracks", func() error { return s.Playlist(ctx).(*playlistRepository).removeOrphans() }), ) if err != nil { - log.Error(ctx, "Error tidying up database", err) + return fmt.Errorf("tidying up database: %w", err) } - return err + return nil } func (s *SQLStore) getDBXBuilder() dbx.Builder { diff --git a/persistence/persistence_test.go b/persistence/persistence_test.go index 13e56bde1..e13f6a231 100644 --- a/persistence/persistence_test.go +++ b/persistence/persistence_test.go @@ -2,7 +2,10 @@ package persistence import ( "context" + "errors" + "time" + "github.com/mattn/go-sqlite3" "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/model" . "github.com/onsi/ginkgo/v2" @@ -55,4 +58,74 @@ var _ = Describe("SQLStore", func() { }) }) }) + + Describe("WithTxRetry", func() { + busy := sqlite3.Error{Code: sqlite3.ErrBusy} + BeforeEach(func() { + DeferCleanup(func(d time.Duration) { txRetryDelay = d }, txRetryDelay) + txRetryDelay = 0 + }) + + It("reruns a busy transaction from a clean rollback", 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()) + if len(attempts) < 3 { + return busy + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(attempts).To(Equal([]bool{true, true, true})) + Expect(ds.Property(ctx).Get("retry-key")).To(Equal("attempt")) + }) + + It("gives up after the last retry, which is not marked as retried", func() { + var attempts []bool + err := ds.WithTxRetry(ctx, func(ctx context.Context, _ model.DataStore) error { + attempts = append(attempts, hasBusyRetry(ctx)) + return busy + }) + Expect(db.IsBusy(err)).To(BeTrue()) + Expect(attempts).To(Equal([]bool{true, true, true, false})) + }) + + It("does not rerun on other errors", func() { + calls := 0 + err := ds.WithTxRetry(ctx, func(context.Context, model.DataStore) error { + calls++ + return sqlite3.Error{Code: sqlite3.ErrConstraint} + }) + Expect(err).To(HaveOccurred()) + Expect(calls).To(Equal(1)) + }) + + It("does not rerun when called inside a transaction", func() { + calls := 0 + err := ds.WithTx(func(tx model.DataStore) error { + return tx.WithTxRetry(ctx, func(context.Context, model.DataStore) error { + calls++ + return busy + }) + }) + Expect(db.IsBusy(err)).To(BeTrue()) + Expect(calls).To(Equal(1)) + }) + + 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.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") + })).To(Succeed()) + return rollback + }) + Expect(err).To(MatchError(rollback)) + _, err = ds.Property(ctx).Get("inner-key") + Expect(err).To(MatchError(model.ErrNotFound)) + }) + }) }) diff --git a/persistence/sql_base_repository.go b/persistence/sql_base_repository.go index bc841db03..026e42b03 100644 --- a/persistence/sql_base_repository.go +++ b/persistence/sql_base_repository.go @@ -643,5 +643,9 @@ 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) { + log.Warn(append(fields, err)...) + return + } log.Error(append(fields, err)...) } diff --git a/scanner/phase_1_folders.go b/scanner/phase_1_folders.go index 6107b3316..feefde032 100644 --- a/scanner/phase_1_folders.go +++ b/scanner/phase_1_folders.go @@ -45,7 +45,9 @@ func createPhaseFolders(ctx context.Context, state *scanState, ds model.DataStor jobs = append(jobs, job) } - return &phaseFolders{jobs: jobs, ctx: ctx, ds: ds, state: state, imageChanges: &imageChangeCollector{ds: ds}} + walkCtx, stopWalk := context.WithCancelCause(ctx) + return &phaseFolders{jobs: jobs, ctx: ctx, walkCtx: walkCtx, stopWalk: stopWalk, ds: ds, state: state, + imageChanges: &imageChangeCollector{ds: ds}} } type scanJob struct { @@ -123,6 +125,8 @@ 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 + stopWalk context.CancelCauseFunc state *scanState prevAlbumPIDConf string imageChanges *imageChangeCollector @@ -144,15 +148,15 @@ func (p *phaseFolders) producer() ppl.Producer[*folderEntry] { var total int64 var totalChanged int64 for _, job := range p.jobs { - if utils.IsCtxDone(p.ctx) { + if utils.IsCtxDone(p.walkCtx) { break } - outputChan, err := walkDirTree(p.ctx, job, job.targetFolders...) + outputChan, err := walkDirTree(p.walkCtx, job, job.targetFolders...) if err != nil { log.Warn(p.ctx, "Scanner: Error scanning library", "lib", job.lib.Name, err) } - for folder := range pl.ReadOrDone(p.ctx, outputChan) { + for folder := range pl.ReadOrDone(p.walkCtx, outputChan) { job.numFolders.Add(1) p.state.sendProgress(&ProgressInfo{ LibID: job.lib.ID, @@ -208,6 +212,9 @@ func (p *phaseFolders) stages() []ppl.Stage[*folderEntry] { func (p *phaseFolders) processFolder(entry *folderEntry) (*folderEntry, error) { defer p.measure(entry)() + if err := context.Cause(p.walkCtx); err != nil { + return entry, err + } // Load children mediafiles from DB cursor, err := p.ds.MediaFile(p.ctx).GetCursor(model.QueryOptions{ @@ -331,128 +338,129 @@ func (p *phaseFolders) persistChanges(entry *folderEntry) (*folderEntry, error) defer p.measure(entry)() p.state.changesDetected.Store(true) - // Collect artwork queue items for changed albums/artists, enqueued in the same transaction - var queueItems []model.ArtworkQueueItem - - err := p.ds.WithTx(func(tx model.DataStore) error { - // Instantiate all repositories just once per folder - folderRepo := tx.Folder(p.ctx) - tagRepo := tx.Tag(p.ctx) - artistRepo := tx.Artist(p.ctx) - libraryRepo := tx.Library(p.ctx) - albumRepo := tx.Album(p.ctx) - mfRepo := tx.MediaFile(p.ctx) - - // A new folder's albums/artists are enqueued below; only pre-existing folders need the diff. - if !entry.isNew() { - if changed, artistImage := entry.imagesChanged(); changed { - p.imageChanges.record(entry.job.lib, imageChangedFolder{ - id: entry.id, path: entry.path, artistImage: artistImage, - }) - } - } - - // Save folder to DB - folder := entry.toFolder() - err := folderRepo.Put(folder) - if err != nil { - log.Error(p.ctx, "Scanner: Error persisting folder to DB", "folder", entry.path, err) - return err - } - - // Save all tags to DB - err = tagRepo.Add(entry.job.lib.ID, entry.tags...) - if err != nil { - log.Error(p.ctx, "Scanner: Error persisting tags to DB", "folder", entry.path, err) - return 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", - "mbz_artist_id", "sort_artist_name", "order_artist_name", "full_text", "search_normalized", "updated_at") - if err != nil { - log.Error(p.ctx, "Scanner: Error persisting artist to DB", "folder", entry.path, "artist", entry.artists[i].Name, err) - return err - } - err = libraryRepo.AddArtist(entry.job.lib.ID, entry.artists[i].ID) - if err != nil { - log.Error(p.ctx, "Scanner: Error adding artist to library", "lib", entry.job.lib.ID, "artist", entry.artists[i].Name, err) - return err - } - if entry.artists[i].Name != consts.UnknownArtist && entry.artists[i].Name != consts.VariousArtists { - queueItems = append(queueItems, scanArtworkItem(model.KindArtistArtwork, entry.artists[i].ID)) - } - } - - // Save all new/modified albums to DB. Their information will be incomplete, but they will be refreshed later - for i := range entry.albums { - err = p.persistAlbum(albumRepo, &entry.albums[i], entry.albumIDMap) - if err != nil { - log.Error(p.ctx, "Scanner: Error persisting album to DB", "folder", entry.path, "album", entry.albums[i], err) - return err - } - if entry.albums[i].Name != consts.UnknownAlbum { - queueItems = append(queueItems, scanArtworkItem(model.KindAlbumArtwork, entry.albums[i].ID)) - } - } - - // Save all tracks to DB - for i := range entry.tracks { - err = mfRepo.Put(&entry.tracks[i]) - if err != nil { - log.Error(p.ctx, "Scanner: Error persisting mediafile to DB", "folder", entry.path, "track", entry.tracks[i], err) - return err - } - } - - // 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(p.ctx).DeleteForItems(model.KindMediaFileArtwork, trackIDs); err != nil { - log.Warn(p.ctx, "Scanner: could not invalidate media_file artwork", "folder", entry.path, err) - } - } - - // Mark all missing tracks as not available - if len(entry.missingTracks) > 0 { - err = mfRepo.MarkMissing(true, entry.missingTracks...) - if err != nil { - log.Error(p.ctx, "Scanner: Error marking missing tracks", "folder", entry.path, err) - return err - } - - // Touch all albums that have missing tracks, so they get refreshed in later phases - groupedMissingTracks := slice.ToMap(entry.missingTracks, func(mf *model.MediaFile) (string, struct{}) { - return mf.AlbumID, struct{}{} - }) - albumsToUpdate := slices.Collect(maps.Keys(groupedMissingTracks)) - err = albumRepo.Touch(albumsToUpdate...) - if err != nil { - log.Error(p.ctx, "Scanner: Error touching album", "folder", entry.path, "albums", albumsToUpdate, err) - return err - } - } - - // 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(p.ctx) - enqueue := queue.Enqueue - if p.state.fullScan { - enqueue = queue.EnqueueIfMissing - } - if err := enqueue(queueItems...); err != nil { - log.Warn(p.ctx, "Scanner: could not enqueue artwork resolution", "folder", entry.path, err) - } - } - return nil + ctx := log.NewContext(p.ctx, "folder", entry.path) + err := p.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { + return p.persistFolder(ctx, tx, entry) }, "scanner: persist changes") if err != nil { - log.Error(p.ctx, "Scanner: Error persisting changes to DB", "folder", entry.path, err) + log.Error(ctx, "Scanner: Error persisting changes to DB", err) + p.stopWalk(err) + return entry, err } - return entry, err + // A new folder's albums/artists were enqueued with it; only pre-existing folders need the diff. + if !entry.isNew() { + if changed, artistImage := entry.imagesChanged(); changed { + p.imageChanges.record(entry.job.lib, imageChangedFolder{ + id: entry.id, path: entry.path, artistImage: artistImage, + }) + } + } + return entry, nil +} + +// persistFolder writes the folder in tx. WithTxRetry may rerun it after a rollback. +func (p *phaseFolders) persistFolder(ctx context.Context, tx model.DataStore, entry *folderEntry) error { + // Collect artwork queue items for changed albums/artists, enqueued in the same transaction + var queueItems []model.ArtworkQueueItem + // persistAlbum consumes the map, so a rerun needs the original + 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) + + // Save folder to DB + folder := entry.toFolder() + err := folderRepo.Put(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...) + 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", + "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) + if err != nil { + return fmt.Errorf("adding artist %q to library: %w", entry.artists[i].Name, err) + } + if entry.artists[i].Name != consts.UnknownArtist && entry.artists[i].Name != consts.VariousArtists { + queueItems = append(queueItems, scanArtworkItem(model.KindArtistArtwork, entry.artists[i].ID)) + } + } + + // Save all new/modified albums to DB. Their information will be incomplete, but they will be refreshed later + for i := range entry.albums { + err = p.persistAlbum(albumRepo, &entry.albums[i], albumIDMap) + if err != nil { + return err + } + if entry.albums[i].Name != consts.UnknownAlbum { + queueItems = append(queueItems, scanArtworkItem(model.KindAlbumArtwork, entry.albums[i].ID)) + } + } + + // Save all tracks to DB + for i := range entry.tracks { + err = mfRepo.Put(&entry.tracks[i]) + if err != nil { + return fmt.Errorf("persisting track %q: %w", entry.tracks[i].Path, err) + } + } + + // 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 { + 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...) + if err != nil { + return fmt.Errorf("marking missing tracks: %w", err) + } + + // Touch all albums that have missing tracks, so they get refreshed in later phases + groupedMissingTracks := slice.ToMap(entry.missingTracks, func(mf *model.MediaFile) (string, struct{}) { + return mf.AlbumID, struct{}{} + }) + albumsToUpdate := slices.Collect(maps.Keys(groupedMissingTracks)) + err = albumRepo.Touch(albumsToUpdate...) + if err != nil { + return fmt.Errorf("touching albums %v: %w", albumsToUpdate, err) + } + } + + // 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) + enqueue := queue.Enqueue + if p.state.fullScan { + enqueue = queue.EnqueueIfMissing + } + if err := enqueue(queueItems...); err != nil { + log.Warn(ctx, "Scanner: could not enqueue artwork resolution", err) + } + } + return nil } // persistAlbum persists the given album to the database, and reassigns annotations from the previous album ID @@ -499,34 +507,32 @@ func (p *phaseFolders) logFolder(entry *folderEntry) (*folderEntry, error) { } func (p *phaseFolders) finalize(err error) error { - errF := p.ds.WithTx(func(tx model.DataStore) error { + p.stopWalk(nil) + defer p.imageChanges.enqueue(p.ctx) + // A failed phase may not have walked every folder, and unvisited ones must not be marked missing + if err != nil { + return err + } + return p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { for _, job := range p.jobs { // Mark all folders that were not updated as missing if len(job.lastUpdates) == 0 { continue } folderIDs := slices.Collect(maps.Keys(job.lastUpdates)) - err := tx.Folder(p.ctx).MarkMissing(true, folderIDs...) - if err != nil { - log.Error(p.ctx, "Scanner: Error marking missing folders", "lib", job.lib.Name, err) - return err + if err := tx.Folder(ctx).MarkMissing(true, folderIDs...); err != nil { + return fmt.Errorf("marking missing folders in %s: %w", job.lib.Name, err) } - err = tx.MediaFile(p.ctx).MarkMissingByFolder(true, folderIDs...) - if err != nil { - log.Error(p.ctx, "Scanner: Error marking tracks in missing folders", "lib", job.lib.Name, err) - return err + if err := tx.MediaFile(ctx).MarkMissingByFolder(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 - _, err = tx.Album(p.ctx).TouchByMissingFolder() - if err != nil { - log.Error(p.ctx, "Scanner: Error touching albums with missing folders", "lib", job.lib.Name, err) - return err + if _, err := tx.Album(ctx).TouchByMissingFolder(); err != nil { + return fmt.Errorf("touching albums with missing folders in %s: %w", job.lib.Name, err) } } return nil }, "scanner: finalize phaseFolders") - p.imageChanges.enqueue(p.ctx) - return errors.Join(err, errF) } var _ phase[*folderEntry] = (*phaseFolders)(nil) diff --git a/scanner/phase_2_missing_tracks.go b/scanner/phase_2_missing_tracks.go index 8c258b833..6ccc9a46c 100644 --- a/scanner/phase_2_missing_tracks.go +++ b/scanner/phase_2_missing_tracks.go @@ -273,66 +273,68 @@ func (p *phaseMissingTracks) findCrossLibraryMatch(missing model.MediaFile) (mod } func (p *phaseMissingTracks) moveMatched(target, missing model.MediaFile) error { - return p.ds.WithTx(func(tx model.DataStore) error { - discardedID := target.ID - oldAlbumID := missing.AlbumID - newAlbumID := target.AlbumID + oldAlbumID := missing.AlbumID + newAlbumID := target.AlbumID + // Use newAlbumID as key since we only care about avoiding duplicate reassignments to the same target. + // Claimed before the transaction so a concurrent move skips it, and released if the move fails. + reassignAlbum := oldAlbumID != newAlbumID + if reassignAlbum { + p.annotationMutex.Lock() + reassignAlbum = !p.processedAlbumAnnotations[newAlbumID] + p.processedAlbumAnnotations[newAlbumID] = true + p.annotationMutex.Unlock() + if !reassignAlbum { + log.Trace(p.ctx, "Scanner: Skipping album annotation reassignment", "from", oldAlbumID, "to", newAlbumID) + } + } + err := p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { + // A rerun must start from the original target, not the one the rolled-back attempt changed + moved := target // Preserve the original created_at from the missing file, so moved tracks // don't appear in "Recently Added" - target.CreatedAt = missing.CreatedAt + moved.CreatedAt = missing.CreatedAt // 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. - target.ID = missing.ID - err := tx.MediaFile(p.ctx).Put(&target) - if err != nil { + moved.ID = missing.ID + if err := tx.MediaFile(ctx).Put(&moved); err != nil { return fmt.Errorf("update matched track: %w", err) } // Discard the new mediafile row (the one that was moved to) - err = tx.MediaFile(p.ctx).Delete(discardedID) - if err != nil { + if err := tx.MediaFile(ctx).Delete(target.ID); err != nil { return fmt.Errorf("delete discarded track: %w", err) } - // Handle album annotation reassignment if AlbumID changed - if oldAlbumID != newAlbumID { - // Use newAlbumID as key since we only care about avoiding duplicate reassignments to the same target - p.annotationMutex.RLock() - alreadyProcessed := p.processedAlbumAnnotations[newAlbumID] - p.annotationMutex.RUnlock() - - if !alreadyProcessed { - p.annotationMutex.Lock() - // Double-check pattern to avoid race conditions - if !p.processedAlbumAnnotations[newAlbumID] { - // Reassign direct album annotations (starred, rating) - log.Debug(p.ctx, "Scanner: Reassigning album annotations", "from", oldAlbumID, "to", newAlbumID) - if err := tx.Album(p.ctx).ReassignAnnotation(oldAlbumID, newAlbumID); err != nil { - log.Warn(p.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(p.ctx).CopyAttributes(oldAlbumID, newAlbumID, "created_at"); err != nil { - if !errors.Is(err, model.ErrNotFound) { - log.Warn(p.ctx, "Scanner: Could not copy album created_at", "from", oldAlbumID, "to", newAlbumID, err) - } - } - - // Note: RefreshPlayCounts will be called in later phases, so we don't need to call it here - p.processedAlbumAnnotations[newAlbumID] = true - } - p.annotationMutex.Unlock() - } else { - log.Trace(p.ctx, "Scanner: Skipping album annotation reassignment", "from", oldAlbumID, "to", newAlbumID) + 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 { + log.Warn(ctx, "Scanner: Could not reassign album annotations", "from", oldAlbumID, "to", newAlbumID, err) } - } - p.state.changesDetected.Store(true) + // 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 !errors.Is(err, model.ErrNotFound) { + log.Warn(ctx, "Scanner: Could not copy album created_at", "from", oldAlbumID, "to", newAlbumID, err) + } + } + // Note: RefreshPlayCounts will be called in later phases, so we don't need to call it here + } return nil - }) + }, "scanner: move matched track") + if err != nil { + if reassignAlbum { + p.annotationMutex.Lock() + delete(p.processedAlbumAnnotations, newAlbumID) + p.annotationMutex.Unlock() + } + return err + } + p.state.changesDetected.Store(true) + return nil } func (p *phaseMissingTracks) finalize(err error) error { @@ -355,7 +357,12 @@ func (p *phaseMissingTracks) finalize(err error) error { } func (p *phaseMissingTracks) purgeMissing() error { - deletedCount, err := p.ds.MediaFile(p.ctx).DeleteAllMissing() + 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() + return err + }, "scanner: purge missing") if err != nil { return fmt.Errorf("error deleting missing files: %w", err) } diff --git a/scanner/phase_2_missing_tracks_test.go b/scanner/phase_2_missing_tracks_test.go index f93b166c1..b7aa52f90 100644 --- a/scanner/phase_2_missing_tracks_test.go +++ b/scanner/phase_2_missing_tracks_test.go @@ -2,6 +2,8 @@ package scanner import ( "context" + "errors" + "maps" "time" "github.com/navidrome/navidrome/conf" @@ -146,6 +148,87 @@ var _ = Describe("phaseMissingTracks", func() { Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) }) + Context("claiming the album annotation reassignment", func() { + var probe *probeTxDS + 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} + BeforeEach(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) + }) + + It("claims the target album before the transaction, so a concurrent move skips it", func() { + probe.during = func() { + phase.annotationMutex.RLock() + defer phase.annotationMutex.RUnlock() + Expect(phase.processedAlbumAnnotations).To(HaveKeyWithValue("new-album", true)) + } + Expect(phase.moveMatched(matchedTrack, missingTrack)).To(Succeed()) + }) + + It("releases the claim when the move fails, so a later move can reassign", func() { + probe.err = errors.New("boom") + Expect(phase.moveMatched(matchedTrack, missingTrack)).To(MatchError("boom")) + Expect(phase.processedAlbumAnnotations).ToNot(HaveKey("new-album")) + }) + }) + + Context("when the move transaction is rerun after a busy rollback", func() { + var rerunDS *rerunTxDS + BeforeEach(func() { + rerunDS = &rerunTxDS{MockDataStore: ds.(*tests.MockDataStore)} + rerunDS.snapshot = func() func() { + saved := maps.Clone(mr.Data) + return func() { mr.Data = saved } + } + phase = createPhaseMissingTracks(ctx, state, rerunDS) + }) + + 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) + + _, err := phase.processMissingTracks(&missingTracks{ + missing: []model.MediaFile{missingTrack}, + matched: []model.MediaFile{matchedTrack}, + }) + Expect(err).ToNot(HaveOccurred()) + + movedTrack, err := ds.MediaFile(ctx).Get("1") + Expect(err).ToNot(HaveOccurred()) + Expect(movedTrack.Path).To(Equal(matchedTrack.Path)) + }) + + It("reassigns the album annotations in the attempt that commits", func() { + albumRepo := tests.CreateMockAlbumRepo() + rerunDS.MockedAlbum = albumRepo + restoreTracks := rerunDS.snapshot + rerunDS.snapshot = func() func() { + restore := restoreTracks() + return func() { + restore() + albumRepo.ReassignAnnotationCalls = nil + } + } + 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) + + _, err := phase.processMissingTracks(&missingTracks{ + missing: []model.MediaFile{missingTrack}, + matched: []model.MediaFile{matchedTrack}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(albumRepo.ReassignAnnotationCalls).To(HaveKeyWithValue("old-album", "new-album")) + }) + }) + It("should move the matched track when the missing track has the same tags and filename", 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} @@ -957,3 +1040,34 @@ var _ = Describe("phaseMissingTracks", func() { }) }) }) + +// rerunTxDS runs every WithTxRetry block twice, as a retry after a rolled-back busy attempt would. +// The mock is not transactional, so snapshot returns the function that plays the rollback. +type rerunTxDS struct { + *tests.MockDataStore + snapshot func() (rollback func()) +} + +func (d *rerunTxDS) WithTxRetry(ctx context.Context, block func(context.Context, model.DataStore) error, _ ...string) error { + rollback := d.snapshot() + _ = block(ctx, d.MockDataStore) + rollback() + return block(ctx, d.MockDataStore) +} + +// probeTxDS runs a hook inside each WithTxRetry block, and can fail the transaction after it. +type probeTxDS struct { + *tests.MockDataStore + during func() + err error +} + +func (d *probeTxDS) WithTxRetry(ctx context.Context, block func(context.Context, model.DataStore) error, _ ...string) error { + if err := block(ctx, d.MockDataStore); err != nil { + return err + } + if d.during != nil { + d.during() + } + return d.err +} diff --git a/scanner/phase_3_refresh_albums.go b/scanner/phase_3_refresh_albums.go index 33e0fed01..964ab7408 100644 --- a/scanner/phase_3_refresh_albums.go +++ b/scanner/phase_3_refresh_albums.go @@ -103,7 +103,9 @@ func (p *phaseRefreshAlbums) refreshAlbum(album *model.Album) (*model.Album, err return nil, nil } start := time.Now() - err := p.ds.Album(p.ctx).Put(album) + err := p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { + return tx.Album(ctx).Put(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 { return nil, fmt.Errorf("refreshing album %s: %w", album.ID, err) @@ -130,7 +132,12 @@ func (p *phaseRefreshAlbums) finalize(err error) error { } // Refresh album annotations start := time.Now() - cnt, err := p.ds.Album(p.ctx).RefreshPlayCounts() + 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() + return txErr + }, "scanner: refresh album play counts") if err != nil { return fmt.Errorf("refreshing album annotations: %w", err) } @@ -138,7 +145,11 @@ func (p *phaseRefreshAlbums) finalize(err error) error { // Refresh artist annotations start = time.Now() - cnt, err = p.ds.Artist(p.ctx).RefreshPlayCounts() + err = p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { + var txErr error + cnt, txErr = tx.Artist(ctx).RefreshPlayCounts() + return txErr + }, "scanner: refresh artist play counts") if err != nil { return fmt.Errorf("refreshing artist annotations: %w", err) } diff --git a/scanner/phase_4_playlists.go b/scanner/phase_4_playlists.go index baa8b749a..4e11fa81d 100644 --- a/scanner/phase_4_playlists.go +++ b/scanner/phase_4_playlists.go @@ -101,7 +101,10 @@ func (p *phasePlaylists) produce(put func(entry *model.Folder)) error { // import the playlists, and returns an error if the flag can't be persisted (so // the scan does not complete as successful without recording the recovery). func (p *phasePlaylists) deferImport() error { - if err := p.ds.Property(p.ctx).Put(consts.PlaylistsImportPendingFlagKey, "1"); err != nil { + err := p.ds.WithTxRetry(p.ctx, func(ctx context.Context, tx model.DataStore) error { + return tx.Property(ctx).Put(consts.PlaylistsImportPendingFlagKey, "1") + }, "scanner: defer playlist import") + if err != nil { return fmt.Errorf("recording pending playlist import: %w", err) } log.Warn(p.ctx, "Playlists will not be imported, as there are no admin users yet. "+ diff --git a/scanner/scanner.go b/scanner/scanner.go index d73007bdd..a8a192771 100644 --- a/scanner/scanner.go +++ b/scanner/scanner.go @@ -217,7 +217,9 @@ func (s *scannerImpl) prepareLibrariesForScan(ctx context.Context, state *scanSt for _, lib := range state.libraries { if lib.LastScanStartedAt.IsZero() { // This is a new scan - mark it as started - err := s.ds.Library(ctx).ScanBegin(lib.ID, state.fullScan) + err := s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { + return tx.Library(ctx).ScanBegin(lib.ID, state.fullScan) + }, "scanner: begin library scan") if err != nil { log.Error(ctx, "Scanner: Error marking scan start", "lib", lib.Name, err) state.sendWarning(err.Error()) @@ -253,7 +255,7 @@ func (s *scannerImpl) prepareLibrariesForScan(ctx context.Context, state *scanSt func (s *scannerImpl) runGC(ctx context.Context, state *scanState) func() error { return func() error { state.sendProgress(&ProgressInfo{ForceUpdate: true}) - return s.ds.WithTx(func(tx model.DataStore) error { + return s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { if state.changesDetected.Load() { start := time.Now() @@ -264,9 +266,7 @@ func (s *scannerImpl) runGC(ctx context.Context, state *scanState) func() error log.Debug(ctx, "Scanner: Running selective GC", "libraryIDs", libraryIDs) } - err := tx.GC(ctx, libraryIDs...) - if err != nil { - log.Error(ctx, "Scanner: Error running GC", err) + if err := tx.GC(ctx, libraryIDs...); err != nil { return fmt.Errorf("running GC: %w", err) } log.Debug(ctx, "Scanner: GC completed", "elapsed", time.Since(start)) @@ -286,10 +286,14 @@ func (s *scannerImpl) runEnqueueMissingArtwork(ctx context.Context, state *scanS return nil } start := time.Now() - queue := s.ds.ArtworkQueue(ctx) var total int64 for _, kind := range []model.Kind{model.KindAlbumArtwork, model.KindArtistArtwork} { - n, err := queue.EnqueueAllMissing(kind, model.ArtworkPriorityScan) + 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) + return err + }, "scanner: enqueue missing artwork") if err != nil { log.Error(ctx, "Scanner: Error enqueueing missing artwork", "kind", kind, err) return fmt.Errorf("enqueueing missing artwork: %w", err) @@ -316,7 +320,9 @@ func (s *scannerImpl) runRefreshStats(ctx context.Context, state *scanState) fun log.Debug(ctx, "Scanner: Refreshed artist stats", "stats", stats, "elapsed", time.Since(start)) start = time.Now() - err = s.ds.Tag(ctx).UpdateCounts() + err = s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { + return tx.Tag(ctx).UpdateCounts() + }, "scanner: update tag counts") if err != nil { log.Error(ctx, "Scanner: Error updating tag counts", err) return fmt.Errorf("updating tag counts: %w", err) @@ -329,28 +335,21 @@ func (s *scannerImpl) runRefreshStats(ctx context.Context, state *scanState) fun func (s *scannerImpl) runUpdateLibraries(ctx context.Context, state *scanState) func() error { return func() error { start := time.Now() - return s.ds.WithTx(func(tx model.DataStore) error { + return s.ds.WithTxRetry(ctx, func(ctx context.Context, tx model.DataStore) error { for _, lib := range state.libraries { - err := tx.Library(ctx).ScanEnd(lib.ID) - if err != nil { - log.Error(ctx, "Scanner: Error updating last scan completed", "lib", lib.Name, err) - return fmt.Errorf("updating last scan completed: %w", err) + if err := tx.Library(ctx).ScanEnd(lib.ID); err != nil { + return fmt.Errorf("updating last scan completed for %s: %w", lib.Name, err) } - err = tx.Property(ctx).Put(consts.PIDTrackKey, conf.Server.PID.Track) - if err != nil { - log.Error(ctx, "Scanner: Error updating track PID conf", err) + if err := tx.Property(ctx).Put(consts.PIDTrackKey, conf.Server.PID.Track); err != nil { return fmt.Errorf("updating track PID conf: %w", err) } - err = tx.Property(ctx).Put(consts.PIDAlbumKey, conf.Server.PID.Album) - if err != nil { - log.Error(ctx, "Scanner: Error updating album PID conf", err) + if err := tx.Property(ctx).Put(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 { - log.Error(ctx, "Scanner: Error refreshing library stats", "lib", lib.Name, err) - return fmt.Errorf("refreshing library stats: %w", err) + return fmt.Errorf("refreshing library stats for %s: %w", lib.Name, err) } } else { log.Debug(ctx, "Scanner: No changes detected, skipping library stats refresh", "lib", lib.Name) diff --git a/scanner/scanner_test.go b/scanner/scanner_test.go index 8542b3ac6..30f4a2b97 100644 --- a/scanner/scanner_test.go +++ b/scanner/scanner_test.go @@ -4,12 +4,16 @@ import ( "context" "database/sql" "errors" + "fmt" + "os" "path/filepath" + "sync/atomic" "testing/fstest" "time" "github.com/Masterminds/squirrel" "github.com/google/uuid" + "github.com/mattn/go-sqlite3" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" @@ -52,7 +56,10 @@ var _ = Describe("Scanner", Ordered, func() { BeforeAll(func() { ctx = request.WithUser(GinkgoT().Context(), model.User{ID: "123", IsAdmin: true}) - tmpDir := GinkgoT().TempDir() + // The DB stays open until the suite ends, and Windows can't delete an open file + tmpDir, err := os.MkdirTemp("", "scanner-test") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = os.RemoveAll(tmpDir) }) conf.Server.DbPath = filepath.Join(tmpDir, "test-scanner.db?_journal_mode=WAL") log.Warn("Using DB at " + conf.Server.DbPath) //conf.Server.DbPath = ":memory:" @@ -1240,8 +1247,55 @@ var _ = Describe("Scanner", Ordered, func() { Expect(albumArtistStats.SongCount).To(Equal(3)) // 3 songs }) }) + + Context("when the database is busy", func() { + var busyDS *busyPersistDS + BeforeEach(func() { + // One album across many folders: the suite's single DB connection deadlocks phase 3 on many albums + album := template(_t{"albumartist": "Artist", "album": "Album"}) + files := fstest.MapFS{} + for i := range 30 { + files[fmt.Sprintf("Artist/Part %02d/%02d - Song.mp3", i, i+1)] = album(track(i+1, fmt.Sprintf("Song %02d", i+1))) + } + createFS(files) + busyDS = &busyPersistDS{MockDataStore: ds} + s = scanner.New(ctx, busyDS, events.NoopBroker(), + playlists.NewPlaylists(busyDS, artwork.NewUploader(busyDS)), metrics.NewNoopInstance()) + }) + + It("gives up and stops walking the library when the database stays busy", func() { + busyDS.failures.Store(1000) + + Expect(runScanner(ctx, true)).To(MatchError(ContainSubstring("database is locked"))) + + Expect(mfRepo.cursorCalls.Load()).To(BeNumerically("<", 30)) + }) + + It("does not mark unvisited folders missing when the scan gives up", func() { + Expect(runScanner(ctx, true)).To(Succeed()) + busyDS.failures.Store(1000) + + 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()) + }) + }) }) +// busyPersistDS fails the scanner's folder saves with SQLITE_BUSY, as if WithTxRetry ran out of retries. +type busyPersistDS struct { + *tests.MockDataStore + failures atomic.Int32 +} + +func (b *busyPersistDS) WithTxRetry(ctx context.Context, block func(context.Context, model.DataStore) error, label ...string) error { + if len(label) > 0 && label[0] == "scanner: persist changes" && b.failures.Add(-1) >= 0 { + return sqlite3.Error{Code: sqlite3.ErrBusy} + } + return b.MockDataStore.WithTxRetry(ctx, block, label...) +} + 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}) @@ -1258,6 +1312,12 @@ func createFindByPath(ctx context.Context, ds model.DataStore) func(string) (*mo type mockMediaFileRepo struct { model.MediaFileRepository GetMissingAndMatchingError error + cursorCalls atomic.Int32 +} + +func (m *mockMediaFileRepo) GetCursor(options ...model.QueryOptions) (model.MediaFileCursor, error) { + m.cursorCalls.Add(1) + return m.MediaFileRepository.GetCursor(options...) } func (m *mockMediaFileRepo) GetMissingAndMatching(libId int) (model.MediaFileCursor, error) { diff --git a/scanner/walk_dir_tree.go b/scanner/walk_dir_tree.go index 20f6d4213..1864a6a44 100644 --- a/scanner/walk_dir_tree.go +++ b/scanner/walk_dir_tree.go @@ -53,6 +53,9 @@ func walkDirTree(ctx context.Context, job *scanJob, targetFolders ...string) (<- // Recursively walk this folder and all its children err = walkFolder(ctx, job, folderPath, checker, results) + if utils.IsCtxDone(ctx) { + return + } if err != nil { log.Error(ctx, "Scanner: Error walking target folder", "path", folderPath, err) continue @@ -87,9 +90,12 @@ func walkFolder(ctx context.Context, job *scanJob, currentFolder string, checker folder.path = dir folder.elapsed.Start() - results <- folder - - return nil + select { + case results <- folder: + return nil + case <-ctx.Done(): + return ctx.Err() + } } func loadDir(ctx context.Context, job *scanJob, dirPath string, checker *IgnoreChecker) (folder *folderEntry, children []string, err error) { @@ -291,6 +297,7 @@ var ignoredDirs = []string{ "$RECYCLE.BIN", "#snapshot", "@Recycle", + "@eaDir", "@Recently-Snapshot", ".git", ".streams", diff --git a/scanner/walk_dir_tree_test.go b/scanner/walk_dir_tree_test.go index 9fb650c4d..43939e5c2 100644 --- a/scanner/walk_dir_tree_test.go +++ b/scanner/walk_dir_tree_test.go @@ -564,6 +564,7 @@ var _ = Describe("walk_dir_tree", func() { Entry("dir starting with ellipsis", "...unhidden_folder", false), Entry("recycle bin", "$Recycle.Bin", true), Entry("snapshot dir", "#snapshot", true), + Entry("synology metadata dir", "@eaDir", true), ) }) diff --git a/tests/mock_data_store.go b/tests/mock_data_store.go index 32f56a4f0..6a0ebbb31 100644 --- a/tests/mock_data_store.go +++ b/tests/mock_data_store.go @@ -329,6 +329,10 @@ func (db *MockDataStore) WithTxImmediate(block func(tx model.DataStore) error, l return block(db) } +func (db *MockDataStore) WithTxRetry(ctx context.Context, block func(ctx context.Context, tx model.DataStore) error, label ...string) error { + return block(ctx, db) +} + func (db *MockDataStore) Resource(ctx context.Context, m any) model.ResourceRepository { switch m.(type) { case model.MediaFile, *model.MediaFile: From 101145742f4762164202cb4858c9126258bdc463 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Wed, 23 Sep 2026 18:42:11 -0400 Subject: [PATCH 06/26] test(plugins): stub DNS in the host SSRF guard tests (#6208) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The dial-time SSRF guard runs on the resolved IP, so the tests that prove a symbolic hostname cannot reach loopback used "localhost." — a trailing dot never matches /etc/hosts, so Go queries real DNS. Machines whose resolver does not answer "localhost." (a VPN DNS, for example) got "no such host" before the dial guard ever ran, failing three specs. Add tests.StubResolver, a net.Resolver backed by an in-memory DNS responder over net.Pipe, and let the plugin dialers take a resolver so tests can inject it. Name resolution in those specs no longer depends on the machine's DNS. --- plugins/host_httpclient.go | 1 + plugins/host_httpclient_test.go | 4 +++ plugins/host_netguard.go | 3 ++ plugins/host_websocket.go | 2 +- plugins/host_websocket_test.go | 1 + plugins/plugins_suite_test.go | 6 ++++ tests/dns_stub.go | 54 +++++++++++++++++++++++++++++++++ 7 files changed, 70 insertions(+), 1 deletion(-) create mode 100644 tests/dns_stub.go diff --git a/plugins/host_httpclient.go b/plugins/host_httpclient.go index 3265ac4b7..6a9a4d2d6 100644 --- a/plugins/host_httpclient.go +++ b/plugins/host_httpclient.go @@ -54,6 +54,7 @@ func newHTTPService(pluginName string, permission *HTTPPermission) *httpServiceI Timeout: 30 * time.Second, KeepAlive: 30 * time.Second, Control: svc.dialControl, + Resolver: dialResolver, }).DialContext // No client timeout: it is set per-request via context deadline. svc.client = &http.Client{Transport: httpclient.NewTransport(svc.transport)} diff --git a/plugins/host_httpclient_test.go b/plugins/host_httpclient_test.go index 6edf75a97..81e3192fd 100644 --- a/plugins/host_httpclient_test.go +++ b/plugins/host_httpclient_test.go @@ -20,6 +20,10 @@ var _ = Describe("httpServiceImpl", func() { ts *httptest.Server ) + BeforeEach(func() { + stubLocalhostDNS() + }) + AfterEach(func() { if ts != nil { ts.Close() diff --git a/plugins/host_netguard.go b/plugins/host_netguard.go index f78aaf6a8..ab092009a 100644 --- a/plugins/host_netguard.go +++ b/plugins/host_netguard.go @@ -8,6 +8,9 @@ import ( "github.com/navidrome/navidrome/utils/netguard" ) +// dialResolver is nil in production (the system resolver); tests swap in a stub to avoid real DNS. +var dialResolver *net.Resolver + // checkPrivateDial runs at dial time on the resolved IP, so hostnames can't reach private addresses unless a // literal IP/CIDR entry or a bare "*" (plugins targeting user-configured LAN services) allows it. func checkPrivateDial(requiredHosts []string, address string) error { diff --git a/plugins/host_websocket.go b/plugins/host_websocket.go index 933e53144..d9d82665c 100644 --- a/plugins/host_websocket.go +++ b/plugins/host_websocket.go @@ -114,7 +114,7 @@ func (s *webSocketServiceImpl) Connect(ctx context.Context, urlStr string, heade // Establish WebSocket connection dialer := websocket.Dialer{ HandshakeTimeout: 30 * time.Second, - NetDialContext: (&net.Dialer{Control: s.dialControl}).DialContext, + NetDialContext: (&net.Dialer{Control: s.dialControl, Resolver: dialResolver}).DialContext, } conn, resp, err := dialer.DialContext(ctx, urlStr, httpHeaders) diff --git a/plugins/host_websocket_test.go b/plugins/host_websocket_test.go index 772d83cc9..9b94e4b70 100644 --- a/plugins/host_websocket_test.go +++ b/plugins/host_websocket_test.go @@ -505,6 +505,7 @@ var _ = Describe("WebSocketService", Ordered, func() { var savedHosts []string BeforeEach(func() { + stubLocalhostDNS() savedHosts = testService.requiredHosts upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }} wsServer = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { diff --git a/plugins/plugins_suite_test.go b/plugins/plugins_suite_test.go index cc1b45f97..aa07eae32 100644 --- a/plugins/plugins_suite_test.go +++ b/plugins/plugins_suite_test.go @@ -46,6 +46,12 @@ func TestPlugins(t *testing.T) { RunSpecs(t, "Plugins Suite") } +// stubLocalhostDNS resolves "localhost." without real DNS: the trailing dot never matches /etc/hosts. +func stubLocalhostDNS() { + dialResolver = tests.StubResolver(map[string]string{"localhost.": "127.0.0.1"}) + DeferCleanup(func() { dialResolver = nil }) +} + // createTestManager creates a new plugin Manager with the given plugin config. // It creates a temp directory, copies the test-metadata-agent plugin, and starts the manager. // Returns the manager, temp directory path, and a cleanup function. diff --git a/tests/dns_stub.go b/tests/dns_stub.go new file mode 100644 index 000000000..7cd7faf6d --- /dev/null +++ b/tests/dns_stub.go @@ -0,0 +1,54 @@ +package tests + +import ( + "context" + "encoding/binary" + "io" + "net" + "net/netip" + + "golang.org/x/net/dns/dnsmessage" +) + +// StubResolver returns a net.Resolver that answers A queries from records (absolute name to IPv4, +// e.g. "localhost." to "127.0.0.1") instead of the system DNS. Other names do not resolve. +func StubResolver(records map[string]string) *net.Resolver { + return &net.Resolver{ + PreferGo: true, + Dial: func(context.Context, string, string) (net.Conn, error) { + client, server := net.Pipe() + go serveStubDNS(server, records) + return client, nil + }, + } +} + +// net.Pipe is not a PacketConn, so the resolver sends one query framed with a 2-byte length prefix. +func serveStubDNS(conn net.Conn, records map[string]string) { + defer conn.Close() + var size [2]byte + if _, err := io.ReadFull(conn, size[:]); err != nil { + return + } + buf := make([]byte, binary.BigEndian.Uint16(size[:])) + if _, err := io.ReadFull(conn, buf); err != nil { + return + } + var msg dnsmessage.Message + if err := msg.Unpack(buf); err != nil || len(msg.Questions) != 1 { + return + } + msg.Response = true + q := msg.Questions[0] + if ip, err := netip.ParseAddr(records[q.Name.String()]); err == nil && q.Type == dnsmessage.TypeA { + msg.Answers = []dnsmessage.Resource{{ + Header: dnsmessage.ResourceHeader{Name: q.Name, Type: q.Type, Class: q.Class}, + Body: &dnsmessage.AResource{A: ip.As4()}, + }} + } + resp, err := msg.Pack() + if err != nil { + return + } + _, _ = conn.Write(append(binary.BigEndian.AppendUint16(nil, uint16(len(resp))), resp...)) +} From 69b496383deb9e58e3930da7628ab563af938946 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Thu, 24 Sep 2026 21:24:46 -0400 Subject: [PATCH 07/26] docs(jellyfin): refresh the README's known limitations (#6220) * docs(jellyfin): refresh the README's known limitations Rewrite the Known limitations list against the current code and Jellyfin 12.1, dropping stale entries (synthetic blurhash, unchecked artist access, global genres) and adding the real gaps: search skipping filters, the one-character minimum, the 2,000-item search cap, position-based playlist entry ids, the rating param, missing endpoints and unemitted Fields. Also move the lyrics description into its own section, add the missing routes and filters to the endpoint table, and drop the stale artist-access TODO in resolveItemByID. * docs(jellyfin): list MaxConcurrentStreams env var and all e2e stubs --- server/jellyfin/README.md | 109 ++++++++++++++++++++++---------------- server/jellyfin/items.go | 3 +- 2 files changed, 63 insertions(+), 49 deletions(-) diff --git a/server/jellyfin/README.md b/server/jellyfin/README.md index e38997569..9315ac4f9 100644 --- a/server/jellyfin/README.md +++ b/server/jellyfin/README.md @@ -40,6 +40,7 @@ ND_JELLYFIN_SERVERNAME="My Music Server" ND_JELLYFIN_EXPOSEDPUBLICUSERS="alice,bob" ND_JELLYFIN_AUTODISCOVERY=true ND_JELLYFIN_QUICKCONNECT=false +ND_JELLYFIN_MAXCONCURRENTSTREAMS=4 ``` Once enabled, the API is mounted at: @@ -140,8 +141,10 @@ Jellyfin has no native concept of multiple music libraries the way Navidrome doe Navidrome library the current user can access is exposed as its own top-level Jellyfin "CollectionFolder" view (`GET /UserViews`), instead of merging every library into a single view. Browsing (`/Items`), artists, and the "Latest" list are all scoped to the libraries the -authenticated user has access to; a library (or item within it) the user cannot access returns -`404`, never `403`, so ids can't be used as an existence oracle. +authenticated user has access to. Fetching an item the user cannot access returns `404`, never +`403`, so ids can't be used as an existence oracle. An inaccessible library id sent as `ParentId` +is not a `404`: it is simply not treated as a library, so none of that library's content is +returned. ### Browsing filters @@ -152,7 +155,7 @@ tracks — Finamp's artist screen sends these *alongside* `ParentId=` album's tracks — Feishin fetches them this way instead of `ParentId`); `GenreIds` (a genre's albums or tracks — Finamp's genre screen sends it the same way; `/Artists/AlbumArtists` and `MusicArtist` queries accept it too, matching artists credited on an album of that genre); -`SearchTerm`; +`Years`; `StudioIds` (record labels, as listed by `GET Studios`); `SearchTerm`; `Filters` (`IsFavorite`, `IsFavoriteOrLikes`, `IsPlayed`, `IsUnplayed`) and the standalone `isFavorite`/`isPlayed` booleans it can also be expressed as — `Filters` wins when both are sent, as in Jellyfin; `Likes`, `Dislikes`, `IsFolder`, `IsNotFolder` and `IsResumable` have no Navidrome @@ -169,9 +172,9 @@ returns direct children only (no tracks — no track is a library's direct child | Quick Connect | `GET QuickConnect/Enabled`, `POST QuickConnect/Initiate`, `GET QuickConnect/Connect`, `POST QuickConnect/Authorize` (authenticated), `POST Users/AuthenticateWithQuickConnect` | | Auth | `POST Users/AuthenticateByName`, `GET Users/Public` | | Users | `GET UserViews`, `GET Users/{userId}/Views`, `GET Users/Me`, `GET Users/{userId}` | -| Browsing | `GET Items`, `GET Users/{userId}/Items`, `GET Items/{itemId}`, `GET Users/{userId}/Items/{itemId}`, `GET Users/{userId}/Items/Latest`, `DELETE Items/{itemId}` (playlists only) | -| Artists / genres | `GET Artists`, `GET Artists/AlbumArtists`, `GET Genres`, `GET MusicGenres` | -| Similar / mixes | `GET Artists/{itemId}/Similar`, `GET Items/{itemId}/Similar`, `GET {Items,Songs,Albums,Artists,Playlists}/{itemId}/InstantMix`, `GET Artists/InstantMix?id=`, `GET MusicGenres/InstantMix?id=` | +| Browsing | `GET Items`, `GET Users/{userId}/Items`, `GET Items/{itemId}`, `GET Users/{userId}/Items/{itemId}`, `GET Items/Latest`, `GET Users/{userId}/Items/Latest`, `DELETE Items/{itemId}` (playlists only) | +| Artists / genres / labels | `GET Artists`, `GET Artists/AlbumArtists`, `GET Genres`, `GET MusicGenres`, `GET Studios`, `GET Items/Filters` | +| Similar / mixes | `GET Artists/{itemId}/Similar`, `GET Items/{itemId}/Similar`, `GET Albums/{itemId}/Similar`, `GET {Items,Songs,Albums,Artists,Playlists}/{itemId}/InstantMix`, `GET Artists/InstantMix?id=`, `GET MusicGenres/InstantMix?id=` | | Images | `GET Items/{itemId}/Images/{type}[/{index}]` (public), `POST`/`DELETE Items/{itemId}/Images/{type}` (playlist cover, authenticated) | | Favorites / ratings for songs, albums, artists, and playlists | `POST`/`DELETE UserFavoriteItems/{itemId}`, `POST`/`DELETE Users/{userId}/FavoriteItems/{itemId}`, `POST`/`DELETE Users/{userId}/Items/{itemId}/Rating`, `GET UserItems/{itemId}/UserData`, `GET Users/{userId}/Items/{itemId}/UserData` | | Streaming | `GET Audio/{itemId}/stream[.{container}]`, `GET Audio/{itemId}/universal`, `GET Audio/{itemId}/main.m3u8`, `GET Items/{itemId}/File`, `GET Items/{itemId}/Download`, `GET`/`POST Items/{itemId}/PlaybackInfo` (`HEAD` too on stream, universal, File, Download and images; a transcode HEAD answers without starting it) | @@ -267,6 +270,22 @@ The stream endpoints reuse the same transcode-decision pipeline as the Subsonic Subsonic. `File`/`Download` stay raw. For HLS clients, force `aac` or `mp3`; other formats are advertised and served but packed-audio players won't decode them. +## Lyrics + +`GET Audio/{id}/Lyrics` serves the main lyric track as a `LyricDto` (`Start` in 100ns ticks, +word-level `Cues` when present), resolved through the full `core/lyrics` pipeline (embedded, `.lrc` +sidecars, plugins per `LyricsPriority`). No lyrics returns `404`, never an empty `200`. + +Results, misses included, are cached for 5 minutes: Jellify fetches lyrics for every played track and +Feishin on every song change, so lyric-less tracks are the hot path. Concurrent misses on the same +track share one pipeline run. That run is detached from the request, so a cancelled request doesn't +fail it for other waiters, and its context has a one-minute deadline. + +Finamp opens its lyrics view only when the track has a `Lyric` `MediaStream` (`HasLyrics` is just a +list badge). Browse lists set both from embedded lyrics only. `PlaybackInfo` runs the full pipeline +per track, so sidecar and plugin lyrics show up there. Feishin also requires server version ≥ 10.9 +(we advertise 12.1.0). + ## AudioMuse-AI compatible endpoints Compatibility shim for Jellyfin front-ends that integrate [AudioMuse-AI](https://github.com/NeptuneHub/audiomuse-ai-plugin) @@ -356,8 +375,9 @@ curl -s -X DELETE "${AUTH[@]}" "$BASE/Items/$PLAYLIST_ID" Handler-level unit tests live alongside each file (`*_test.go`). A full end-to-end suite in [`e2e/`](e2e) exercises every endpoint through the real router against a real SQLite database and -real repositories (only artwork/streaming/ffmpeg are stubbed), with per-`Describe` snapshot -isolation — mirroring the Subsonic `server/subsonic/e2e` suite. Run it with: +real repositories (only artwork, streaming, ffmpeg, external metadata agents and sonic similarity +are stubbed), with per-`Describe` snapshot isolation — mirroring the Subsonic `server/subsonic/e2e` +suite. Run it with: ```bash make test PKG=./server/jellyfin/... @@ -365,42 +385,37 @@ make test PKG=./server/jellyfin/... ## Known limitations -- **Genres are global.** `GET Genres`/`MusicGenres` is not scoped to the current user's - libraries (genre tags aren't per-library entities in Navidrome's model). -- **Artist item-access relies on list-time scoping.** Unlike albums and songs (which each - belong to exactly one library and are checked against `user.HasLibraryAccess` on every - fetch), an artist can have content across multiple libraries via `library_artist`, so there's - no single library id to gate a direct `GET Items/{artistId}` or favorite/rating call against. - Access control for artists is enforced by scoping the `Artists`/`Items?IncludeItemTypes=MusicArtist` - *list* to the user's libraries, plus the persistence layer's own defense-in-depth; a client - that already has an artist id from elsewhere is not re-checked against library membership. -- **Blurhashes are synthetic, not computed from the artwork (follow-up).** `ImageBlurHashes` is - populated by `dto/blurhash.go`, which derives a well-formed **1-component (solid color)** - blurhash by hashing the item id — it never looks at the actual image. Real Jellyfin computes a - multi-component blurhash from the cover's pixels (downscaled to 128×128) once at scan time and - stores it per image, so its placeholder approximates the art. Ours satisfies the protocol - (Finamp gets a valid value to use as a de-dup key and a placeholder, no missing-blurhash - warning) but renders as a flat color while art loads. A proper implementation would compute the - real blurhash in the `core/artwork` pipeline (where the image is already decoded), cache it - keyed like the artwork, and have the mappers read it — keeping the synthetic value as a fallback - for art that hasn't been rendered yet. -- **The WebSocket only keep-alives; it pushes no events (follow-up).** `GET socket` sends a - `ForceKeepAlive` and answers `KeepAlive` pings so real-time clients (Finamp) settle into a - working session instead of 404-loop-reconnecting, but it never pushes anything. A follow-up - would broadcast real session/playstate and library-change events over it (via `server/events`), - mirroring Jellyfin's session messages. -- **Lyrics.** `GET Audio/{id}/Lyrics` serves the main lyric track as a `LyricDto` (`Start` in - 100ns ticks, word-level `Cues` when present), resolved through the full `core/lyrics` pipeline - (embedded, `.lrc` sidecars, plugins per `LyricsPriority`) behind a 5-minute TTL cache that also - caches misses — Jellify fetches for every played track, Feishin per song change, so lyric-less - tracks are the hot path. No lyrics → 404 (never an empty 200), which all three clients degrade - gracefully. Finamp gates its lyrics view on a `Lyric` `MediaStream` (not `HasLyrics`, which is - just a list badge): browse lists advertise it from embedded lyrics only (the `"[]"` sentinel - check — the column is never `""` post-scan), while `PlaybackInfo` runs the full pipeline per - track so sidecar/plugin lyrics also light up. Feishin additionally requires server version - ≥ 10.9 (`jellyfinVersion` advertises 12.1.0). - Concurrent misses on the same track share one pipeline invocation (`SimpleCache.GetWithLoader` - is singleflighted), and the load runs detached from the request context with a one-minute bound, - so a cancelled request or hung plugin can't fail or pin the load for other waiters. - Follow-up: tracks whose only lyrics are sidecar/plugin-sourced show no `HasLyrics` badge in - lists (request-time sources can't be known at list time without per-row I/O). +- **Genre lists ignore `ParentId`.** `GET Genres`/`MusicGenres` and + `Items?IncludeItemTypes=MusicGenre` list the genres of every library the user can access, never + only the `ParentId` library. `Items?IncludeItemTypes=MusicGenre` also ignores `SearchTerm`. + `GET Items/Filters` is the exception: its genres follow `ParentId`. +- **Search skips some filters.** With `SearchTerm`, `Filters=IsFavorite`, `IsPlayed`, `IsUnplayed` + and `isFavorite`/`isPlayed` are skipped, so the response holds every search match. The full-text + search's first phase has no annotation join to filter on. Artist searches also skip the role (album + artist vs. artist) and `GenreIds`. Playlist lists ignore `SearchTerm`. +- **Search pages are capped at 2,000 items.** A larger `Limit` is lowered to 2,000. Jellyfin has no + cap. +- **One-character searches return nothing.** The shared search layer ignores terms shorter than two + characters, so album, artist and song searches for `a` come back empty. Real Jellyfin has no + minimum. +- **Playlist entry ids are positions, not song ids.** `PlaylistItemId` encodes the entry's position + in the playlist, so ids change when entries are inserted or removed; re-read the list after an + edit. Real Jellyfin uses the song id. Clients that echo `PlaylistItemId` back (Finamp) work; + clients that send the song id as `EntryIds` or to `Move` (Jellify) get `404`. +- **Rating uses the `rating` param.** `POST Users/{userId}/Items/{itemId}/Rating` reads a `rating` + (0-10) and stores it as 0-5 stars. Jellyfin's `likes` boolean is not read, so a client that sends + only `likes` clears the rating. +- **Missing endpoints.** `UserPlayedItems` (mark played/unplayed), `Search/Hints` and the non-legacy + `UserItems/{itemId}/Rating` are not implemented. Play counts only change through playback + reporting. +- **Some `Fields` are never emitted.** `ProviderIds`, `People`, `Etag` and `DateLastMediaAdded` are + accepted but never returned. +- **Blurhashes can be missing.** `ImageBlurHashes` carries the real blurhash that the artwork + pipeline computes and stores. Until that happens for an item, the field is omitted: a fake value + would pin the wrong placeholder in clients that key their image cache on it. +- **The WebSocket only keeps alive; it pushes no events.** `GET socket` sends `ForceKeepAlive` and + answers `KeepAlive` pings, so real-time clients (Finamp) keep a working session instead of + reconnecting in a loop. It never pushes session, playstate or library-change events. +- **List lyric badges only see embedded lyrics.** Tracks with only sidecar or plugin lyrics show no + `HasLyrics` badge in lists, because those sources can't be known at list time without per-row I/O. + `PlaybackInfo` and `Audio/{id}/Lyrics` still find them (see "Lyrics"). diff --git a/server/jellyfin/items.go b/server/jellyfin/items.go index 124d17920..23e35dd59 100644 --- a/server/jellyfin/items.go +++ b/server/jellyfin/items.go @@ -893,8 +893,7 @@ func (api *Router) resolveItemByID(ctx context.Context, id string, fields dto.Fi return dto.AlbumToBaseItem(*al, fields), true } if ar, err := api.ds.Artist(ctx).Get(id); err == nil { - // TODO: an artist spans multiple libraries (library_artist), so there's no single - // LibraryID to gate here; artist access relies on list-time scoping and persistence. + // 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 { From a0c416037655a2d4476d3e055907540605f5bf87 Mon Sep 17 00:00:00 2001 From: Deluan Date: Thu, 24 Sep 2026 21:40:49 -0400 Subject: [PATCH 08/26] chore: update golangci-lint version to v2.14.0 --- Makefile | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Makefile b/Makefile index 81a609422..eccccbffb 100644 --- a/Makefile +++ b/Makefile @@ -20,7 +20,7 @@ IMAGE_PLATFORMS ?= $(shell echo $(SUPPORTED_PLATFORMS) | tr ',' '\n' | grep "lin PLATFORMS ?= $(SUPPORTED_PLATFORMS) DOCKER_TAG ?= deluan/navidrome:develop -GOLANGCI_LINT_VERSION ?= v2.13.2 +GOLANGCI_LINT_VERSION ?= v2.14.0 UI_SRC_FILES := $(shell find ui -type f -not -path "ui/build/*" -not -path "ui/node_modules/*") From 3f89baaec8f3285429496d5f0b796165034c70ed Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Thu, 24 Sep 2026 22:07:17 -0400 Subject: [PATCH 09/26] test(server): make the handleM3U spec independent of spec order (#6221) It relied on another spec having set auth.PublicTokenAuth, so it panicked whenever Ginkgo ran it first. --- server/public/handle_shares_test.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/server/public/handle_shares_test.go b/server/public/handle_shares_test.go index 1bf631fd4..4de2a75c1 100644 --- a/server/public/handle_shares_test.go +++ b/server/public/handle_shares_test.go @@ -4,7 +4,9 @@ import ( "net/http" "net/http/httptest" + "github.com/go-chi/jwtauth/v5" "github.com/navidrome/navidrome/core" + "github.com/navidrome/navidrome/core/auth" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" @@ -17,6 +19,7 @@ var _ = Describe("handleM3U", func() { var pub *Router BeforeEach(func() { + auth.PublicTokenAuth = jwtauth.New("HS256", []byte("test-secret"), nil) ds = &tests.MockDataStore{} shareRepo = &tests.MockShareRepo{} ds.MockedShare = shareRepo From 336fa482f84f2600f7cf79464fd828549e4aa43b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Fri, 25 Sep 2026 12:19:33 -0400 Subject: [PATCH 10/26] fix(plugins): skip the startup error reset when no plugin has an error (#6223) The plugin manager clears last_error on every startup with an UPDATE that takes the SQLite write lock even when no row matches. On slow storage the startup scan often holds the lock at that moment, so the reset waited out the busy timeout and logged "database is locked", even with no plugins installed. ClearErrors now checks for errors with a read first and only writes when there is something to clear. --- persistence/plugin_repository.go | 8 ++++++++ persistence/plugin_repository_test.go | 15 +++++++++++++++ 2 files changed, 23 insertions(+) diff --git a/persistence/plugin_repository.go b/persistence/plugin_repository.go index c1e36f0b1..7d5781f49 100644 --- a/persistence/plugin_repository.go +++ b/persistence/plugin_repository.go @@ -35,6 +35,14 @@ func (r *pluginRepository) ClearErrors() error { if !r.isPermitted() { return rest.ErrPermissionDenied } + // An UPDATE takes the write lock even when nothing matches, so only run it when there is an error to clear + var hasErrors bool + if err := r.db.NewQuery("SELECT EXISTS (SELECT 1 FROM plugin WHERE last_error != '')").Row(&hasErrors); err != nil { + return err + } + if !hasErrors { + return nil + } _, err := r.db.NewQuery("UPDATE plugin SET last_error = '' WHERE last_error != ''").Execute() return err } diff --git a/persistence/plugin_repository_test.go b/persistence/plugin_repository_test.go index dc68b0892..9b135057e 100644 --- a/persistence/plugin_repository_test.go +++ b/persistence/plugin_repository_test.go @@ -1,7 +1,10 @@ package persistence import ( + "context" + "github.com/deluan/rest" + "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" . "github.com/onsi/ginkgo/v2" @@ -198,6 +201,18 @@ var _ = Describe("PluginRepository", func() { err := repo.ClearErrors() 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"}) + conn, err := db.Db().Conn(GinkgoT().Context()) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(conn.Close) + _, err = conn.ExecContext(GinkgoT().Context(), "BEGIN IMMEDIATE") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _, _ = conn.ExecContext(context.Background(), "ROLLBACK") }) + + Expect(repo.ClearErrors()).To(Succeed()) + }) }) }) From 659d067abaaf31d4a261836331a1f6bc00746c88 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Fri, 25 Sep 2026 16:19:30 -0400 Subject: [PATCH 11/26] fix(archiver): give same-named albums their own folder in artist zips (#6225) * refactor: add Tags.First and slice.GroupOrdered helpers Tags.First returns the first value of a tag or an empty string, replacing the inline len-check-then-index pattern in FullTitle, FullAlbumName, Album.FullName and the Subsonic album version mapping. slice.GroupOrdered is slice.Group returning the groups in first-seen order, for callers that need a deterministic order the map-based Group cannot give. * fix(archiver): give same-named albums their own folder in artist zips Artist zips put every album in a folder named after the album, so two albums with the same name (an original and a deluxe edition, or names that only differ in characters the sanitizer replaces) were merged into one folder, with tracks mixed together and duplicate zip entries when file names collided. The folder is now named after FullAlbumName(), so with AppendAlbumVersion on (the default) the version is part of the name, matching what clients display. Albums whose sanitized names still clash get a " [suffix]" taken from the first field that has a distinct, non-empty value for all of them: album version, year, release type, record label, catalog number, then a short album id. This follows the shape of beets' %aunique{} path function. Albums are also grouped with slice.GroupOrdered instead of a map, so the zip is deterministic. * fix(archiver): use the release year to tell same-named albums apart Taggers often write an edition's date to the Date tag next to an original date, and the scanner then stores the original year in Year and the edition's year in ReleaseYear. Reissues of the same album therefore share Year, so the year disambiguator could not tell them apart and they fell through to the album id suffix. Prefer ReleaseYear and fall back to Year when it is not set. Found by downloading an artist zip from a live server built from this branch. * fix(archiver): let one clashing album keep the plain folder name A disambiguator was only accepted when every clashing album had a non-empty value, so an original and its deluxe edition (with the version not appended to the name) fell through to the album id suffix. Accept a field whose values are distinct across the group even when one of them is empty, as beets' %aunique{} does: that album keeps the plain name, which the suffixed folders cannot clash with. Two or more empty values still count as a tie. * test(archiver): refactor tests for album naming conventions and query order --- core/archiver.go | 70 +++++++++++++++++-- core/archiver_test.go | 136 +++++++++++++++++++++++++++++++++++++ model/album.go | 4 +- model/mediafile.go | 8 +-- model/tag.go | 7 ++ model/tag_test.go | 4 ++ server/subsonic/helpers.go | 4 +- utils/slice/slice.go | 17 +++++ utils/slice/slice_test.go | 12 ++++ 9 files changed, 246 insertions(+), 16 deletions(-) diff --git a/core/archiver.go b/core/archiver.go index c9436279e..6406ca075 100644 --- a/core/archiver.go +++ b/core/archiver.go @@ -2,12 +2,14 @@ package core import ( "archive/zip" + "cmp" "context" "errors" "fmt" "io" "os" "path/filepath" + "strconv" "strings" "github.com/Masterminds/squirrel" @@ -58,16 +60,16 @@ func (a *archiver) zipAlbums(ctx context.Context, id string, format string, bitr } z := createZipWriter(out, format, bitrate) - albums := slice.Group(mfs, func(mf model.MediaFile) string { - return mf.AlbumID - }) + albums := slice.GroupOrdered(mfs, func(mf model.MediaFile) string { return mf.AlbumID }) + folders := albumFolders(albums) for _, album := range albums { discs := slice.Group(album, func(mf model.MediaFile) int { return mf.DiscNumber }) isMultiDisc := len(discs) > 1 - log.Debug(ctx, "Zipping album", "name", album[0].Album, "artist", album[0].AlbumArtist, + folder := folders[album[0].AlbumID] + log.Debug(ctx, "Zipping album", "name", album[0].Album, "artist", album[0].AlbumArtist, "folder", folder, "format", format, "bitrate", bitrate, "isMultiDisc", isMultiDisc, "numTracks", len(album)) for _, mf := range album { - file := a.albumFilename(mf, format, isMultiDisc) + file := a.albumFilename(mf, format, isMultiDisc, folder) if addErr := a.addFileToZip(ctx, z, mf, format, bitrate, file); errors.Is(addErr, stream.ErrTooManyTranscodes) { // Stop iterating: continuing would just rack up more // rejections from the limiter. Close finalises whatever @@ -96,7 +98,61 @@ func createZipWriter(out io.Writer, format string, bitrate int) *zip.Writer { return z } -func (a *archiver) albumFilename(mf model.MediaFile, format string, isMultiDisc bool) string { +// Tried in order; the first one whose values are distinct across the clashing albums wins. +// One album may have an empty value: it keeps the plain name, which the others can't clash with. +var albumDisambiguators = []func(model.MediaFile) string{ + func(mf model.MediaFile) string { return mf.Tags.First(model.TagAlbumVersion) }, + func(mf model.MediaFile) string { + // Reissues share Year (often the original's) but not ReleaseYear. + if y := cmp.Or(mf.ReleaseYear, mf.Year); y != 0 { + return strconv.Itoa(y) + } + return "" + }, + func(mf model.MediaFile) string { return mf.MbzAlbumType }, + func(mf model.MediaFile) string { return mf.Tags.First(model.TagRecordLabel) }, + func(mf model.MediaFile) string { return mf.CatalogNum }, + func(mf model.MediaFile) string { return mf.AlbumID[:min(6, len(mf.AlbumID))] }, + func(mf model.MediaFile) string { return mf.AlbumID }, +} + +// albumFolders maps each album id to its zip folder. Albums whose names sanitize to the +// same folder get a " [suffix]" from the first disambiguator that tells them all apart. +func albumFolders(albums [][]model.MediaFile) map[string]string { + byName := map[string][]model.MediaFile{} + for _, album := range albums { + name := str.SanitizeFilename(album[0].FullAlbumName()) + byName[name] = append(byName[name], album[0]) + } + folders := make(map[string]string, len(albums)) + for name, group := range byName { + if len(group) == 1 { + folders[group[0].AlbumID] = name + continue + } + fields: + for _, field := range albumDisambiguators { + ids := make(map[string]string, len(group)) // suffix -> album id + for _, mf := range group { + s := str.SanitizeFilename(field(mf)) + if _, dup := ids[s]; dup { + continue fields + } + ids[s] = mf.AlbumID + } + for s, id := range ids { + folders[id] = name + if s != "" { + folders[id] = fmt.Sprintf("%s [%s]", name, s) + } + } + break + } + } + return folders +} + +func (a *archiver) albumFilename(mf model.MediaFile, format string, isMultiDisc bool, folder string) string { _, file := filepath.Split(mf.Path) if format != "raw" { file = strings.TrimSuffix(file, mf.Suffix) + format @@ -104,7 +160,7 @@ func (a *archiver) albumFilename(mf model.MediaFile, format string, isMultiDisc if isMultiDisc { file = fmt.Sprintf("Disc %02d/%s", mf.DiscNumber, file) } - return fmt.Sprintf("%s/%s", str.SanitizeFilename(mf.Album), file) + return fmt.Sprintf("%s/%s", folder, file) } // ZipShare takes an already-loaded share: Share.Load records a visit, so diff --git a/core/archiver_test.go b/core/archiver_test.go index 461af1800..86833717a 100644 --- a/core/archiver_test.go +++ b/core/archiver_test.go @@ -8,6 +8,8 @@ import ( "strings" "github.com/Masterminds/squirrel" + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/core" "github.com/navidrome/navidrome/core/stream" "github.com/navidrome/navidrome/model" @@ -91,6 +93,140 @@ var _ = Describe("Archiver", func() { Expect(zr.File[0].Name).To(Equal("Album 1/01 - track1.mp3")) Expect(zr.File[1].Name).To(Equal("Album 1/02 - track2.mp3")) }) + + When("albums that share a name", func() { + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + }) + + // zipArtistEntries zips the given tracks as artist "1" and returns the entry names in zip order. + zipArtistEntries := func(mfs model.MediaFiles) []string { + mfRepo := &mockMediaFileRepository{} + mfRepo.On("GetAll", mock.Anything).Return(mfs, nil) + ds.On("MediaFile", mock.Anything).Return(mfRepo) + ms.On("NewStream", mock.Anything, mock.Anything, mock.Anything).Return(io.NopCloser(strings.NewReader("test")), nil) + + out := new(bytes.Buffer) + Expect(arch.ZipArtist(context.Background(), "1", "mp3", 128, out)).To(Succeed()) + zr, err := zip.NewReader(bytes.NewReader(out.Bytes()), int64(out.Len())) + Expect(err).To(BeNil()) + names := make([]string, len(zr.File)) + for i, f := range zr.File { + names[i] = f.Name + } + return names + } + + It("keeps the albums in query order", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "3", Album: "Album C"}, + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Album A"}, + {Path: "a/02.mp3", Suffix: "mp3", AlbumID: "1", Album: "Album A"}, + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Album B"}, + }) + Expect(names).To(Equal([]string{"Album C/01.mp3", "Album A/01.mp3", "Album A/02.mp3", "Album B/01.mp3"})) + }) + + It("suffixes the year when it tells the albums apart", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01 - Intro.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001}, + {Path: "b/01 - Intro.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 2005}, + }) + Expect(names).To(Equal([]string{"Greatest Hits [2001]/01 - Intro.mp3", "Greatest Hits [2005]/01 - Intro.mp3"})) + }) + + It("prefers the release year, so reissues of the same original are told apart", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 1996, ReleaseYear: 2001}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 1996, ReleaseYear: 2011}, + {Path: "c/01.mp3", Suffix: "mp3", AlbumID: "3", Album: "Greatest Hits", Year: 1996}, + }) + Expect(names).To(Equal([]string{"Greatest Hits [2001]/01.mp3", "Greatest Hits [2011]/01.mp3", "Greatest Hits [1996]/01.mp3"})) + }) + + It("names the folder after the full album name", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001, + Tags: model.Tags{model.TagAlbumVersion: {"Original"}}}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 2005, + Tags: model.Tags{model.TagAlbumVersion: {"CD/Digital"}}}, + }) + Expect(names).To(Equal([]string{"Greatest Hits (Original)/01.mp3", "Greatest Hits (CD_Digital)/01.mp3"})) + }) + + It("prefers the album version over the year when it is not part of the name", func() { + conf.Server.Subsonic.AppendAlbumVersion = false + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001, + Tags: model.Tags{model.TagAlbumVersion: {"Original"}}}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 2005, + Tags: model.Tags{model.TagAlbumVersion: {"Deluxe Edition"}}}, + }) + Expect(names).To(Equal([]string{"Greatest Hits [Original]/01.mp3", "Greatest Hits [Deluxe Edition]/01.mp3"})) + }) + + It("leaves the one album without the field unsuffixed", func() { + conf.Server.Subsonic.AppendAlbumVersion = false + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 2005, + Tags: model.Tags{model.TagAlbumVersion: {"Deluxe Edition"}}}, + }) + Expect(names).To(Equal([]string{"Greatest Hits/01.mp3", "Greatest Hits [Deluxe Edition]/01.mp3"})) + }) + + It("skips a field that is empty on more than one album", func() { + conf.Server.Subsonic.AppendAlbumVersion = false + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 2005}, + {Path: "c/01.mp3", Suffix: "mp3", AlbumID: "3", Album: "Greatest Hits", Year: 2010, + Tags: model.Tags{model.TagAlbumVersion: {"Deluxe Edition"}}}, + }) + Expect(names).To(Equal([]string{"Greatest Hits [2001]/01.mp3", "Greatest Hits [2005]/01.mp3", "Greatest Hits [2010]/01.mp3"})) + }) + + It("skips a field that is the same on every album", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Live", Year: 2001, MbzAlbumType: "album", CatalogNum: "CAT-1"}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Live", Year: 2001, MbzAlbumType: "album", CatalogNum: "CAT-2"}, + }) + Expect(names).To(Equal([]string{"Live [CAT-1]/01.mp3", "Live [CAT-2]/01.mp3"})) + }) + + It("falls back to the album id when nothing differs", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "0123456789abcdef", Album: "Greatest Hits", Year: 2001}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "fedcba9876543210", Album: "Greatest Hits", Year: 2001}, + }) + Expect(names).To(Equal([]string{"Greatest Hits [012345]/01.mp3", "Greatest Hits [fedcba]/01.mp3"})) + }) + + It("treats names that sanitize to the same folder as a clash", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "A/B", Year: 2001}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: `A\B`, Year: 2005}, + }) + Expect(names).To(Equal([]string{"A_B [2001]/01.mp3", "A_B [2005]/01.mp3"})) + }) + + It("sanitizes the suffix", func() { + conf.Server.Subsonic.AppendAlbumVersion = false + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Hits", Tags: model.Tags{model.TagAlbumVersion: {"Vinyl"}}}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Hits", Tags: model.Tags{model.TagAlbumVersion: {"CD/Digital"}}}, + }) + Expect(names).To(Equal([]string{"Hits [Vinyl]/01.mp3", "Hits [CD_Digital]/01.mp3"})) + }) + + It("leaves the folder name alone when only one album has it", func() { + names := zipArtistEntries(model.MediaFiles{ + {Path: "a/01.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001}, + {Path: "b/01.mp3", Suffix: "mp3", AlbumID: "2", Album: "Other", Year: 2005}, + }) + Expect(names).To(Equal([]string{"Greatest Hits/01.mp3", "Other/01.mp3"})) + }) + }) }) Context("when the transcode limiter rejects a file", func() { diff --git a/model/album.go b/model/album.go index ee24bfa96..114d19e1e 100644 --- a/model/album.go +++ b/model/album.go @@ -76,8 +76,8 @@ func (a Album) CoverArtID() ArtworkID { } func (a Album) FullName() string { - if conf.Server.Subsonic.AppendAlbumVersion && len(a.Tags[TagAlbumVersion]) > 0 { - return appendSuffix(a.Name, a.Tags[TagAlbumVersion][0]) + if v := a.Tags.First(TagAlbumVersion); conf.Server.Subsonic.AppendAlbumVersion && v != "" { + return appendSuffix(a.Name, v) } return a.Name } diff --git a/model/mediafile.go b/model/mediafile.go index 7cbdb583c..1147eaa35 100644 --- a/model/mediafile.go +++ b/model/mediafile.go @@ -105,15 +105,15 @@ type MediaFile struct { } func (mf MediaFile) FullTitle() string { - if conf.Server.Subsonic.AppendSubtitle && len(mf.Tags[TagSubtitle]) > 0 { - return appendSuffix(mf.Title, mf.Tags[TagSubtitle][0]) + if s := mf.Tags.First(TagSubtitle); conf.Server.Subsonic.AppendSubtitle && s != "" { + return appendSuffix(mf.Title, s) } return mf.Title } func (mf MediaFile) FullAlbumName() string { - if conf.Server.Subsonic.AppendAlbumVersion && len(mf.Tags[TagAlbumVersion]) > 0 { - return appendSuffix(mf.Album, mf.Tags[TagAlbumVersion][0]) + if v := mf.Tags.First(TagAlbumVersion); conf.Server.Subsonic.AppendAlbumVersion && v != "" { + return appendSuffix(mf.Album, v) } return mf.Album } diff --git a/model/tag.go b/model/tag.go index 234cfb359..152b7164e 100644 --- a/model/tag.go +++ b/model/tag.go @@ -80,6 +80,13 @@ func (t Tags) Values(name TagName) []string { return t[name] } +func (t Tags) First(name TagName) string { + if v := t[name]; len(v) > 0 { + return v[0] + } + return "" +} + func (t Tags) IDs() []string { var ids []string for name, tag := range t { diff --git a/model/tag_test.go b/model/tag_test.go index 4dc99019b..c5a941b4e 100644 --- a/model/tag_test.go +++ b/model/tag_test.go @@ -43,6 +43,10 @@ var _ = Describe("Tag", func() { Expect(tags.Values("genre")).To(ConsistOf("Rock", "Pop")) Expect(tags.Values("artist")).To(ConsistOf("The Beatles")) }) + It("should get the first value by name", func() { + Expect(tags.First("genre")).To(Equal("Rock")) + Expect(tags.First("missing")).To(BeEmpty()) + }) Describe("Hash", func() { It("should always return the same value for the same tags ", func() { diff --git a/server/subsonic/helpers.go b/server/subsonic/helpers.go index 2d9d53b18..55b1b213e 100644 --- a/server/subsonic/helpers.go +++ b/server/subsonic/helpers.go @@ -515,9 +515,7 @@ func buildOSAlbumID3(ctx context.Context, album model.Album) *responses.OpenSubs dir.IsCompilation = album.Compilation dir.DiscTitles = buildDiscSubtitles(album) dir.ExplicitStatus = mapExplicitStatus(album.ExplicitStatus) - if len(album.Tags.Values(model.TagAlbumVersion)) > 0 { - dir.Version = album.Tags.Values(model.TagAlbumVersion)[0] - } + dir.Version = album.Tags.First(model.TagAlbumVersion) return &dir } diff --git a/utils/slice/slice.go b/utils/slice/slice.go index 73537c8f8..e3677e754 100644 --- a/utils/slice/slice.go +++ b/utils/slice/slice.go @@ -33,6 +33,23 @@ func Group[T any, K comparable](s []T, keyFunc func(T) K) map[K][]T { return m } +// GroupOrdered is Group with the groups in first-seen order. +func GroupOrdered[T any, K comparable](s []T, keyFunc func(T) K) [][]T { + var groups [][]T + index := map[K]int{} + for _, item := range s { + k := keyFunc(item) + i, ok := index[k] + if !ok { + i = len(groups) + index[k] = i + groups = append(groups, nil) + } + groups[i] = append(groups[i], item) + } + return groups +} + func ToMap[T any, K comparable, V any](s []T, transformFunc func(T) (K, V)) map[K]V { m := make(map[K]V, len(s)) for _, item := range s { diff --git a/utils/slice/slice_test.go b/utils/slice/slice_test.go index 27548d693..506b61b9f 100644 --- a/utils/slice/slice_test.go +++ b/utils/slice/slice_test.go @@ -63,6 +63,18 @@ var _ = Describe("Slice Utils", func() { }) }) + Describe("GroupOrdered", func() { + It("returns nil for an empty input", func() { + Expect(slice.GroupOrdered([]int{}, func(v int) int { return v })).To(BeNil()) + }) + + It("keeps groups in first-seen order and items in input order", func() { + keyFunc := func(v int) int { return v % 3 } + result := slice.GroupOrdered([]int{2, 1, 4, 3, 5, 6, 8}, keyFunc) + Expect(result).To(Equal([][]int{{2, 5, 8}, {1, 4}, {3, 6}})) + }) + }) + Describe("ToMap", func() { It("returns empty map for an empty input", func() { transformFunc := func(v int) (int, string) { return v, strconv.Itoa(v) } From 83fff44c822d47d1275b9dfb0142bdfefac90cd4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Fri, 25 Sep 2026 16:42:49 -0400 Subject: [PATCH 12/26] feat(archiver): add folder cover image to downloaded zips (#6224) * feat(archiver): add folder cover image to downloaded zips Album, artist, playlist and share downloads now include the item's cover as folder., the image most players and car stereos show for the files next to it. This helps users who copy transcoded downloads to offline devices, since transcoding drops the embedded artwork (#5841). Album folders get the album cover, the artist zip root gets the artist image, and playlist and share zips get the playlist or shared item's cover at the root. The image is the same one getCoverArt serves, resized to 500px (not square), and named by its detected type. Items without artwork get no image, and a cover that fails to load is logged and skipped so it never breaks the archive. The archiver reads covers through a new core.CoverArtReader interface, implemented by artwork.CoverArtReader, to avoid an import cycle between core and core/artwork. * refactor(archiver): read covers through artwork.Artwork directly The archiver no longer needs a local CoverArtReader interface and adapter. The only reason core/artwork imported core was a core.AbsolutePath call in loadArtistFolder, which now reads the library path from the repository and cleans it the same way. With the import cycle gone, the archiver takes artwork.Artwork and treats ErrUnavailable and ErrNotFound as no cover. Covers are now written after each album's tracks (and after all tracks for the archive root), so a slow artwork lookup does not delay the first bytes of the download. * fix(archiver): read share covers as admin so private playlists keep theirs Public share downloads run with an anonymous context, and the playlist repository hides private playlists from anonymous users, so a zip of a shared private playlist silently had no folder image. Like the public image handler, the share itself is the authorization: the cover lookup now runs with an admin user. Only the cover read is elevated; streaming keeps the anonymous context, so the transcode limiter still keys public downloads the same way. * style(archiver): trim comments Shorten the comments added by the folder cover image change and drop the ones the names already explain. * fix(archiver): add one cover per album folder in artist zips Albums with the same name share a zip folder (pre-existing naming), so an artist zip with two such albums wrote two folder. entries at the same path. Keep the first cover and skip the rest for that folder. --- cmd/wire_gen.go | 4 +- core/archiver.go | 96 ++++++++++-- core/archiver_test.go | 178 +++++++++++++++++++++- core/artwork/folders_artist.go | 6 +- core/artwork/folders_artist_paths_test.go | 15 +- 5 files changed, 270 insertions(+), 29 deletions(-) diff --git a/cmd/wire_gen.go b/cmd/wire_gen.go index 43ba808d0..cb06cc047 100644 --- a/cmd/wire_gen.go +++ b/cmd/wire_gen.go @@ -95,7 +95,7 @@ func CreateSubsonicAPIRouter(ctx context.Context) *subsonic.Router { transcodingCache := stream.GetTranscodingCache() mediaStreamer := stream.NewMediaStreamer(dataStore, fFmpeg, transcodingCache) share := core.NewShare(dataStore) - archiver := core.NewArchiver(mediaStreamer, dataStore, share) + archiver := core.NewArchiver(mediaStreamer, dataStore, share, artworkArtwork) players := core.NewPlayers(dataStore) broker := events.GetBroker() metricsMetrics := metrics.GetPrometheusInstance(dataStore) @@ -152,7 +152,7 @@ func CreatePublicRouter() *public.Router { transcodingCache := stream.GetTranscodingCache() mediaStreamer := stream.NewMediaStreamer(dataStore, fFmpeg, transcodingCache) share := core.NewShare(dataStore) - archiver := core.NewArchiver(mediaStreamer, dataStore, share) + archiver := core.NewArchiver(mediaStreamer, dataStore, share, artworkArtwork) router := public.New(dataStore, artworkArtwork, mediaStreamer, share, archiver) return router } diff --git a/core/archiver.go b/core/archiver.go index 6406ca075..33236d889 100644 --- a/core/archiver.go +++ b/core/archiver.go @@ -7,20 +7,27 @@ import ( "errors" "fmt" "io" + "net/http" "os" + "path" "path/filepath" "strconv" "strings" + "time" "github.com/Masterminds/squirrel" + "github.com/navidrome/navidrome/core/artwork" "github.com/navidrome/navidrome/core/stream" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/persistence" "github.com/navidrome/navidrome/utils/slice" "github.com/navidrome/navidrome/utils/str" ) +const archiveCoverArtSize = 500 + type Archiver interface { ZipAlbum(ctx context.Context, id string, format string, bitrate int, w io.Writer) error ZipArtist(ctx context.Context, id string, format string, bitrate int, w io.Writer) error @@ -28,18 +35,19 @@ type Archiver interface { ZipPlaylist(ctx context.Context, id string, format string, bitrate int, w io.Writer) error } -func NewArchiver(ms stream.MediaStreamer, ds model.DataStore, shares Share) Archiver { - return &archiver{ds: ds, ms: ms, shares: shares} +func NewArchiver(ms stream.MediaStreamer, ds model.DataStore, shares Share, artwork artwork.Artwork) Archiver { + return &archiver{ds: ds, ms: ms, shares: shares, artwork: artwork} } type archiver struct { - ds model.DataStore - ms stream.MediaStreamer - shares Share + ds model.DataStore + ms stream.MediaStreamer + shares Share + artwork artwork.Artwork } func (a *archiver) ZipAlbum(ctx context.Context, id string, format string, bitrate int, out io.Writer) error { - return a.zipAlbums(ctx, id, format, bitrate, out, squirrel.Eq{"album_id": id}) + return a.zipAlbums(ctx, id, format, bitrate, out, squirrel.Eq{"album_id": id}, model.ArtworkID{}) } func (a *archiver) ZipArtist(ctx context.Context, id string, format string, bitrate int, out io.Writer) error { @@ -49,10 +57,11 @@ func (a *archiver) ZipArtist(ctx context.Context, id string, format string, bitr persistence.ParticipantIDFilter("media_file", id, model.RoleAlbumArtist), squirrel.Eq{"missing": false}, } - return a.zipAlbums(ctx, id, format, bitrate, out, filter) + return a.zipAlbums(ctx, id, format, bitrate, out, filter, model.Artist{ID: id}.CoverArtID()) } -func (a *archiver) zipAlbums(ctx context.Context, id string, format string, bitrate int, out io.Writer, filters squirrel.Sqlizer) error { +// 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"}) if err != nil { log.Error(ctx, "Error loading mediafiles from artist", "id", id, err) @@ -80,7 +89,10 @@ func (a *archiver) zipAlbums(ctx context.Context, id string, format string, bitr return addErr } } + // After the tracks, so a slow artwork lookup doesn't delay the first bytes. + a.addCoverArtToZip(ctx, z, album[0].AlbumCoverArtID(), folder) } + a.addCoverArtToZip(ctx, z, rootArt, "") err = z.Close() if err != nil { log.Error(ctx, "Error closing zip file", "id", id, err) @@ -170,7 +182,10 @@ func (a *archiver) ZipShare(ctx context.Context, s *model.Share, out io.Writer) return model.ErrNotAuthorized } log.Debug(ctx, "Zipping share", "name", s.ID, "format", s.Format, "bitrate", s.MaxBitRate, "numTracks", len(s.Tracks)) - return a.zipMediaFiles(ctx, s.ID, s.ID, s.Format, s.MaxBitRate, out, s.Tracks, false) + // The share is the authorization (as in the public image handler): an anonymous lookup would + // hide a private playlist. Only the cover read is elevated. + coverCtx := request.WithUser(ctx, model.User{IsAdmin: true}) + return a.zipMediaFiles(ctx, s.ID, s.ID, s.Format, s.MaxBitRate, out, s.Tracks, coverCtx, s.CoverArtID(), false) } func (a *archiver) ZipPlaylist(ctx context.Context, id string, format string, bitrate int, out io.Writer) error { @@ -181,10 +196,10 @@ func (a *archiver) ZipPlaylist(ctx context.Context, id string, format string, bi } mfs := pls.MediaFiles() log.Debug(ctx, "Zipping playlist", "name", pls.Name, "format", format, "bitrate", bitrate, "numTracks", len(mfs)) - return a.zipMediaFiles(ctx, id, pls.Name, format, bitrate, out, mfs, true) + return a.zipMediaFiles(ctx, id, pls.Name, format, bitrate, out, mfs, ctx, pls.CoverArtID(), true) } -func (a *archiver) zipMediaFiles(ctx context.Context, id, name string, format string, bitrate int, out io.Writer, mfs model.MediaFiles, addM3U bool) error { +func (a *archiver) zipMediaFiles(ctx context.Context, id, name string, format string, bitrate int, out io.Writer, mfs model.MediaFiles, coverCtx context.Context, coverArt model.ArtworkID, addM3U bool) error { z := createZipWriter(out, format, bitrate) zippedMfs := make(model.MediaFiles, len(mfs)) @@ -199,6 +214,7 @@ func (a *archiver) zipMediaFiles(ctx context.Context, id, name string, format st mf.Path = file zippedMfs[idx] = mf } + a.addCoverArtToZip(coverCtx, z, coverArt, "") // Add M3U file if requested if addM3U && len(zippedMfs) > 0 { @@ -276,3 +292,61 @@ func (a *archiver) addFileToZip(ctx context.Context, z *zip.Writer, mf model.Med return nil } + +// addCoverArtToZip adds the cover as dir/folder.. Errors are logged, never returned. +func (a *archiver) addCoverArtToZip(ctx context.Context, z *zip.Writer, artID model.ArtworkID, dir string) { + if artID.ID == "" { + return + } + // Buffered so a failed read leaves no empty entry. + data, err := a.readCoverArt(ctx, artID) + if errors.Is(err, artwork.ErrUnavailable) || errors.Is(err, model.ErrNotFound) { + log.Debug(ctx, "No cover art to add to zip", "artID", artID) + return + } + if err != nil { + log.Warn(ctx, "Error reading cover art for zipping", "artID", artID, err) + return + } + ext := coverArtExtension(data) + if ext == "" { + log.Warn(ctx, "Unknown cover art image type, not adding it to zip", "artID", artID) + return + } + w, err := z.CreateHeader(&zip.FileHeader{ + Name: path.Join(dir, "folder."+ext), + Modified: time.Now(), + Method: zip.Store, + }) + if err != nil { + log.Warn(ctx, "Error creating cover art zip entry", "artID", artID, err) + return + } + if _, err = w.Write(data); err != nil { + log.Warn(ctx, "Error zipping cover art", "artID", artID, err) + } +} + +func (a *archiver) readCoverArt(ctx context.Context, artID model.ArtworkID) ([]byte, error) { + img, err := a.artwork.Get(ctx, artID, archiveCoverArtSize, false) + if err != nil { + return nil, err + } + defer img.Close() + return io.ReadAll(img) +} + +// Resizing may re-encode the image, so the type comes from its bytes. +func coverArtExtension(data []byte) string { + switch http.DetectContentType(data) { + case "image/jpeg": + return "jpg" + case "image/png": + return "png" + case "image/webp": + return "webp" + case "image/gif": + return "gif" + } + return "" +} diff --git a/core/archiver_test.go b/core/archiver_test.go index 86833717a..9dab44cef 100644 --- a/core/archiver_test.go +++ b/core/archiver_test.go @@ -4,6 +4,7 @@ import ( "archive/zip" "bytes" "context" + "errors" "io" "strings" @@ -11,8 +12,10 @@ import ( "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/core" + "github.com/navidrome/navidrome/core/artwork" "github.com/navidrome/navidrome/core/stream" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/persistence" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -25,13 +28,15 @@ var _ = Describe("Archiver", func() { ms *mockMediaStreamer ds *mockDataStore sh *mockShare + ca *mockCoverArt ) BeforeEach(func() { ms = &mockMediaStreamer{} sh = &mockShare{} ds = &mockDataStore{} - arch = core.NewArchiver(ms, ds, sh) + ca = &mockCoverArt{images: map[string][]byte{}} + arch = core.NewArchiver(ms, ds, sh, ca) }) Context("ZipAlbum", func() { @@ -332,8 +337,179 @@ var _ = Describe("Archiver", func() { Expect(string(m3uContent)).To(Equal(expectedM3U)) }) }) + Context("cover art", func() { + var ( + jpegData = []byte("\xff\xd8\xff\xe0 fake jpeg") + pngData = []byte("\x89PNG\x0d\x0a\x1a\x0a fake png") + ) + + mockAlbumTracks := func(filter squirrel.Sqlizer, mfs model.MediaFiles) { + mfRepo := &mockMediaFileRepository{} + mfRepo.On("GetAll", []model.QueryOptions{{Filters: filter, Sort: "album"}}).Return(mfs, nil) + ds.On("MediaFile", mock.Anything).Return(mfRepo) + ms.On("NewStream", mock.Anything, mock.Anything, mock.Anything).Return(io.NopCloser(strings.NewReader("test")), nil) + } + + It("adds the album cover to the album folder", func() { + ca.images["al-1"] = jpegData + mockAlbumTracks(squirrel.Eq{"album_id": "1"}, model.MediaFiles{ + {Path: "test_data/01 - track1.mp3", Suffix: "mp3", AlbumID: "1", Album: "Album/Promo", DiscNumber: 1}, + }) + + out := new(bytes.Buffer) + Expect(arch.ZipAlbum(context.Background(), "1", "mp3", 128, out)).To(Succeed()) + + files := readZip(out) + Expect(files).To(HaveLen(2)) + Expect(files).To(HaveKeyWithValue("Album_Promo/folder.jpg", jpegData)) + Expect(ca.requests).To(ConsistOf(coverRequest{id: "al-1", size: 500, square: false})) + }) + + It("adds the artist image to the root and each album cover to its folder", func() { + ca.images["ar-1"] = pngData + ca.images["al-1"] = jpegData + ca.images["al-2"] = jpegData + mockAlbumTracks(squirrel.And{ + persistence.ParticipantIDFilter("media_file", "1", model.RoleAlbumArtist), + squirrel.Eq{"missing": false}, + }, model.MediaFiles{ + {Path: "test_data/01 - track1.mp3", Suffix: "mp3", AlbumID: "1", Album: "Album 1", DiscNumber: 1}, + {Path: "test_data/02 - track2.mp3", Suffix: "mp3", AlbumID: "2", Album: "Album 2", DiscNumber: 1}, + }) + + out := new(bytes.Buffer) + Expect(arch.ZipArtist(context.Background(), "1", "mp3", 128, out)).To(Succeed()) + + files := readZip(out) + Expect(files).To(HaveLen(5)) + Expect(files).To(HaveKeyWithValue("folder.png", pngData)) + Expect(files).To(HaveKeyWithValue("Album 1/folder.jpg", jpegData)) + Expect(files).To(HaveKeyWithValue("Album 2/folder.jpg", jpegData)) + }) + + It("puts each same-named album's cover in that album's own folder", func() { + ca.images["al-1"] = jpegData + ca.images["al-2"] = pngData + mockAlbumTracks(squirrel.And{ + persistence.ParticipantIDFilter("media_file", "1", model.RoleAlbumArtist), + squirrel.Eq{"missing": false}, + }, model.MediaFiles{ + {Path: "test_data/01 - track1.mp3", Suffix: "mp3", AlbumID: "1", Album: "Greatest Hits", Year: 2001, DiscNumber: 1}, + {Path: "test_data/02 - track2.mp3", Suffix: "mp3", AlbumID: "2", Album: "Greatest Hits", Year: 2005, DiscNumber: 1}, + }) + + out := new(bytes.Buffer) + Expect(arch.ZipArtist(context.Background(), "1", "mp3", 128, out)).To(Succeed()) + + files := readZip(out) + Expect(files).To(HaveKeyWithValue("Greatest Hits [2001]/folder.jpg", jpegData)) + Expect(files).To(HaveKeyWithValue("Greatest Hits [2005]/folder.png", pngData)) + }) + + It("adds the playlist cover to the root", func() { + ca.images["pl-1"] = jpegData + plRepo := &mockPlaylistRepository{} + plRepo.On("GetWithTracks", "1", true, false).Return(&model.Playlist{ + ID: "1", + Name: "Test Playlist", + Tracks: []model.PlaylistTrack{ + {MediaFile: model.MediaFile{Path: "test_data/01 - track1.mp3", Suffix: "mp3", AlbumID: "1", Artist: "Artist 1", Title: "track1"}}, + }, + }, nil) + ds.On("Playlist", mock.Anything).Return(plRepo) + ms.On("NewStream", mock.Anything, mock.Anything, mock.Anything).Return(io.NopCloser(strings.NewReader("test")), nil) + + out := new(bytes.Buffer) + Expect(arch.ZipPlaylist(context.Background(), "1", "mp3", 128, out)).To(Succeed()) + + files := readZip(out) + Expect(files).To(HaveLen(3)) + Expect(files).To(HaveKeyWithValue("folder.jpg", jpegData)) + Expect(files).To(HaveKey("Test Playlist.m3u")) + }) + + It("adds the shared item's cover to the root, even for a private playlist", func() { + ca.images["pl-10"] = jpegData + ms.On("NewStream", mock.Anything, mock.Anything, mock.Anything).Return(io.NopCloser(strings.NewReader("test")), nil) + share := &model.Share{ + ID: "1", + Downloadable: true, + Format: "mp3", + MaxBitRate: 128, + ResourceType: "playlist", + ResourceIDs: "10", + Tracks: model.MediaFiles{ + {ID: "1", Path: "test_data/01 - track1.mp3", Suffix: "mp3", Artist: "Artist 1", Title: "track1"}, + }, + } + + out := new(bytes.Buffer) + Expect(arch.ZipShare(context.Background(), share, out)).To(Succeed()) + + files := readZip(out) + Expect(files).To(HaveLen(2)) + Expect(files).To(HaveKeyWithValue("folder.jpg", jpegData)) + Expect(ca.requests).To(ConsistOf(coverRequest{id: "pl-10", size: 500, square: false, admin: true})) + }) + + It("still builds the archive when the cover cannot be read", func() { + ca.err = errors.New("boom") + mockAlbumTracks(squirrel.Eq{"album_id": "1"}, model.MediaFiles{ + {Path: "test_data/01 - track1.mp3", Suffix: "mp3", AlbumID: "1", Album: "Album", DiscNumber: 1}, + }) + + out := new(bytes.Buffer) + Expect(arch.ZipAlbum(context.Background(), "1", "mp3", 128, out)).To(Succeed()) + + files := readZip(out) + Expect(files).To(HaveLen(1)) + Expect(files).To(HaveKey("Album/01 - track1.mp3")) + }) + }) }) +func readZip(out *bytes.Buffer) map[string][]byte { + zr, err := zip.NewReader(bytes.NewReader(out.Bytes()), int64(out.Len())) + Expect(err).ToNot(HaveOccurred()) + files := make(map[string][]byte, len(zr.File)) + for _, f := range zr.File { + r, err := f.Open() + Expect(err).ToNot(HaveOccurred()) + data, err := io.ReadAll(r) + Expect(err).ToNot(HaveOccurred()) + _ = r.Close() + files[f.Name] = data + } + return files +} + +type coverRequest struct { + id string + size int + square bool + admin bool +} + +type mockCoverArt struct { + artwork.Artwork + images map[string][]byte + err error + requests []coverRequest +} + +func (m *mockCoverArt) Get(ctx context.Context, artID model.ArtworkID, size int, square bool) (*artwork.Image, error) { + user, _ := request.UserFrom(ctx) + m.requests = append(m.requests, coverRequest{id: artID.String(), size: size, square: square, admin: user.IsAdmin}) + if m.err != nil { + return nil, m.err + } + data, ok := m.images[artID.String()] + if !ok { + return nil, artwork.ErrUnavailable + } + return &artwork.Image{ReadCloser: io.NopCloser(bytes.NewReader(data))}, nil +} + type mockDataStore struct { mock.Mock model.DataStore diff --git a/core/artwork/folders_artist.go b/core/artwork/folders_artist.go index 1ca1ce034..efef42f81 100644 --- a/core/artwork/folders_artist.go +++ b/core/artwork/folders_artist.go @@ -14,7 +14,6 @@ import ( "time" "github.com/Masterminds/squirrel" - "github.com/navidrome/navidrome/core" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils" @@ -169,8 +168,9 @@ func loadArtistFolder(ctx context.Context, ds model.DataStore, albums model.Albu folderPath = filepath.Dir(folderPath) } - // TODO: Hacky, but the easiest way to get the folder ID ATM - libPath := core.AbsolutePath(ctx, ds, libID, "") + // Cleaned like the album paths; Join keeps an empty path empty, Clean would return ".". + libPath, _ := ds.Library(ctx).GetPath(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, diff --git a/core/artwork/folders_artist_paths_test.go b/core/artwork/folders_artist_paths_test.go index cc8af3e63..127661891 100644 --- a/core/artwork/folders_artist_paths_test.go +++ b/core/artwork/folders_artist_paths_test.go @@ -6,7 +6,6 @@ import ( "path/filepath" "time" - "github.com/navidrome/navidrome/core" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" @@ -47,11 +46,11 @@ var _ = Describe("loadArtistFolder", func() { BeforeEach(func() { ctx = context.Background() - DeferCleanup(stubCoreAbsolutePath()) - updatedAt = time.Now().Truncate(time.Second).Add(5 * time.Minute) repo = &fakeFolderRepo{result: []model.Folder{{ImagesUpdatedAt: updatedAt}}} - ds = &tests.MockDataStore{MockedFolder: repo} + libRepo := &tests.MockLibraryRepo{} + libRepo.SetData(model.Libraries{{ID: 1, Path: filepath.FromSlash("/music")}}) + ds = &tests.MockDataStore{MockedFolder: repo, MockedLibrary: libRepo} albums = model.Albums{{LibraryID: 1, ID: "album1", Name: "Album 1"}} }) @@ -107,11 +106,3 @@ var _ = Describe("loadArtistFolder", func() { Expect(upd).To(BeZero()) }) }) - -func stubCoreAbsolutePath() func() { - original := core.AbsolutePath - core.AbsolutePath = func(context.Context, model.DataStore, int, string) string { - return filepath.FromSlash("/music") - } - return func() { core.AbsolutePath = original } -} From b293b962566019d33c3e094164649b75c777b657 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Fri, 25 Sep 2026 18:06:10 -0400 Subject: [PATCH 13/26] refactor(persistence): stateless repositories with per-call context (#6149) * refactor(persistence): adopt generic deluan/rest repository API Pin deluan/rest to the refactor branch. REST-facing repository methods take a context and return typed values. Drop DataStore.Resource and ResourceRepository; the native API names typed repositories directly through a per-request adapter that later commits remove. * refactor(persistence): base repository helpers take a context * refactor(persistence): LibraryRepository takes a context per call * refactor(persistence): PropertyRepository takes a context per call * refactor(persistence): UserPropsRepository takes a context per call * refactor(persistence): TranscodingRepository takes a context per call * refactor(persistence): ShareRepository takes a context per call * refactor(persistence): PlayerRepository takes a context per call * refactor(persistence): RadioRepository takes a context per call * refactor(persistence): PlayQueueRepository takes a context per call * refactor(persistence): Tag and Genre repositories take a context per call * refactor(persistence): PluginRepository takes a context per call * refactor(persistence): Scrobble repositories take a context per call * refactor(persistence): FolderRepository takes a context per call * refactor(persistence): Artwork repositories take a context per call * refactor(persistence): UserRepository takes a context per call * refactor(persistence): ArtistRepository takes a context per call ReadAll no longer rewrites the shared sort mappings for the role filter; it works on a per-call copy. * test(persistence): assert artist role sort sanitization in ReadAll * refactor(persistence): AlbumRepository takes a context per call * test(persistence): pass the test context to album repository helpers * refactor(persistence): MediaFileRepository takes a context per call * refactor(persistence): Playlist repositories take a context per call * refactor(persistence): build all repositories once per store * refactor(core): REST repository wrappers are built once * refactor(persistence): repositories are stateless Remove the context field from the base repository and the per-request REST adapter. Enable the containedctx linter so no repository can hold a request context again. * chore(lint): skip containedctx in test files * refactor: share simplifications from the stateless repositories sweep Add deleteOwnedAll on sqlRepository and use it in player/share Delete to remove the duplicated bulk-delete loop; have Share.Repository() return model.ShareRepository so subsonic sharing.go drops its repeated type assertions. * chore(core): assert REST wrappers implement Persistable * chore: reformat imports * perf(persistence): build repositories on first use Each transaction store used to construct all 21 repositories up front, paying for filter and sort mapping setup the block never touched. Fields are now sync.OnceValue thunks, so a store only builds what it uses. * fix(persistence): clean plugin references per deleted user A bulk user delete that fails on a later id had already removed the earlier rows but skipped their plugin cleanup. Cleanup now runs right after each successful delete. * fix(core): unload disabled plugins even when a user delete fails A bulk delete can fail on a later id after earlier users were removed and their plugins auto-disabled. The wrapper returned before unloading, leaving those plugins running until the next successful delete or a restart. * chore(deps): pin deluan/rest to v1.0.1 Replaces the pseudo-version of the refactor branch with the tagged release. REST error messages now name the bare type (Artist, not model.Artist). * test: use the spec context instead of context.Background() Replace the context.Background()/context.TODO() calls this branch added to tests with the spec's ctx, GinkgoT().Context(), or t/b.Context(), so repository calls are bound to the running spec's lifetime. * test: declare the spec context once per Describe Set ctx from GinkgoT().Context() first in each top-level BeforeEach and reuse it, building user contexts on top of it instead of repeating inline calls. --- .golangci.yml | 4 + adapters/lastfm/agent_test.go | 2 +- adapters/lastfm/auth_router_test.go | 2 +- adapters/listenbrainz/agent_test.go | 2 +- cmd/artwork.go | 34 +- cmd/artwork_test.go | 90 ++-- cmd/missing.go | 6 +- cmd/pls.go | 10 +- cmd/plugin.go | 8 +- cmd/root.go | 10 +- cmd/svc.go | 2 +- cmd/user.go | 26 +- cmd/utils.go | 4 +- core/agents/local_agent.go | 6 +- core/agents/session_keys.go | 6 +- core/archiver.go | 4 +- core/archiver_test.go | 24 +- core/artwork/artwork.go | 24 +- core/artwork/artwork_suite_test.go | 7 +- core/artwork/artwork_test.go | 44 +- core/artwork/disc.go | 6 +- core/artwork/e2e/acquire_serve_test.go | 36 +- core/artwork/e2e/artist_test.go | 4 +- core/artwork/e2e/e2e_suite_test.go | 10 +- core/artwork/e2e/mediafile_test.go | 2 +- core/artwork/e2e/playlist_test.go | 10 +- core/artwork/e2e/radio_test.go | 4 +- core/artwork/e2e/resolution_harness_test.go | 18 +- core/artwork/folders_album.go | 6 +- core/artwork/folders_artist.go | 4 +- core/artwork/housekeeping.go | 24 +- core/artwork/housekeeping_test.go | 8 +- core/artwork/library_fs.go | 2 +- core/artwork/library_fs_test.go | 4 +- core/artwork/processor.go | 10 +- core/artwork/processor_test.go | 54 +-- core/artwork/prune.go | 10 +- core/artwork/prune_test.go | 48 +- core/artwork/resolve.go | 16 +- core/artwork/uploader_test.go | 25 +- core/artwork/worker.go | 18 +- core/artwork/worker_soak_test.go | 10 +- core/artwork/worker_test.go | 119 ++--- core/auth/auth.go | 8 +- core/common.go | 2 +- core/external/extdata_helper_test.go | 18 +- core/external/provider.go | 16 +- core/external/provider_refreshinfo_test.go | 8 +- core/external/provider_similarsongs.go | 8 +- .../external/provider_updatealbuminfo_test.go | 2 +- .../provider_updateartistinfo_test.go | 2 +- core/library.go | 140 +++--- core/library_test.go | 110 ++--- core/lyrics/lyrics.go | 2 +- core/maintenance.go | 36 +- core/maintenance_test.go | 20 +- core/matcher/matcher.go | 10 +- core/matcher/matcher_test.go | 8 +- core/metrics/insights.go | 28 +- core/metrics/prometheus.go | 8 +- core/playback/device.go | 2 +- core/playback/playbackserver.go | 2 +- core/players.go | 10 +- core/players_test.go | 6 +- core/playlists/import.go | 14 +- core/playlists/import_test.go | 6 +- core/playlists/parse_m3u.go | 6 +- core/playlists/parse_m3u_test.go | 2 +- core/playlists/playlists.go | 61 +-- core/playlists/playlists_test.go | 6 +- core/playlists/rest_adapter.go | 45 +- core/playlists/rest_adapter_test.go | 120 ++--- core/scrobbler/buffered_scrobbler.go | 16 +- core/scrobbler/buffered_scrobbler_test.go | 24 +- core/scrobbler/play_tracker.go | 24 +- core/scrobbler/play_tracker_test.go | 42 +- core/share.go | 91 ++-- core/share_test.go | 41 +- core/sonic/sonic.go | 6 +- core/stream/decider.go | 6 +- core/stream/media_streamer.go | 2 +- core/stream/media_streamer_test.go | 6 +- core/user.go | 56 +-- core/user_test.go | 42 +- db/db.go | 2 +- go.mod | 2 +- go.sum | 18 +- model/album.go | 34 +- model/annotation.go | 13 +- model/artist.go | 23 +- model/artwork.go | 55 ++- model/bookmark.go | 11 +- model/datastore.go | 49 +- model/folder.go | 21 +- model/genre.go | 11 +- model/get_entity.go | 10 +- model/get_entity_test.go | 6 +- model/library.go | 29 +- model/mediafile.go | 51 +- model/player.go | 15 +- model/playlist.go | 54 ++- model/playqueue.go | 9 +- model/plugin.go | 21 +- model/properties.go | 10 +- model/radio.go | 16 +- model/scrobble.go | 16 +- model/scrobble_buffer.go | 17 +- model/searchable.go | 4 +- model/share.go | 12 +- model/tag.go | 9 +- model/transcoding.go | 16 +- model/user.go | 29 +- model/user_props.go | 10 +- persistence/album_repository.go | 137 +++--- persistence/album_repository_test.go | 338 +++++++------ persistence/artist_repository.go | 177 +++---- persistence/artist_repository_test.go | 434 ++++++++++------- persistence/artwork_hydration.go | 2 +- persistence/artwork_hydration_test.go | 120 ++--- persistence/artwork_queue_repository.go | 71 ++- persistence/artwork_queue_repository_test.go | 228 ++++----- persistence/artwork_repository.go | 48 +- persistence/artwork_repository_test.go | 140 +++--- persistence/criteria_sql_benchmark_test.go | 6 +- persistence/e2e/e2e_suite_test.go | 36 +- persistence/e2e/smartplaylist_test.go | 4 +- persistence/folder_repository.go | 87 ++-- persistence/folder_repository_test.go | 98 ++-- persistence/genre_repository.go | 36 +- persistence/genre_repository_test.go | 79 ++- persistence/item_tags_test.go | 35 +- persistence/library_repository.go | 136 +++--- persistence/library_repository_test.go | 81 ++-- persistence/mediafile_repository.go | 211 ++++---- persistence/mediafile_repository_test.go | 450 +++++++++--------- persistence/persistence.go | 207 ++++---- persistence/persistence_suite_test.go | 62 +-- persistence/persistence_test.go | 36 +- persistence/player_repository.go | 88 ++-- persistence/player_repository_test.go | 109 ++--- persistence/playlist_repository.go | 214 ++++----- persistence/playlist_repository_test.go | 176 +++---- persistence/playlist_track_repository.go | 140 +++--- persistence/playlist_track_repository_test.go | 153 +++--- persistence/playqueue_repository.go | 51 +- persistence/playqueue_repository_test.go | 124 ++--- persistence/plugin_cleanup_test.go | 81 ++-- persistence/plugin_repository.go | 67 ++- persistence/plugin_repository_test.go | 99 ++-- persistence/property_repository.go | 21 +- persistence/property_repository_test.go | 14 +- persistence/radio_repository.go | 97 ++-- persistence/radio_repository_test.go | 74 +-- persistence/scrobble_buffer_repository.go | 29 +- .../scrobble_buffer_repository_test.go | 54 +-- persistence/scrobble_repository.go | 53 +-- persistence/scrobble_repository_test.go | 37 +- persistence/share_repository.go | 116 +++-- persistence/share_repository_test.go | 260 +++++----- persistence/smart_playlist_repository.go | 73 +-- persistence/smart_playlist_repository_test.go | 191 ++++---- persistence/sort_index_coverage_test.go | 6 +- persistence/sql_annotations.go | 55 +-- persistence/sql_annotations_test.go | 75 +-- persistence/sql_base_repository.go | 151 +++--- persistence/sql_base_repository_test.go | 61 +-- persistence/sql_bookmarks.go | 77 +-- persistence/sql_bookmarks_test.go | 58 +-- persistence/sql_participations.go | 13 +- persistence/sql_restful.go | 14 +- persistence/sql_search.go | 21 +- persistence/sql_search_fts.go | 5 +- persistence/sql_search_fts_test.go | 42 +- persistence/sql_search_like.go | 5 +- persistence/sql_search_like_test.go | 12 +- persistence/sql_tags.go | 49 +- persistence/tag_library_filtering_test.go | 16 +- persistence/tag_repository.go | 31 +- persistence/tag_repository_test.go | 65 +-- persistence/transcoding_repository.go | 81 ++-- persistence/transcoding_repository_test.go | 78 ++- persistence/user_props_repository.go | 21 +- persistence/user_repository.go | 188 ++++---- persistence/user_repository_test.go | 218 +++++---- plugins/host_library.go | 4 +- plugins/host_library_test.go | 22 +- plugins/host_matcher_test.go | 44 +- plugins/host_scrobbleretriever.go | 6 +- plugins/host_scrobbleretriever_test.go | 28 +- plugins/host_storage_test.go | 2 +- plugins/host_subsonicapi.go | 2 +- plugins/host_subsonicapi_test.go | 33 +- plugins/host_taskqueue.go | 2 +- plugins/host_users.go | 2 +- plugins/host_users_test.go | 26 +- plugins/host_websocket.go | 2 +- plugins/manager.go | 36 +- plugins/manager_loader.go | 10 +- plugins/manager_plugin.go | 2 +- plugins/manager_readonly_test.go | 22 +- plugins/manager_sync.go | 22 +- plugins/manager_sync_test.go | 28 +- plugins/manager_watcher.go | 14 +- plugins/manager_watcher_test.go | 20 +- scanner/controller.go | 16 +- scanner/controller_test.go | 6 +- scanner/image_changes.go | 6 +- scanner/phase_1_folders.go | 54 +-- scanner/phase_2_missing_tracks.go | 18 +- scanner/phase_2_missing_tracks_test.go | 132 ++--- scanner/phase_3_refresh_albums.go | 12 +- scanner/phase_3_refresh_albums_test.go | 4 +- scanner/phase_4_playlists.go | 16 +- scanner/phase_4_playlists_test.go | 22 +- scanner/scanner.go | 32 +- scanner/scanner_benchmark_test.go | 2 +- scanner/scanner_multilibrary_test.go | 112 ++--- scanner/scanner_selective_test.go | 42 +- scanner/scanner_test.go | 148 +++--- scanner/watcher.go | 4 +- server/auth.go | 31 +- server/auth_test.go | 28 +- server/events/sse.go | 2 +- server/initial_setup.go | 18 +- server/initial_setup_test.go | 24 +- server/jellyfin/annotations.go | 12 +- server/jellyfin/annotations_test.go | 44 +- server/jellyfin/api_test.go | 13 +- server/jellyfin/auth.go | 4 +- server/jellyfin/auth_test.go | 18 +- server/jellyfin/browsing.go | 6 +- server/jellyfin/browsing_test.go | 30 +- server/jellyfin/e2e/auth_test.go | 4 +- server/jellyfin/e2e/e2e_suite_test.go | 10 +- server/jellyfin/e2e/multiuser_test.go | 4 +- server/jellyfin/e2e/playlists_test.go | 18 +- server/jellyfin/e2e/sessions_test.go | 6 +- server/jellyfin/e2e/similar_test.go | 4 +- server/jellyfin/images.go | 8 +- server/jellyfin/images_test.go | 14 +- server/jellyfin/items.go | 52 +- server/jellyfin/items_test.go | 158 +++--- server/jellyfin/lyrics_test.go | 2 +- server/jellyfin/middlewares.go | 2 +- server/jellyfin/middlewares_test.go | 10 +- server/jellyfin/playlists.go | 16 +- server/jellyfin/playlists_test.go | 14 +- server/jellyfin/quickconnect.go | 4 +- server/jellyfin/quickconnect_test.go | 7 +- server/jellyfin/similar.go | 2 +- server/jellyfin/similar_test.go | 8 +- server/jellyfin/socket_test.go | 5 +- server/jellyfin/stream.go | 2 +- server/jellyfin/stream_test.go | 38 +- server/jellyfin/system.go | 4 +- server/jellyfin/system_test.go | 8 +- server/jellyfin/users.go | 4 +- server/jellyfin/users_test.go | 14 +- server/middlewares.go | 2 +- server/middlewares_test.go | 10 +- server/nativeapi/artists.go | 16 +- server/nativeapi/config_test.go | 9 +- server/nativeapi/inspect.go | 2 +- server/nativeapi/library_test.go | 15 +- server/nativeapi/metadata_test.go | 21 +- server/nativeapi/missing.go | 29 +- server/nativeapi/missing_test.go | 5 +- server/nativeapi/native_api.go | 57 +-- server/nativeapi/native_api_song_test.go | 12 +- server/nativeapi/playlists.go | 10 +- server/nativeapi/playlists_test.go | 25 +- server/nativeapi/plugin.go | 20 +- server/nativeapi/plugin_test.go | 42 +- server/nativeapi/queue.go | 12 +- server/nativeapi/queue_test.go | 2 +- server/nativeapi/radios.go | 22 +- server/nativeapi/translations.go | 20 +- .../user_password_token_refresh_test.go | 9 +- server/public/handle_streams.go | 6 +- server/public/handle_streams_test.go | 6 +- server/serve_index.go | 2 +- server/serve_index_test.go | 3 +- server/subsonic/album_lists.go | 14 +- server/subsonic/album_lists_test.go | 20 +- server/subsonic/bookmarks.go | 30 +- server/subsonic/bookmarks_test.go | 2 +- server/subsonic/browsing.go | 20 +- server/subsonic/browsing_test.go | 8 +- server/subsonic/e2e/e2e_suite_test.go | 6 +- .../subsonic/e2e/subsonic_album_lists_test.go | 4 +- server/subsonic/e2e/subsonic_artwork_test.go | 14 +- .../subsonic/e2e/subsonic_bookmarks_test.go | 4 +- server/subsonic/e2e/subsonic_browsing_test.go | 24 +- .../e2e/subsonic_media_annotation_test.go | 14 +- .../e2e/subsonic_media_retrieval_test.go | 4 +- .../e2e/subsonic_multilibrary_test.go | 16 +- .../subsonic/e2e/subsonic_playlists_test.go | 30 +- server/subsonic/e2e/subsonic_sharing_test.go | 14 +- .../e2e/subsonic_sonic_similarity_test.go | 4 +- server/subsonic/e2e/subsonic_stream_test.go | 2 +- .../subsonic/e2e/subsonic_transcode_test.go | 22 +- server/subsonic/library_scanning.go | 4 +- server/subsonic/library_scanning_test.go | 12 +- server/subsonic/media_annotation.go | 22 +- server/subsonic/media_annotation_test.go | 2 +- server/subsonic/media_retrieval.go | 4 +- server/subsonic/media_retrieval_test.go | 6 +- server/subsonic/middlewares.go | 4 +- server/subsonic/middlewares_test.go | 26 +- server/subsonic/radio.go | 8 +- server/subsonic/searching.go | 10 +- server/subsonic/searching_test.go | 6 +- server/subsonic/sharing.go | 22 +- server/subsonic/stream.go | 2 +- server/subsonic/transcode.go | 4 +- tests/harness/harness.go | 8 +- tests/mock_album_repo.go | 35 +- tests/mock_artist_repo.go | 37 +- tests/mock_artwork_queue_repo.go | 31 +- tests/mock_artwork_repo.go | 21 +- tests/mock_data_store.go | 123 ++--- tests/mock_genre_repo.go | 25 +- tests/mock_library_repo.go | 55 +-- tests/mock_library_service.go | 15 +- tests/mock_mediafile_repo.go | 81 ++-- tests/mock_playlist_repo.go | 37 +- tests/mock_playlist_track_repo.go | 28 +- tests/mock_playqueue_repo.go | 9 +- tests/mock_plugin_repo.go | 45 +- tests/mock_property_repo.go | 16 +- tests/mock_radio_repository.go | 25 +- tests/mock_scrobble_buffer_repo.go | 15 +- tests/mock_scrobble_repo.go | 24 +- tests/mock_share_repo.go | 16 +- tests/mock_tag_repo.go | 4 +- tests/mock_transcoding_repo.go | 10 +- tests/mock_user_props_repo.go | 16 +- tests/mock_user_repo.go | 48 +- tests/mock_user_service.go | 7 +- 339 files changed, 6220 insertions(+), 6173 deletions(-) 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} } From 6b05189190a924b66042e32da13792ed53abdcf4 Mon Sep 17 00:00:00 2001 From: Deluan Date: Fri, 25 Sep 2026 18:18:58 -0400 Subject: [PATCH 14/26] chore: update Go dependencies to latest versions --- go.mod | 18 +++++++++--------- go.sum | 36 ++++++++++++++++++------------------ 2 files changed, 27 insertions(+), 27 deletions(-) diff --git a/go.mod b/go.mod index 01fb0f73d..464038d7e 100644 --- a/go.mod +++ b/go.mod @@ -8,7 +8,7 @@ replace go.senan.xyz/taglib => github.com/deluan/go-taglib v0.0.0-20260913142955 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/bmatcuk/doublestar/v4 v4.10.2 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 @@ -16,7 +16,7 @@ require ( github.com/djherbis/fscache v0.10.2-0.20260829235704-6d85d5878c22 github.com/djherbis/stream v1.5.1 github.com/djherbis/times v1.6.0 - github.com/dustin/go-humanize v1.0.1 + github.com/dustin/go-humanize v1.1.0 github.com/extism/go-sdk v1.7.1 github.com/fatih/structs v1.1.0 github.com/gen2brain/webp v0.6.4 @@ -39,8 +39,8 @@ require ( github.com/mattn/go-sqlite3 v1.14.52 github.com/microcosm-cc/bluemonday v1.0.27 github.com/mileusna/useragent v1.3.5 - github.com/onsi/ginkgo/v2 v2.32.2 - github.com/onsi/gomega v1.43.0 + github.com/onsi/ginkgo/v2 v2.33.0 + github.com/onsi/gomega v1.44.0 github.com/pelletier/go-toml/v2 v2.4.3 github.com/pmezard/go-difflib v1.0.0 github.com/pocketbase/dbx v1.12.0 @@ -81,15 +81,15 @@ require ( github.com/creack/pty v1.1.24 // indirect github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.1 // indirect github.com/dylibso/observe-sdk/go v0.0.0-20240828172851-9145d8ad07e1 // indirect - github.com/ebitengine/purego v0.10.2 // indirect + github.com/ebitengine/purego v0.11.1 // indirect github.com/fsnotify/fsnotify v1.10.1 // indirect github.com/go-logr/logr v1.4.4 // indirect github.com/go-task/slim-sprig/v3 v3.0.0 // indirect - github.com/gobwas/glob v0.2.3 // indirect + github.com/gobwas/glob v1.0.0 // indirect github.com/goccy/go-json v0.10.6 // indirect github.com/goccy/go-yaml v1.19.2 // indirect github.com/google/go-cmp v0.7.0 // indirect - github.com/google/pprof v0.0.0-20260825171938-4d453200e7d9 // indirect + github.com/google/pprof v0.0.0-20260906184651-6331bc6350fe // indirect github.com/google/subcommands v1.2.0 // indirect github.com/gorilla/css v1.0.1 // indirect github.com/hashicorp/errwrap v1.1.0 // indirect @@ -133,8 +133,8 @@ require ( go.yaml.in/yaml/v3 v3.0.5 // indirect golang.org/x/crypto v0.57.0 // indirect golang.org/x/mod v0.41.0 // indirect - golang.org/x/telemetry v0.0.0-20260811182544-a038080d80e5 // indirect - golang.org/x/tools v0.49.0 // indirect + golang.org/x/telemetry v0.0.0-20260908163034-4bcc4b2ee518 // indirect + golang.org/x/tools v0.50.0 // indirect google.golang.org/protobuf v1.36.12 // indirect gopkg.in/ini.v1 v1.67.3 // indirect gopkg.in/natefinch/npipe.v2 v2.0.0-20160621034901-c1b8fa8bdcce // indirect diff --git a/go.sum b/go.sum index bcc1395ba..929366a76 100644 --- a/go.sum +++ b/go.sum @@ -14,8 +14,8 @@ github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuP github.com/aymerick/douceur v0.2.0/go.mod h1:wlT5vV2O3h55X9m7iVYN0TBM0NH/MmbLnd30/FjWUq4= github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= -github.com/bmatcuk/doublestar/v4 v4.10.0 h1:zU9WiOla1YA122oLM6i4EXvGW62DvKZVxIe6TYWexEs= -github.com/bmatcuk/doublestar/v4 v4.10.0/go.mod h1:xBQ8jztBU6kakFMg+8WGxn0c6z1fTSPVIjEY1Wr7jzc= +github.com/bmatcuk/doublestar/v4 v4.10.2 h1:eF7W7HWKg3z9NrWV9pTLnNeoXaqq3Tq9DNKXVMfoCnw= +github.com/bmatcuk/doublestar/v4 v4.10.2/go.mod h1:xBQ8jztBU6kakFMg+8WGxn0c6z1fTSPVIjEY1Wr7jzc= github.com/cespare/reflex v0.3.2 h1:SBN/trM94Ifs/ozz77cR3KxKm4dNE22zfG+0+54y5bQ= github.com/cespare/reflex v0.3.2/go.mod h1:3hfHPnuDWHtNWk0aLKwwP6pomRkS3r2nM127108jY/4= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= @@ -47,12 +47,12 @@ github.com/djherbis/times v1.6.0 h1:w2ctJ92J8fBvWPxugmXIv7Nz7Q3iDMKNx9v5ocVH20c= github.com/djherbis/times v1.6.0/go.mod h1:gOHeRAz2h+VJNZ5Gmc/o7iD9k4wW7NMVqieYCY99oc0= github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI= github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= -github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= -github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/dustin/go-humanize v1.1.0 h1:dbKTrvD0klcbBV/h4AWJdMuZogJACoMlvWIWZ5b2xWg= +github.com/dustin/go-humanize v1.1.0/go.mod h1:hc1CvRkJMsgxqjmjMQF3QNRAZBwY8AXBAzKYoSX9sFI= github.com/dylibso/observe-sdk/go v0.0.0-20240828172851-9145d8ad07e1 h1:idfl8M8rPW93NehFw5H1qqH8yG158t5POr+LX9avbJY= github.com/dylibso/observe-sdk/go v0.0.0-20240828172851-9145d8ad07e1/go.mod h1:C8DzXehI4zAbrdlbtOByKX6pfivJTBiV9Jjqv56Yd9Q= -github.com/ebitengine/purego v0.10.2 h1:W809HbnvzAxgdm+aOvlSekrM16wGCdT/e76+9tS7gzE= -github.com/ebitengine/purego v0.10.2/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI0mRfJZTmQ= +github.com/ebitengine/purego v0.11.1 h1:2zpWRSQNVKN4eKsKO9eM1ILDgWfYMY9GwqRmK6XeQ/0= +github.com/ebitengine/purego v0.11.1/go.mod h1:DCHPP08djqhNSoTfImcnHYQRZmd0qhakvrozqaEYhGQ= github.com/extism/go-sdk v1.7.1 h1:lWJos6uY+tRFdlIHR+SJjwFDApY7OypS/2nMhiVQ9Sw= github.com/extism/go-sdk v1.7.1/go.mod h1:IT+Xdg5AZM9hVtpFUA+uZCJMge/hbvshl8bwzLtFyKA= github.com/fatih/structs v1.1.0 h1:Q7juDM0QtcnhCpeyLGQKyg4TOIghuNXrkL32pHAUMxo= @@ -88,8 +88,8 @@ github.com/go-viper/encoding/ini v0.1.1 h1:MVWY7B2XNw7lnOqHutGRc97bF3rP7omOdgjdM github.com/go-viper/encoding/ini v0.1.1/go.mod h1:Pfi4M2V1eAGJVZ5q6FrkHPhtHED2YgLlXhvgMVrB+YQ= github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= -github.com/gobwas/glob v0.2.3 h1:A4xDbljILXROh+kObIiy5kIaPYD8e96x1tgBhUI5J+Y= -github.com/gobwas/glob v0.2.3/go.mod h1:d3Ez4x06l9bZtSvzIay5+Yzi0fmZzPgnTbPcKjJAkT8= +github.com/gobwas/glob v1.0.0 h1:p+FKbLEIsK1yZ39/OINwFvqNb5oyPY4H8xcy6uYu8dg= +github.com/gobwas/glob v1.0.0/go.mod h1:oWCdo522i2P1n/hMXGNWs7yoV4wy/ciZuUIbvKj5rkc= github.com/goccy/go-json v0.10.6 h1:p8HrPJzOakx/mn/bQtjgNjdTcN+/S6FcG2CTtQOrHVU= github.com/goccy/go-json v0.10.6/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= @@ -101,8 +101,8 @@ github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-pipeline v0.0.0-20230411140531-6cbedfc1d3fc h1:hd+uUVsB1vdxohPneMrhGH2YfQuH5hRIK9u4/XCeUtw= github.com/google/go-pipeline v0.0.0-20230411140531-6cbedfc1d3fc/go.mod h1:SL66SJVysrh7YbDCP9tH30b8a9o/N2HeiQNUm85EKhc= -github.com/google/pprof v0.0.0-20260825171938-4d453200e7d9 h1:dl4UZiszMU+NKHirOiCKTC+hRuNAQ0moHPxSg6WcU1o= -github.com/google/pprof v0.0.0-20260825171938-4d453200e7d9/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk= +github.com/google/pprof v0.0.0-20260906184651-6331bc6350fe h1:QAinXoAFJdGQYztXn3VpFey7KCwpedbZ/EkzbplQ0cY= +github.com/google/pprof v0.0.0-20260906184651-6331bc6350fe/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk= github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE= github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -180,10 +180,10 @@ github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOF github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/ogier/pflag v0.0.1 h1:RW6JSWSu/RkSatfcLtogGfFgpim5p7ARQ10ECk5O750= github.com/ogier/pflag v0.0.1/go.mod h1:zkFki7tvTa0tafRvTBIZTvzYyAu6kQhPZFnshFFPE+g= -github.com/onsi/ginkgo/v2 v2.32.2 h1:2o6vyFvR6snrJWgRVztC+OwuqqPEMI1UzYl2s2iU7Cg= -github.com/onsi/ginkgo/v2 v2.32.2/go.mod h1:+aXOY+vzZ5mu2iI2HpTZUPmM//oQfsNFX6gU9kNcA44= -github.com/onsi/gomega v1.43.0 h1:VlG/1FxqNxhSO+lq/OHBNaaqwiBK/mO8JbVkX9Y+FeU= -github.com/onsi/gomega v1.43.0/go.mod h1:REff/hsDsodHoKlWsP2mAPhu1+5/6hVYNf9rIEBpeSg= +github.com/onsi/ginkgo/v2 v2.33.0 h1:C8gBA6Uc2ZEubiV+SXiu5tZnMTwEmXHgkJwGozKtZf8= +github.com/onsi/ginkgo/v2 v2.33.0/go.mod h1:+aXOY+vzZ5mu2iI2HpTZUPmM//oQfsNFX6gU9kNcA44= +github.com/onsi/gomega v1.44.0 h1:eAiGl3Pw5jz5GQdDff0BcxYpAX1JxW8xD7mFUuwNfZQ= +github.com/onsi/gomega v1.44.0/go.mod h1:e/C2HwaZ1DhvjzXXuFhcR7hY7Sh9pl7MmoWKEjzwcdA= github.com/pelletier/go-toml/v2 v2.4.3 h1:GTRvJQutkOSftxIFD5xw9aepkYNuPWmVJpffdDPYVpY= github.com/pelletier/go-toml/v2 v2.4.3/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA= @@ -309,8 +309,8 @@ golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5h 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= -golang.org/x/telemetry v0.0.0-20260811182544-a038080d80e5 h1:ZUSxONxc981v7AW7QUg+I9WwZzSTTJ019ENBYr5pV/Q= -golang.org/x/telemetry v0.0.0-20260811182544-a038080d80e5/go.mod h1:LVehoXe41cL5SCVQilsV7Gg6BNG+Js6P9PhSbYTIUkQ= +golang.org/x/telemetry v0.0.0-20260908163034-4bcc4b2ee518 h1:F5BWKvW126NXR74uxkxuc1jQHhm/rwm/J3rSiFyuRs4= +golang.org/x/telemetry v0.0.0-20260908163034-4bcc4b2ee518/go.mod h1:i+ivNqjDnTF3WTElsdk5g9V5DTSBYgdNo7xTU9SDwYA= golang.org/x/term v0.46.0 h1:3+OXuTbaKDgwk8jTi3aSLHRlmWqHEUDUtxnbFigO4YE= golang.org/x/term v0.46.0/go.mod h1:+K02xbkittuwc0Am4abfA3Fc+XRGXkvBXNO88NCXPoc= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= @@ -320,8 +320,8 @@ 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.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI= -golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo= +golang.org/x/tools v0.50.0 h1:c2ifzfcuY7L90lZ2aKd8S4K2NpASF08SZx9ZuJkHmSU= +golang.org/x/tools v0.50.0/go.mod h1:7ulVMw3831Mwi5EZD6RomGyffr4VFjuNYXf2BbCEAV0= google.golang.org/appengine v1.6.5/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc= google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= From 4113d954ca0cbfe5e9f89a919c09f0fe6f7d1339 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sat, 26 Sep 2026 11:43:00 -0400 Subject: [PATCH 15/26] fix(scanner): register .webm as audio/webm (#6230) Go 1.27 changed the built-in MIME type for .webm from audio/webm to video/webm, so the scanner stopped treating WebM files as audio after the Go bump in 0.64.0. Map .webm to audio/webm in mime_types.yaml so it no longer depends on the Go version. --- model/file_types_test.go | 4 ++++ resources/mime_types.yaml | 1 + 2 files changed, 5 insertions(+) diff --git a/model/file_types_test.go b/model/file_types_test.go index 93301e151..07dac645e 100644 --- a/model/file_types_test.go +++ b/model/file_types_test.go @@ -18,6 +18,10 @@ var _ = Describe("File Types()", func() { Expect(model.IsAudioFile("test.flac")).To(BeTrue()) }) + It("returns true for a WebM file", func() { + Expect(model.IsAudioFile("test.webm")).To(BeTrue()) + }) + It("returns false for a non-audio file", func() { Expect(model.IsAudioFile("test.jpg")).To(BeFalse()) }) diff --git a/resources/mime_types.yaml b/resources/mime_types.yaml index 18a2c22b5..1a963bdf7 100644 --- a/resources/mime_types.yaml +++ b/resources/mime_types.yaml @@ -29,6 +29,7 @@ types: .wvp: audio/x-wavpack .tak: audio/tak .mka: audio/x-matroska + .webm: audio/webm # Image .gif: image/gif From bb7d81a5ea1b4ad7d5479c82c9abbabc6c66a359 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sat, 26 Sep 2026 12:30:52 -0400 Subject: [PATCH 16/26] refactor(scanner): remove the unused legacy ffmpeg metadata extractor (#6231) * refactor(scanner): remove the legacy ffmpeg metadata extractor The ffmpeg extractor in scanner/metadata_old has not been wired into the scanner since the taglib-only rewrite, so it was only exercised by its own tests. Remove the package, the FFmpeg.Probe method and its ffmetadata command that only it used, and the startup fallback for Scanner.Extractor="ffmpeg". Configs that still set it keep working: unknown extractors already fall back to taglib with a warning. * fix(conf): warn and fall back to taglib for an unknown Scanner.Extractor Validate the option when loading the config, so invalid values such as the removed "ffmpeg" extractor are reported once at startup instead of only when a library storage is created. --- conf/configuration.go | 5 + conf/configuration_test.go | 21 + core/ffmpeg/ffmpeg.go | 27 -- core/ffmpeg/ffmpeg_test.go | 15 +- scanner/metadata_old/ffmpeg/ffmpeg.go | 211 --------- .../metadata_old/ffmpeg/ffmpeg_suite_test.go | 17 - scanner/metadata_old/ffmpeg/ffmpeg_test.go | 375 ---------------- scanner/metadata_old/metadata.go | 411 ------------------ .../metadata_old/metadata_internal_test.go | 144 ------ scanner/metadata_old/metadata_suite_test.go | 17 - scanner/metadata_old/metadata_test.go | 95 ---- server/initial_setup.go | 4 - tests/harness/harness.go | 2 - tests/mock_ffmpeg.go | 6 - 14 files changed, 30 insertions(+), 1320 deletions(-) delete mode 100644 scanner/metadata_old/ffmpeg/ffmpeg.go delete mode 100644 scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go delete mode 100644 scanner/metadata_old/ffmpeg/ffmpeg_test.go delete mode 100644 scanner/metadata_old/metadata.go delete mode 100644 scanner/metadata_old/metadata_internal_test.go delete mode 100644 scanner/metadata_old/metadata_suite_test.go delete mode 100644 scanner/metadata_old/metadata_test.go diff --git a/conf/configuration.go b/conf/configuration.go index 6f295ebeb..0173f5628 100644 --- a/conf/configuration.go +++ b/conf/configuration.go @@ -514,6 +514,11 @@ func Load(noConfigDump bool) { Server.UICoverArtSize = newValue } + if Server.Scanner.Extractor != consts.DefaultScannerExtractor { + log.Warn("Invalid Scanner.Extractor, using default", "value", Server.Scanner.Extractor, "default", consts.DefaultScannerExtractor) + Server.Scanner.Extractor = consts.DefaultScannerExtractor + } + // Floor MaxImageSize at MaxImageUploadSize so accepted uploads can always be read back. imgSize, _ := humanize.ParseBytes(Server.MaxImageSize) uploadSize, _ := humanize.ParseBytes(Server.MaxImageUploadSize) diff --git a/conf/configuration_test.go b/conf/configuration_test.go index 9404f01bc..8c4c8ab86 100644 --- a/conf/configuration_test.go +++ b/conf/configuration_test.go @@ -402,6 +402,27 @@ var _ = Describe("Configuration", func() { }) }) + Describe("Scanner.Extractor", func() { + BeforeEach(func() { + viper.Reset() + conf.SetViperDefaults() + viper.SetDefault("datafolder", GinkgoT().TempDir()) + viper.SetDefault("loglevel", "error") + conf.ResetConf() + }) + + It("falls back to taglib for an unknown extractor", func() { + viper.SetDefault("scanner.extractor", "ffmpeg") + conf.Load(true) + Expect(conf.Server.Scanner.Extractor).To(Equal("taglib")) + }) + + It("keeps taglib", func() { + conf.Load(true) + Expect(conf.Server.Scanner.Extractor).To(Equal("taglib")) + }) + }) + Describe("EnforceNonRootUser", func() { It("defaults to false", func() { conf.Load(true) diff --git a/core/ffmpeg/ffmpeg.go b/core/ffmpeg/ffmpeg.go index cc38dd9de..af59178af 100644 --- a/core/ffmpeg/ffmpeg.go +++ b/core/ffmpeg/ffmpeg.go @@ -49,7 +49,6 @@ type FFmpeg interface { Transcode(ctx context.Context, opts TranscodeOptions) (io.ReadCloser, error) ExtractImage(ctx context.Context, path string) (io.ReadCloser, error) ConvertAnimatedImage(ctx context.Context, reader io.Reader, maxSize int, quality int) (io.ReadCloser, error) - Probe(ctx context.Context, files []string) (string, error) ProbeAudioStream(ctx context.Context, filePath string) (*AudioProbeResult, error) CmdPath() (string, error) IsAvailable() bool @@ -68,7 +67,6 @@ var ErrAnimatedWebPUnsupported = errors.New("ffmpeg lacks libwebp_anim encoder const ( extractImageCmd = "ffmpeg -i %s -map 0:v -map -0:V -vcodec copy -f image2pipe -" - probeCmd = "ffmpeg %s -f ffmetadata" probeAudioStreamCmd = "ffprobe -v error -select_streams a:0 -print_format json -show_streams -show_format %s" ) @@ -149,17 +147,6 @@ func fileExists(path string) error { return nil } -func (e *ffmpeg) Probe(ctx context.Context, files []string) (string, error) { - if _, err := ffmpegCmd(); err != nil { - return "", err - } - args := createProbeCommand(probeCmd, files) - log.Trace(ctx, "Executing ffmpeg command", "args", args) - cmd := exec.CommandContext(ctx, args[0], args[1:]...) // #nosec - output, _ := cmd.CombinedOutput() - return string(output), nil -} - func (e *ffmpeg) ProbeAudioStream(ctx context.Context, filePath string) (*AudioProbeResult, error) { if _, err := ffmpegCmd(); err != nil { return nil, err @@ -593,20 +580,6 @@ func createFFmpegCommand(cmd, path string, maxBitRate, offset int) []string { return args } -func createProbeCommand(cmd string, inputs []string) []string { - var args []string - for _, s := range fixCmd(cmd) { - if s == "%s" { - for _, inp := range inputs { - args = append(args, "-i", inp) - } - } else { - args = append(args, s) - } - } - return args -} - func fixCmd(cmd string) []string { split := strings.Fields(cmd) cmdPath, _ := ffmpegCmd() diff --git a/core/ffmpeg/ffmpeg_test.go b/core/ffmpeg/ffmpeg_test.go index dbc8fa3c8..46684fe14 100644 --- a/core/ffmpeg/ffmpeg_test.go +++ b/core/ffmpeg/ffmpeg_test.go @@ -62,23 +62,16 @@ var _ = Describe("ffmpeg", func() { }) }) - Describe("createProbeCommand", func() { - It("creates a valid command line", func() { - args := createProbeCommand(probeCmd, []string{"/music library/one.mp3", "/music library/two.mp3"}) - Expect(args).To(Equal([]string{"ffmpeg", "-i", "/music library/one.mp3", "-i", "/music library/two.mp3", "-f", "ffmetadata"})) - }) - }) - When("ffmpegPath is set", func() { It("returns the correct ffmpeg path", func() { ffmpegPath = "/usr/bin/ffmpeg" - args := createProbeCommand(probeCmd, []string{"one.mp3"}) - Expect(args).To(Equal([]string{"/usr/bin/ffmpeg", "-i", "one.mp3", "-f", "ffmetadata"})) + args := createFFmpegCommand("ffmpeg -i %s -f mp3 -", "one.mp3", 0, 0) + Expect(args).To(Equal([]string{"/usr/bin/ffmpeg", "-i", "one.mp3", "-f", "mp3", "-"})) }) It("returns the correct ffmpeg path with spaces", func() { ffmpegPath = "/usr/bin/with spaces/ffmpeg.exe" - args := createProbeCommand(probeCmd, []string{"one.mp3"}) - Expect(args).To(Equal([]string{"/usr/bin/with spaces/ffmpeg.exe", "-i", "one.mp3", "-f", "ffmetadata"})) + args := createFFmpegCommand("ffmpeg -i %s -f mp3 -", "one.mp3", 0, 0) + Expect(args).To(Equal([]string{"/usr/bin/with spaces/ffmpeg.exe", "-i", "one.mp3", "-f", "mp3", "-"})) }) }) diff --git a/scanner/metadata_old/ffmpeg/ffmpeg.go b/scanner/metadata_old/ffmpeg/ffmpeg.go deleted file mode 100644 index 8fc496c02..000000000 --- a/scanner/metadata_old/ffmpeg/ffmpeg.go +++ /dev/null @@ -1,211 +0,0 @@ -package ffmpeg - -import ( - "bufio" - "context" - "errors" - "regexp" - "strconv" - "strings" - "time" - - "github.com/navidrome/navidrome/core/ffmpeg" - "github.com/navidrome/navidrome/log" - "github.com/navidrome/navidrome/scanner/metadata_old" -) - -const ExtractorID = "ffmpeg" - -type Extractor struct { - ffmpeg ffmpeg.FFmpeg -} - -func (e *Extractor) Parse(files ...string) (map[string]metadata_old.ParsedTags, error) { - output, err := e.ffmpeg.Probe(context.TODO(), files) - if err != nil { - log.Error("Cannot use ffmpeg to extract tags. Aborting", err) - return nil, err - } - fileTags := map[string]metadata_old.ParsedTags{} - if len(output) == 0 { - return fileTags, errors.New("error extracting metadata files") - } - infos := e.parseOutput(output) - for file, info := range infos { - tags, err := e.extractMetadata(file, info) - // Skip files with errors - if err == nil { - fileTags[file] = tags - } - } - return fileTags, nil -} - -func (e *Extractor) CustomMappings() metadata_old.ParsedTags { - return metadata_old.ParsedTags{ - "disc": {"tpa"}, - "has_picture": {"metadata_block_picture"}, - "originaldate": {"tdor"}, - } -} - -func (e *Extractor) Version() string { - return e.ffmpeg.Version() -} - -func (e *Extractor) extractMetadata(filePath, info string) (metadata_old.ParsedTags, error) { - tags := e.parseInfo(info) - if len(tags) == 0 { - log.Trace("Not a media file. Skipping", "filePath", filePath) - return nil, errors.New("not a media file") - } - - return tags, nil -} - -var ( - // Input #0, mp3, from 'groovin.mp3': - inputRegex = regexp.MustCompile(`(?m)^Input #\d+,.*,\sfrom\s'(.*)'`) - - // TITLE : Back In Black - tagsRx = regexp.MustCompile(`(?i)^\s{4,6}([\w\s-]+)\s*:(.*)`) - - // : Second comment line - continuationRx = regexp.MustCompile(`(?i)^\s+:(.*)`) - - // Duration: 00:04:16.00, start: 0.000000, bitrate: 995 kb/s` - durationRx = regexp.MustCompile(`^\s\sDuration: ([\d.:]+).*bitrate: (\d+)`) - - // Stream #0:0: Audio: mp3, 44100 Hz, stereo, fltp, 192 kb/s - bitRateRx = regexp.MustCompile(`^\s{2,4}Stream #\d+:\d+: Audio:.*, (\d+) kb/s`) - - // Stream #0:0: Audio: mp3, 44100 Hz, stereo, fltp, 192 kb/s - // Stream #0:0: Audio: flac, 44100 Hz, stereo, s16 - // Stream #0:0: Audio: dsd_lsbf_planar, 352800 Hz, stereo, fltp, 5644 kb/s - audioStreamRx = regexp.MustCompile(`^\s{2,4}Stream #\d+:\d+.*: Audio: (.*), (.*) Hz, ([\w.]+),*(.*.,)*`) - - // Stream #0:1: Video: mjpeg, yuvj444p(pc, bt470bg/unknown/unknown), 600x600 [SAR 1:1 DAR 1:1], 90k tbr, 90k tbn, 90k tbc` - coverRx = regexp.MustCompile(`^\s{2,4}Stream #\d+:.+: (Video):.*`) -) - -func (e *Extractor) parseOutput(output string) map[string]string { - outputs := map[string]string{} - all := inputRegex.FindAllStringSubmatchIndex(output, -1) - for i, loc := range all { - // Filename is the first captured group - file := output[loc[2]:loc[3]] - - // File info is everything from the match, up until the beginning of the next match - info := "" - initial := loc[1] - if i < len(all)-1 { - end := all[i+1][0] - 1 - info = output[initial:end] - } else { - // if this is the last match - info = output[initial:] - } - outputs[file] = info - } - return outputs -} - -func (e *Extractor) parseInfo(info string) map[string][]string { - tags := map[string][]string{} - - reader := strings.NewReader(info) - scanner := bufio.NewScanner(reader) - lastTag := "" - for scanner.Scan() { - line := scanner.Text() - if len(line) == 0 { - continue - } - match := tagsRx.FindStringSubmatch(line) - if len(match) > 0 { - tagName := strings.TrimSpace(strings.ToLower(match[1])) - if tagName != "" { - tagValue := strings.TrimSpace(match[2]) - tags[tagName] = append(tags[tagName], tagValue) - lastTag = tagName - continue - } - } - - if lastTag != "" { - match = continuationRx.FindStringSubmatch(line) - if len(match) > 0 { - if tags[lastTag] == nil { - tags[lastTag] = []string{""} - } - tagValue := tags[lastTag][0] - tags[lastTag][0] = tagValue + "\n" + strings.TrimSpace(match[1]) - continue - } - } - - lastTag = "" - match = coverRx.FindStringSubmatch(line) - if len(match) > 0 { - tags["has_picture"] = []string{"true"} - continue - } - - match = durationRx.FindStringSubmatch(line) - if len(match) > 0 { - tags["duration"] = []string{e.parseDuration(match[1])} - if len(match) > 1 { - tags["bitrate"] = []string{match[2]} - } - continue - } - - match = bitRateRx.FindStringSubmatch(line) - if len(match) > 0 { - tags["bitrate"] = []string{match[1]} - } - - match = audioStreamRx.FindStringSubmatch(line) - if len(match) > 0 { - tags["samplerate"] = []string{match[2]} - tags["channels"] = []string{e.parseChannels(match[3])} - } - } - - comment := tags["comment"] - if len(comment) > 0 && comment[0] == "Cover (front)" { - delete(tags, "comment") - } - - return tags -} - -var zeroTime = time.Date(0000, time.January, 1, 0, 0, 0, 0, time.UTC) - -func (e *Extractor) parseDuration(tag string) string { - d, err := time.Parse("15:04:05", tag) - if err != nil { - return "0" - } - return strconv.FormatFloat(d.Sub(zeroTime).Seconds(), 'f', 2, 32) -} - -func (e *Extractor) parseChannels(tag string) string { - switch tag { - case "mono": - return "1" - case "stereo": - return "2" - case "5.1": - return "6" - case "7.1": - return "8" - default: - return "0" - } -} - -// Inputs will always be absolute paths -func init() { - metadata_old.RegisterExtractor(ExtractorID, &Extractor{ffmpeg: ffmpeg.New()}) -} diff --git a/scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go b/scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go deleted file mode 100644 index 815940381..000000000 --- a/scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go +++ /dev/null @@ -1,17 +0,0 @@ -package ffmpeg - -import ( - "testing" - - "github.com/navidrome/navidrome/log" - "github.com/navidrome/navidrome/tests" - . "github.com/onsi/ginkgo/v2" - . "github.com/onsi/gomega" -) - -func TestFFMpeg(t *testing.T) { - tests.Init(t, true) - log.SetLevel(log.LevelFatal) - RegisterFailHandler(Fail) - RunSpecs(t, "FFMpeg Suite") -} diff --git a/scanner/metadata_old/ffmpeg/ffmpeg_test.go b/scanner/metadata_old/ffmpeg/ffmpeg_test.go deleted file mode 100644 index 6c7f43a5d..000000000 --- a/scanner/metadata_old/ffmpeg/ffmpeg_test.go +++ /dev/null @@ -1,375 +0,0 @@ -package ffmpeg - -import ( - . "github.com/onsi/ginkgo/v2" - . "github.com/onsi/gomega" -) - -var _ = Describe("Extractor", func() { - var e *Extractor - BeforeEach(func() { - e = &Extractor{} - }) - - Context("extractMetadata", func() { - It("extracts MusicBrainz custom tags", func() { - const output = ` -Input #0, ape, from './Capture/02 01 - Symphony No. 5 in C minor, Op. 67 I. Allegro con brio - Ludwig van Beethoven.ape': - Metadata: - ALBUM : Forever Classics - ARTIST : Ludwig van Beethoven - TITLE : Symphony No. 5 in C minor, Op. 67: I. Allegro con brio - MUSICBRAINZ_ALBUMSTATUS: official - MUSICBRAINZ_ALBUMTYPE: album - MusicBrainz_AlbumComment: MP3 - Musicbrainz_Albumid: 71eb5e4a-90e2-4a31-a2d1-a96485fcb667 - musicbrainz_trackid: ffe06940-727a-415a-b608-b7e45737f9d8 - Musicbrainz_Artistid: 1f9df192-a621-4f54-8850-2c5373b7eac9 - Musicbrainz_Albumartistid: 89ad4ac3-39f7-470e-963a-56509c546377 - Musicbrainz_Releasegroupid: 708b1ae1-2d3d-34c7-b764-2732b154f5b6 - musicbrainz_releasetrackid: 6fee2e35-3049-358f-83be-43b36141028b - CatalogNumber : PLD 1201 -` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(SatisfyAll( - HaveKeyWithValue("catalognumber", []string{"PLD 1201"}), - HaveKeyWithValue("musicbrainz_trackid", []string{"ffe06940-727a-415a-b608-b7e45737f9d8"}), - HaveKeyWithValue("musicbrainz_albumid", []string{"71eb5e4a-90e2-4a31-a2d1-a96485fcb667"}), - HaveKeyWithValue("musicbrainz_artistid", []string{"1f9df192-a621-4f54-8850-2c5373b7eac9"}), - HaveKeyWithValue("musicbrainz_albumartistid", []string{"89ad4ac3-39f7-470e-963a-56509c546377"}), - HaveKeyWithValue("musicbrainz_albumtype", []string{"album"}), - HaveKeyWithValue("musicbrainz_albumcomment", []string{"MP3"}), - )) - }) - - It("detects embedded cover art correctly", func() { - const output = ` -Input #0, mp3, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.mp3': - Metadata: - compilation : 1 - Duration: 00:00:01.02, start: 0.000000, bitrate: 477 kb/s - Stream #0:0: Audio: mp3, 44100 Hz, stereo, fltp, 192 kb/s - Stream #0:1: Video: mjpeg, yuvj444p(pc, bt470bg/unknown/unknown), 600x600 [SAR 1:1 DAR 1:1], 90k tbr, 90k tbn, 90k tbc` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("has_picture", []string{"true"})) - }) - - It("detects embedded cover art in ffmpeg 4.4 output", func() { - const output = ` -Input #0, flac, from '/run/media/naomi/Archivio/Musica/Katy Perry/Chained to the Rhythm/01 Katy Perry featuring Skip Marley - Chained to the Rhythm.flac': - Metadata: - ARTIST : Katy Perry featuring Skip Marley - Duration: 00:03:57.91, start: 0.000000, bitrate: 983 kb/s - Stream #0:0: Audio: flac, 44100 Hz, stereo, s16 - Stream #0:1: Video: mjpeg (Baseline), yuvj444p(pc, bt470bg/unknown/unknown), 599x518, 90k tbr, 90k tbn, 90k tbc (attached pic) - Metadata: - comment : Cover (front)` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("has_picture", []string{"true"})) - }) - - It("detects embedded cover art in ogg containers", func() { - const output = ` -Input #0, ogg, from '/Users/deluan/Music/iTunes/iTunes Media/Music/_Testes/Jamaican In New York/01-02 Jamaican In New York (Album Version).opus': - Duration: 00:04:28.69, start: 0.007500, bitrate: 139 kb/s - Stream #0:0(eng): Audio: opus, 48000 Hz, stereo, fltp - Metadata: - ALBUM : Jamaican In New York - metadata_block_picture: AAAAAwAAAAppbWFnZS9qcGVnAAAAAAAAAAAAAAAAAAAAAAAAAAAAA4Id/9j/4AAQSkZJRgABAQEAYABgAAD/2wBDAAMCAgMCAgMDAwMEAwMEBQgFBQQEBQoHBwYIDAoMDAsKCwsNDhIQDQ4RDgsLEBYQERMUFRUVDA8XGBYUGBIUFRT/2wBDAQMEBAUEBQkFBQkUDQsNFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQ - TITLE : Jamaican In New York (Album Version)` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKey("metadata_block_picture")) - md = md.Map(e.CustomMappings()) - Expect(md).To(HaveKey("has_picture")) - }) - - It("detects embedded cover art in m4a containers", func() { - const output = ` -Input #0, mov,mp4,m4a,3gp,3g2,mj2, from 'Putumayo Presents_ Euro Groove/01 Destins et Désirs.m4a': - Metadata: - album : Putumayo Presents: Euro Groove - Duration: 00:05:15.81, start: 0.047889, bitrate: 133 kb/s - Stream #0:0[0x1](und): Audio: aac (LC) (mp4a / 0x6134706D), 44100 Hz, stereo, fltp, 125 kb/s (default) - Metadata: - creation_time : 2008-03-11T21:03:23.000000Z - vendor_id : [0][0][0][0] - Stream #0:1[0x0]: Video: png, rgb24(pc), 350x350, 90k tbr, 90k tbn (attached pic) -` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("has_picture", []string{"true"})) - }) - - It("gets bitrate from the stream, if available", func() { - const output = ` -Input #0, mp3, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.mp3': - Duration: 00:00:01.02, start: 0.000000, bitrate: 477 kb/s - Stream #0:0: Audio: mp3, 44100 Hz, stereo, fltp, 192 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("bitrate", []string{"192"})) - }) - - It("parses duration with milliseconds", func() { - const output = ` -Input #0, mp3, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.mp3': - Duration: 00:05:02.63, start: 0.000000, bitrate: 140 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("duration", []string{"302.63"})) - }) - - It("parse flac bitrates", func() { - const output = ` -Input #0, mp3, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.mp3': - Duration: 00:00:01.02, start: 0.000000, bitrate: 477 kb/s - Stream #0:0: Audio: mp3, 44100 Hz, stereo, fltp, 192 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("channels", []string{"2"})) - }) - - It("parse channels from the stream with bitrate", func() { - const output = ` -Input #0, flac, from '/Users/deluan/Music/Music/Media/__/Crazy For You/01-01 Crazy For You.flac': - Metadata: - TITLE : Crazy For You - Duration: 00:04:13.00, start: 0.000000, bitrate: 852 kb/s - Stream #0:0: Audio: flac, 44100 Hz, stereo, s16 - Stream #0:1: Video: mjpeg (Progressive), yuvj444p(pc, bt470bg/unknown/unknown), 600x600, 90k tbr, 90k tbn, 90k tbc (attached pic) -` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("bitrate", []string{"852"})) - }) - - It("parse 7.1 channels from the stream", func() { - const output = ` -Input #0, wav, from '/Users/deluan/Music/Music/Media/_/multichannel/Nums_7dot1_24_48000.wav': - Duration: 00:00:09.05, bitrate: 9216 kb/s - Stream #0:0: Audio: pcm_s24le ([1][0][0][0] / 0x0001), 48000 Hz, 7.1, s32 (24 bit), 9216 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("channels", []string{"8"})) - }) - - It("parse channels from the stream without bitrate", func() { - const output = ` -Input #0, flac, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.flac': - Duration: 00:00:01.02, start: 0.000000, bitrate: 1371 kb/s - Stream #0:0: Audio: flac, 44100 Hz, stereo, fltp, s32 (24 bit)` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("channels", []string{"2"})) - }) - - It("parse channels from the stream with lang", func() { - const output = ` -Input #0, flac, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.m4a': - Duration: 00:00:01.02, start: 0.000000, bitrate: 1371 kb/s - Stream #0:0(eng): Audio: aac (LC) (mp4a / 0x6134706D), 44100 Hz, stereo, fltp, 262 kb/s (default)` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("channels", []string{"2"})) - }) - - It("parse channels from the stream with lang 2", func() { - const output = ` -Input #0, flac, from '/Users/deluan/Music/iTunes/iTunes Media/Music/Compilations/Putumayo Presents Blues Lounge/09 Pablo's Blues.m4a': - Duration: 00:00:01.02, start: 0.000000, bitrate: 1371 kb/s - Stream #0:0(eng): Audio: vorbis, 44100 Hz, stereo, fltp, 192 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("channels", []string{"2"})) - }) - - It("parse sampleRate from the stream", func() { - const output = ` -Input #0, dsf, from '/Users/deluan/Downloads/06-04 Perpetual Change.dsf': - Duration: 00:14:19.46, start: 0.000000, bitrate: 5644 kb/s - Stream #0:0: Audio: dsd_lsbf_planar, 352800 Hz, stereo, fltp, 5644 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("samplerate", []string{"352800"})) - }) - - It("parse sampleRate from the stream", func() { - const output = ` -Input #0, wav, from '/Users/deluan/Music/Music/Media/_/multichannel/Nums_7dot1_24_48000.wav': - Duration: 00:00:09.05, bitrate: 9216 kb/s - Stream #0:0: Audio: pcm_s24le ([1][0][0][0] / 0x0001), 48000 Hz, 7.1, s32 (24 bit), 9216 kb/s` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("samplerate", []string{"48000"})) - }) - - It("parses stream level tags", func() { - const output = ` -Input #0, ogg, from './01-02 Drive (Teku).opus': - Metadata: - ALBUM : Hot Wheels Acceleracers Soundtrack - Duration: 00:03:37.37, start: 0.007500, bitrate: 135 kb/s - Stream #0:0(eng): Audio: opus, 48000 Hz, stereo, fltp - Metadata: - TITLE : Drive (Teku)` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("title", []string{"Drive (Teku)"})) - }) - - It("does not overlap top level tags with the stream level tags", func() { - const output = ` -Input #0, mp3, from 'groovin.mp3': - Metadata: - title : Groovin' (feat. Daniel Sneijers, Susanne Alt) - Duration: 00:03:34.28, start: 0.025056, bitrate: 323 kb/s - Metadata: - title : garbage` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("title", []string{"Groovin' (feat. Daniel Sneijers, Susanne Alt)", "garbage"})) - }) - - It("parses multiline tags", func() { - const outputWithMultilineComment = ` -Input #0, mov,mp4,m4a,3gp,3g2,mj2, from 'modulo.m4a': - Metadata: - comment : https://www.mixcloud.com/codigorock/30-minutos-com-saara-saara/ - : - : Tracklist: - : - : 01. Saara Saara - : 02. Carta Corrente - : 03. X - : 04. Eclipse Lunar - : 05. Vírus de Sírius - : 06. Doktor Fritz - : 07. Wunderbar - : 08. Quarta Dimensão - Duration: 00:26:46.96, start: 0.052971, bitrate: 69 kb/s` - const expectedComment = `https://www.mixcloud.com/codigorock/30-minutos-com-saara-saara/ - -Tracklist: - -01. Saara Saara -02. Carta Corrente -03. X -04. Eclipse Lunar -05. Vírus de Sírius -06. Doktor Fritz -07. Wunderbar -08. Quarta Dimensão` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", outputWithMultilineComment) - Expect(md).To(HaveKeyWithValue("comment", []string{expectedComment})) - }) - - It("parses sort tags correctly", func() { - const output = ` -Input #0, mp3, from '/Users/deluan/Downloads/椎名林檎 - 加爾基 精液 栗ノ花 - 2003/02 - ドツペルゲンガー.mp3': - Metadata: - title-sort : Dopperugengā - album : 加爾基 精液 栗ノ花 - artist : 椎名林檎 - album_artist : 椎名林檎 - title : ドツペルゲンガー - albumsort : Kalk Samen Kuri No Hana - artist_sort : Shiina, Ringo - ALBUMARTISTSORT : Shiina, Ringo -` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(SatisfyAll( - HaveKeyWithValue("title", []string{"ドツペルゲンガー"}), - HaveKeyWithValue("album", []string{"加爾基 精液 栗ノ花"}), - HaveKeyWithValue("artist", []string{"椎名林檎"}), - HaveKeyWithValue("album_artist", []string{"椎名林檎"}), - HaveKeyWithValue("title-sort", []string{"Dopperugengā"}), - HaveKeyWithValue("albumsort", []string{"Kalk Samen Kuri No Hana"}), - HaveKeyWithValue("artist_sort", []string{"Shiina, Ringo"}), - HaveKeyWithValue("albumartistsort", []string{"Shiina, Ringo"}), - )) - }) - - It("ignores cover comment", func() { - const output = ` -Input #0, mp3, from './Edie Brickell/Picture Perfect Morning/01-01 Tomorrow Comes.mp3': - Metadata: - title : Tomorrow Comes - artist : Edie Brickell - Duration: 00:03:56.12, start: 0.000000, bitrate: 332 kb/s - Stream #0:0: Audio: mp3, 44100 Hz, stereo, s16p, 320 kb/s - Stream #0:1: Video: mjpeg, yuvj420p(pc, bt470bg/unknown/unknown), 1200x1200 [SAR 72:72 DAR 1:1], 90k tbr, 90k tbn, 90k tbc - Metadata: - comment : Cover (front)` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).ToNot(HaveKey("comment")) - }) - - It("parses tags with spaces in the name", func() { - const output = ` -Input #0, mp3, from '/Users/deluan/Music/Music/Media/_/Wyclef Jean - From the Hut, to the Projects, to the Mansion/10 - The Struggle (interlude).mp3': - Metadata: - ALBUM ARTIST : Wyclef Jean -` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("album artist", []string{"Wyclef Jean"})) - }) - }) - - It("parses an integer TBPM tag", func() { - const output = ` - Input #0, mp3, from 'tests/fixtures/test.mp3': - Metadata: - TBPM : 123` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("tbpm", []string{"123"})) - }) - - It("parses and rounds a floating point fBPM tag", func() { - const output = ` - Input #0, ogg, from 'tests/fixtures/test.ogg': - Metadata: - FBPM : 141.7` - md, _ := e.extractMetadata("tests/fixtures/test.ogg", output) - Expect(md).To(HaveKeyWithValue("fbpm", []string{"141.7"})) - }) - - It("parses replaygain data correctly", func() { - const output = ` - Input #0, mp3, from 'test.mp3': - Metadata: - REPLAYGAIN_ALBUM_PEAK: 0.9125 - REPLAYGAIN_TRACK_PEAK: 0.4512 - REPLAYGAIN_TRACK_GAIN: -1.48 dB - REPLAYGAIN_ALBUM_GAIN: +3.21518 dB - Side data: - replaygain: track gain - -1.480000, track peak - 0.000011, album gain - 3.215180, album peak - 0.000021, - ` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(SatisfyAll( - HaveKeyWithValue("replaygain_track_gain", []string{"-1.48 dB"}), - HaveKeyWithValue("replaygain_track_peak", []string{"0.4512"}), - HaveKeyWithValue("replaygain_album_gain", []string{"+3.21518 dB"}), - HaveKeyWithValue("replaygain_album_peak", []string{"0.9125"}), - )) - }) - - It("parses lyrics with language code", func() { - const output = ` - Input #0, mp3, from 'test.mp3': - Metadata: - lyrics-eng : [00:00.00]This is - : [00:02.50]English - lyrics-xxx : [00:00.00]This is - : [00:02.50]unspecified - ` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(SatisfyAll( - HaveKeyWithValue("lyrics-eng", []string{ - "[00:00.00]This is\n[00:02.50]English", - }), - HaveKeyWithValue("lyrics-xxx", []string{ - "[00:00.00]This is\n[00:02.50]unspecified", - }), - )) - }) - - It("parses normal LYRICS tag", func() { - const output = ` - Input #0, mp3, from 'test.mp3': - Metadata: - LYRICS : [00:00.00]This is - : [00:02.50]English - ` - md, _ := e.extractMetadata("tests/fixtures/test.mp3", output) - Expect(md).To(HaveKeyWithValue("lyrics", []string{ - "[00:00.00]This is\n[00:02.50]English", - })) - }) -}) diff --git a/scanner/metadata_old/metadata.go b/scanner/metadata_old/metadata.go deleted file mode 100644 index 2906a2c09..000000000 --- a/scanner/metadata_old/metadata.go +++ /dev/null @@ -1,411 +0,0 @@ -package metadata_old - -import ( - "context" - "encoding/json" - "fmt" - "math" - "os" - "path" - "regexp" - "strconv" - "strings" - "time" - - "github.com/djherbis/times" - "github.com/google/uuid" - "github.com/navidrome/navidrome/conf" - "github.com/navidrome/navidrome/consts" - "github.com/navidrome/navidrome/log" - "github.com/navidrome/navidrome/model" -) - -type Extractor interface { - Parse(files ...string) (map[string]ParsedTags, error) - CustomMappings() ParsedTags - Version() string -} - -var extractors = map[string]Extractor{} - -func RegisterExtractor(id string, parser Extractor) { - extractors[id] = parser -} - -func LogExtractors() { - for id, p := range extractors { - log.Debug("Registered metadata extractor", "id", id, "version", p.Version()) - } -} - -func Extract(files ...string) (map[string]Tags, error) { - p, ok := extractors[conf.Server.Scanner.Extractor] - if !ok { - log.Warn("Invalid 'Scanner.Extractor' option. Using default", "requested", conf.Server.Scanner.Extractor, - "validOptions", "ffmpeg,taglib", "default", consts.DefaultScannerExtractor) - p = extractors[consts.DefaultScannerExtractor] - } - - extractedTags, err := p.Parse(files...) - if err != nil { - return nil, err - } - - result := map[string]Tags{} - for filePath, tags := range extractedTags { - fileInfo, err := os.Stat(filePath) - if err != nil { - log.Warn("Error stating file. Skipping", "filePath", filePath, err) - continue - } - - tags = tags.Map(p.CustomMappings()) - result[filePath] = NewTag(filePath, fileInfo, tags) - } - - return result, nil -} - -func NewTag(filePath string, fileInfo os.FileInfo, tags ParsedTags) Tags { - for t, values := range tags { - values = removeDuplicatesAndEmpty(values) - if len(values) == 0 { - delete(tags, t) - continue - } - tags[t] = values - } - return Tags{ - filePath: filePath, - fileInfo: fileInfo, - Tags: tags, - } -} - -func removeDuplicatesAndEmpty(values []string) []string { - encountered := map[string]struct{}{} - empty := true - result := make([]string, 0, len(values)) - for _, v := range values { - if _, ok := encountered[v]; ok { - continue - } - encountered[v] = struct{}{} - empty = empty && v == "" - result = append(result, v) - } - if empty { - return nil - } - return result -} - -type ParsedTags map[string][]string - -func (p ParsedTags) Map(customMappings ParsedTags) ParsedTags { - if customMappings == nil { - return p - } - for tagName, alternatives := range customMappings { - for _, altName := range alternatives { - if altValue, ok := p[altName]; ok { - p[tagName] = append(p[tagName], altValue...) - delete(p, altName) - } - } - } - return p -} - -type Tags struct { - filePath string - fileInfo os.FileInfo - Tags ParsedTags -} - -// Common tags - -func (t Tags) Title() string { return t.getFirstTagValue("title", "sort_name", "titlesort") } -func (t Tags) Album() string { return t.getFirstTagValue("album", "sort_album", "albumsort") } -func (t Tags) Artist() string { return t.getFirstTagValue("artist", "sort_artist", "artistsort") } -func (t Tags) AlbumArtist() string { - return t.getFirstTagValue("album_artist", "album artist", "albumartist") -} -func (t Tags) SortTitle() string { return t.getSortTag("tsot", "title", "name") } -func (t Tags) SortAlbum() string { return t.getSortTag("tsoa", "album") } -func (t Tags) SortArtist() string { return t.getSortTag("tsop", "artist") } -func (t Tags) SortAlbumArtist() string { return t.getSortTag("tso2", "albumartist", "album_artist") } -func (t Tags) Genres() []string { return t.getAllTagValues("genre") } -func (t Tags) Date() (int, string) { return t.getDate("date") } -func (t Tags) OriginalDate() (int, string) { return t.getDate("originaldate") } -func (t Tags) ReleaseDate() (int, string) { return t.getDate("releasedate") } -func (t Tags) Comment() string { return t.getFirstTagValue("comment") } -func (t Tags) Compilation() bool { return t.getBool("tcmp", "compilation", "wm/iscompilation") } -func (t Tags) TrackNumber() (int, int) { return t.getTuple("track", "tracknumber") } -func (t Tags) DiscNumber() (int, int) { return t.getTuple("disc", "discnumber") } -func (t Tags) DiscSubtitle() string { - return t.getFirstTagValue("tsst", "discsubtitle", "setsubtitle") -} -func (t Tags) CatalogNum() string { return t.getFirstTagValue("catalognumber") } -func (t Tags) Bpm() int { return (int)(math.Round(t.getFloat("tbpm", "bpm", "fbpm"))) } -func (t Tags) HasPicture() bool { return t.getFirstTagValue("has_picture") != "" } - -// MusicBrainz Identifiers - -func (t Tags) MbzReleaseTrackID() string { - return t.getMbzID("musicbrainz_releasetrackid", "musicbrainz release track id") -} - -func (t Tags) MbzRecordingID() string { - return t.getMbzID("musicbrainz_trackid", "musicbrainz track id") -} -func (t Tags) MbzAlbumID() string { return t.getMbzID("musicbrainz_albumid", "musicbrainz album id") } -func (t Tags) MbzArtistID() string { - return t.getMbzID("musicbrainz_artistid", "musicbrainz artist id") -} -func (t Tags) MbzAlbumArtistID() string { - return t.getMbzID("musicbrainz_albumartistid", "musicbrainz album artist id") -} -func (t Tags) MbzAlbumType() string { - return t.getFirstTagValue("musicbrainz_albumtype", "musicbrainz album type") -} -func (t Tags) MbzAlbumComment() string { - return t.getFirstTagValue("musicbrainz_albumcomment", "musicbrainz album comment") -} - -// Gain Properties - -func (t Tags) RGAlbumGain() float64 { - return t.getGainValue("replaygain_album_gain", "r128_album_gain") -} -func (t Tags) RGAlbumPeak() float64 { return t.getPeakValue("replaygain_album_peak") } -func (t Tags) RGTrackGain() float64 { - return t.getGainValue("replaygain_track_gain", "r128_track_gain") -} -func (t Tags) RGTrackPeak() float64 { return t.getPeakValue("replaygain_track_peak") } - -// File properties - -func (t Tags) Duration() float32 { return float32(t.getFloat("duration")) } -func (t Tags) SampleRate() int { return t.getInt("samplerate") } -func (t Tags) BitRate() int { return t.getInt("bitrate") } -func (t Tags) Channels() int { return t.getInt("channels") } -func (t Tags) ModificationTime() time.Time { return t.fileInfo.ModTime() } -func (t Tags) Size() int64 { return t.fileInfo.Size() } -func (t Tags) FilePath() string { return t.filePath } -func (t Tags) Suffix() string { return strings.ToLower(strings.TrimPrefix(path.Ext(t.filePath), ".")) } -func (t Tags) BirthTime() time.Time { - if ts := times.Get(t.fileInfo); ts.HasBirthTime() { - return ts.BirthTime() - } - return time.Now() -} - -func (t Tags) Lyrics() string { - lyricList := model.LyricList{} - basicLyrics := t.getAllTagValues("lyrics", "unsynced_lyrics", "unsynced lyrics", "unsyncedlyrics") - - for _, value := range basicLyrics { - parsed, err := model.ParseLyrics(context.Background(), ".lrc", "xxx", []byte(value)) - if err != nil { - log.Warn("Unexpected failure occurred when parsing lyrics", "file", t.filePath, "error", err) - continue - } - if main, ok := parsed.Main(); ok { - lyricList = append(lyricList, main) - } - } - - for tag, value := range t.Tags { - if after, ok := strings.CutPrefix(tag, "lyrics-"); ok { - language := strings.TrimSpace(after) - - if language == "" { - language = "xxx" - } - - for _, text := range value { - parsed, err := model.ParseLyrics(context.Background(), ".lrc", language, []byte(text)) - if err != nil { - log.Warn("Unexpected failure occurred when parsing lyrics", "file", t.filePath, "error", err) - continue - } - if main, ok := parsed.Main(); ok { - lyricList = append(lyricList, main) - } - } - } - } - - res, err := json.Marshal(lyricList) - if err != nil { - log.Warn("Unexpected error occurred when serializing lyrics", "file", t.filePath, "error", err) - return "" - } - return string(res) -} - -func (t Tags) getGainValue(rgTagName, r128TagName string) float64 { - // Check for ReplayGain first - // ReplayGain is in the form [-]a.bb dB and normalized to -18dB - var tag = t.getFirstTagValue(rgTagName) - if tag != "" { - tag = strings.TrimSpace(strings.Replace(tag, "dB", "", 1)) - var value, err = strconv.ParseFloat(tag, 64) - if err != nil || value == math.Inf(-1) || value == math.Inf(1) { - return 0 - } - return value - } - - // If ReplayGain is not found, check for R128 gain - // R128 gain is a Q7.8 fixed point number normalized to -23dB - tag = t.getFirstTagValue(r128TagName) - if tag != "" { - var iValue, err = strconv.Atoi(tag) - if err != nil { - return 0 - } - // Convert Q7.8 to float - var value = float64(iValue) / 256.0 - // Adding 5 dB to normalize with ReplayGain level - return value + 5 - } - - return 0 -} - -func (t Tags) getPeakValue(tagName string) float64 { - var tag = t.getFirstTagValue(tagName) - var value, err = strconv.ParseFloat(tag, 64) - if err != nil || value == math.Inf(-1) || value == math.Inf(1) { - // A default of 1 for peak value results in no changes - return 1 - } - return value -} - -func (t Tags) getTags(tagNames ...string) []string { - for _, tag := range tagNames { - if v, ok := t.Tags[tag]; ok { - return v - } - } - return nil -} - -func (t Tags) getFirstTagValue(tagNames ...string) string { - ts := t.getTags(tagNames...) - if len(ts) > 0 { - return ts[0] - } - return "" -} - -func (t Tags) getAllTagValues(tagNames ...string) []string { - values := make([]string, 0, len(tagNames)*2) - for _, tag := range tagNames { - if v, ok := t.Tags[tag]; ok { - values = append(values, v...) - } - } - return values -} - -func (t Tags) getSortTag(originalTag string, tagNames ...string) string { - formats := []string{"sort%s", "sort_%s", "sort-%s", "%ssort", "%s_sort", "%s-sort"} - all := make([]string, 1, len(tagNames)*len(formats)+1) - all[0] = originalTag - for _, tag := range tagNames { - for _, format := range formats { - name := fmt.Sprintf(format, tag) - all = append(all, name) - } - } - return t.getFirstTagValue(all...) -} - -var dateRegex = regexp.MustCompile(`([12]\d\d\d)`) - -func (t Tags) getDate(tagNames ...string) (int, string) { - tag := t.getFirstTagValue(tagNames...) - if len(tag) < 4 { - return 0, "" - } - // first get just the year - match := dateRegex.FindStringSubmatch(tag) - if len(match) == 0 { - log.Warn("Error parsing "+tagNames[0]+" field for year", "file", t.filePath, "date", tag) - return 0, "" - } - year, _ := strconv.Atoi(match[1]) - - if len(tag) < 5 { - return year, match[1] - } - - //then try YYYY-MM-DD - if len(tag) > 10 { - tag = tag[:10] - } - layout := "2006-01-02" - _, err := time.Parse(layout, tag) - if err != nil { - layout = "2006-01" - _, err = time.Parse(layout, tag) - if err != nil { - log.Warn("Error parsing "+tagNames[0]+" field for month + day", "file", t.filePath, "date", tag) - return year, match[1] - } - } - return year, tag -} - -func (t Tags) getBool(tagNames ...string) bool { - tag := t.getFirstTagValue(tagNames...) - if tag == "" { - return false - } - i, _ := strconv.Atoi(strings.TrimSpace(tag)) - return i == 1 -} - -func (t Tags) getTuple(tagNames ...string) (int, int) { - tag := t.getFirstTagValue(tagNames...) - if tag == "" { - return 0, 0 - } - tuple := strings.Split(tag, "/") - t1, t2 := 0, 0 - t1, _ = strconv.Atoi(tuple[0]) - if len(tuple) > 1 { - t2, _ = strconv.Atoi(tuple[1]) - } else { - t2tag := t.getFirstTagValue(tagNames[0] + "total") - t2, _ = strconv.Atoi(t2tag) - } - return t1, t2 -} - -func (t Tags) getMbzID(tagNames ...string) string { - tag := t.getFirstTagValue(tagNames...) - if _, err := uuid.Parse(tag); err != nil { - return "" - } - return tag -} - -func (t Tags) getInt(tagNames ...string) int { - tag := t.getFirstTagValue(tagNames...) - i, _ := strconv.Atoi(tag) - return i -} - -func (t Tags) getFloat(tagNames ...string) float64 { - var tag = t.getFirstTagValue(tagNames...) - var value, err = strconv.ParseFloat(tag, 64) - if err != nil { - return 0 - } - return value -} diff --git a/scanner/metadata_old/metadata_internal_test.go b/scanner/metadata_old/metadata_internal_test.go deleted file mode 100644 index aff1ede9c..000000000 --- a/scanner/metadata_old/metadata_internal_test.go +++ /dev/null @@ -1,144 +0,0 @@ -package metadata_old - -import ( - . "github.com/onsi/ginkgo/v2" - . "github.com/onsi/gomega" -) - -var _ = Describe("Tags", func() { - DescribeTable("getDate", - func(tag string, expectedYear int, expectedDate string) { - md := &Tags{} - md.Tags = map[string][]string{"date": {tag}} - testYear, testDate := md.Date() - Expect(testYear).To(Equal(expectedYear)) - Expect(testDate).To(Equal(expectedDate)) - }, - Entry(nil, "1985", 1985, "1985"), - Entry(nil, "2002-01", 2002, "2002-01"), - Entry(nil, "1969.06", 1969, "1969"), - Entry(nil, "1980.07.25", 1980, "1980"), - Entry(nil, "2004-00-00", 2004, "2004"), - Entry(nil, "2016-12-31", 2016, "2016-12-31"), - Entry(nil, "2013-May-12", 2013, "2013"), - Entry(nil, "May 12, 2016", 2016, "2016"), - Entry(nil, "01/10/1990", 1990, "1990"), - Entry(nil, "invalid", 0, ""), - ) - - Describe("getMbzID", func() { - It("return a valid MBID", func() { - md := &Tags{} - md.Tags = map[string][]string{ - "musicbrainz_trackid": {"8f84da07-09a0-477b-b216-cc982dabcde1"}, - "musicbrainz_releasetrackid": {"6caf16d3-0b20-3fe6-8020-52e31831bc11"}, - "musicbrainz_albumid": {"f68c985d-f18b-4f4a-b7f0-87837cf3fbf9"}, - "musicbrainz_artistid": {"89ad4ac3-39f7-470e-963a-56509c546377"}, - "musicbrainz_albumartistid": {"ada7a83c-e3e1-40f1-93f9-3e73dbc9298a"}, - } - Expect(md.MbzRecordingID()).To(Equal("8f84da07-09a0-477b-b216-cc982dabcde1")) - Expect(md.MbzReleaseTrackID()).To(Equal("6caf16d3-0b20-3fe6-8020-52e31831bc11")) - Expect(md.MbzAlbumID()).To(Equal("f68c985d-f18b-4f4a-b7f0-87837cf3fbf9")) - Expect(md.MbzArtistID()).To(Equal("89ad4ac3-39f7-470e-963a-56509c546377")) - Expect(md.MbzAlbumArtistID()).To(Equal("ada7a83c-e3e1-40f1-93f9-3e73dbc9298a")) - }) - It("return empty string for invalid MBID", func() { - md := &Tags{} - md.Tags = map[string][]string{ - "musicbrainz_trackid": {"11406732-6"}, - "musicbrainz_albumid": {"11406732"}, - "musicbrainz_artistid": {"200455"}, - "musicbrainz_albumartistid": {"194"}, - } - Expect(md.MbzRecordingID()).To(Equal("")) - Expect(md.MbzAlbumID()).To(Equal("")) - Expect(md.MbzArtistID()).To(Equal("")) - Expect(md.MbzAlbumArtistID()).To(Equal("")) - }) - }) - - Describe("getAllTagValues", func() { - It("returns values from all tag names", func() { - md := &Tags{} - md.Tags = map[string][]string{ - "genre": {"Rock", "Pop", "New Wave"}, - } - - Expect(md.Genres()).To(ConsistOf("Rock", "Pop", "New Wave")) - }) - }) - - Describe("removeDuplicatesAndEmpty", func() { - It("removes duplicates", func() { - md := NewTag("/music/artist/album01/Song.mp3", nil, ParsedTags{ - "genre": []string{"pop", "rock", "pop"}, - "date": []string{"2023-03-01", "2023-03-01"}, - "mood": []string{"happy", "sad"}, - }) - Expect(md.Tags).To(HaveKeyWithValue("genre", []string{"pop", "rock"})) - Expect(md.Tags).To(HaveKeyWithValue("date", []string{"2023-03-01"})) - Expect(md.Tags).To(HaveKeyWithValue("mood", []string{"happy", "sad"})) - }) - It("removes empty tags", func() { - md := NewTag("/music/artist/album01/Song.mp3", nil, ParsedTags{ - "genre": []string{"pop", "rock", "pop"}, - "mood": []string{"", ""}, - }) - Expect(md.Tags).To(HaveKeyWithValue("genre", []string{"pop", "rock"})) - Expect(md.Tags).ToNot(HaveKey("mood")) - }) - }) - - Describe("BPM", func() { - var t *Tags - BeforeEach(func() { - t = &Tags{Tags: map[string][]string{ - "fbpm": {"141.7"}, - }} - }) - - It("rounds a floating point fBPM tag", func() { - Expect(t.Bpm()).To(Equal(142)) - }) - }) - - Describe("ReplayGain", func() { - DescribeTable("getGainValue", - func(tag string, expected float64) { - md := &Tags{} - md.Tags = map[string][]string{"replaygain_track_gain": {tag}} - Expect(md.RGTrackGain()).To(Equal(expected)) - - }, - Entry("0", "0", 0.0), - Entry("1.2dB", "1.2dB", 1.2), - Entry("Infinity", "Infinity", 0.0), - Entry("Invalid value", "INVALID VALUE", 0.0), - ) - DescribeTable("getPeakValue", - func(tag string, expected float64) { - md := &Tags{} - md.Tags = map[string][]string{"replaygain_track_peak": {tag}} - Expect(md.RGTrackPeak()).To(Equal(expected)) - - }, - Entry("0", "0", 0.0), - Entry("0.5", "0.5", 0.5), - Entry("Invalid dB suffix", "0.7dB", 1.0), - Entry("Infinity", "Infinity", 1.0), - Entry("Invalid value", "INVALID VALUE", 1.0), - ) - DescribeTable("getR128GainValue", - func(tag string, expected float64) { - md := &Tags{} - md.Tags = map[string][]string{"r128_track_gain": {tag}} - Expect(md.RGTrackGain()).To(Equal(expected)) - - }, - Entry("0", "0", 5.0), - Entry("-3776", "-3776", -9.75), - Entry("Infinity", "Infinity", 0.0), - Entry("Invalid value", "INVALID VALUE", 0.0), - ) - }) -}) diff --git a/scanner/metadata_old/metadata_suite_test.go b/scanner/metadata_old/metadata_suite_test.go deleted file mode 100644 index 03ec3c847..000000000 --- a/scanner/metadata_old/metadata_suite_test.go +++ /dev/null @@ -1,17 +0,0 @@ -package metadata_old - -import ( - "testing" - - "github.com/navidrome/navidrome/log" - "github.com/navidrome/navidrome/tests" - . "github.com/onsi/ginkgo/v2" - . "github.com/onsi/gomega" -) - -func TestMetadata(t *testing.T) { - tests.Init(t, true) - log.SetLevel(log.LevelFatal) - RegisterFailHandler(Fail) - RunSpecs(t, "Metadata Suite") -} diff --git a/scanner/metadata_old/metadata_test.go b/scanner/metadata_old/metadata_test.go deleted file mode 100644 index 444bb7fc4..000000000 --- a/scanner/metadata_old/metadata_test.go +++ /dev/null @@ -1,95 +0,0 @@ -package metadata_old_test - -import ( - "cmp" - "encoding/json" - "slices" - - "github.com/navidrome/navidrome/conf" - "github.com/navidrome/navidrome/conf/configtest" - "github.com/navidrome/navidrome/core/ffmpeg" - "github.com/navidrome/navidrome/model" - "github.com/navidrome/navidrome/scanner/metadata_old" - _ "github.com/navidrome/navidrome/scanner/metadata_old/ffmpeg" - . "github.com/onsi/ginkgo/v2" - . "github.com/onsi/gomega" -) - -var _ = Describe("Tags", func() { - var zero int64 = 0 - var secondTs int64 = 2500 - - makeLyrics := func(synced bool, lang, secondLine string) model.Lyrics { - lines := []model.Line{ - {Value: "This is"}, - {Value: secondLine}, - } - - if synced { - lines[0].Start = &zero - lines[1].Start = &secondTs - } - - lyrics := model.Lyrics{ - Lang: lang, - Line: lines, - Synced: synced, - } - - return lyrics - } - - sortLyrics := func(lines model.LyricList) model.LyricList { - slices.SortFunc(lines, func(a, b model.Lyrics) int { - langDiff := cmp.Compare(a.Lang, b.Lang) - if langDiff != 0 { - return langDiff - } - return cmp.Compare(a.Line[1].Value, b.Line[1].Value) - }) - - return lines - } - - compareLyrics := func(m metadata_old.Tags, expected model.LyricList) { - lyrics := model.LyricList{} - Expect(json.Unmarshal([]byte(m.Lyrics()), &lyrics)).To(BeNil()) - Expect(sortLyrics(lyrics)).To(Equal(sortLyrics(expected))) - } - - // Only run these tests if FFmpeg is available - FFmpegContext := XContext - if ffmpeg.New().IsAvailable() { - FFmpegContext = Context - } - FFmpegContext("Extract with FFmpeg", func() { - BeforeEach(func() { - DeferCleanup(configtest.SetupConfig()) - conf.Server.Scanner.Extractor = "ffmpeg" - }) - - DescribeTable("Lyrics test", - func(file string) { - path := "tests/fixtures/" + file - mds, err := metadata_old.Extract(path) - Expect(err).ToNot(HaveOccurred()) - Expect(mds).To(HaveLen(1)) - - m := mds[path] - compareLyrics(m, model.LyricList{ - makeLyrics(true, "eng", "English"), - makeLyrics(true, "xxx", "unspecified"), - }) - }, - - Entry("Parses AIFF file", "test.aiff"), - Entry("Parses MP3 files", "test.mp3"), - // Disabled, because it fails in pipeline - // Entry("Parses WAV files", "test.wav"), - - // FFMPEG behaves very weirdly for multivalued tags for non-ID3 - // Specifically, they are separated by ";, which is indistinguishable - // from other fields - ) - }) -}) diff --git a/server/initial_setup.go b/server/initial_setup.go index be9e14ae9..462e22e54 100644 --- a/server/initial_setup.go +++ b/server/initial_setup.go @@ -72,10 +72,6 @@ func checkFFmpegInstallation() { _, err := f.CmdPath() if err != nil { log.Warn("Unable to find ffmpeg. Transcoding will fail if used", err) - if conf.Server.Scanner.Extractor == "ffmpeg" { - log.Warn("ffmpeg cannot be used for metadata extraction. Falling back to taglib") - conf.Server.Scanner.Extractor = "taglib" - } return } if !f.IsProbeAvailable() { diff --git a/tests/harness/harness.go b/tests/harness/harness.go index 92196ef26..67f3150a2 100644 --- a/tests/harness/harness.go +++ b/tests/harness/harness.go @@ -176,8 +176,6 @@ func (NoopFFmpeg) ExtractImage(context.Context, string) (io.ReadCloser, error) { return nil, errors.New("noop ffmpeg: extract image not supported") } -func (NoopFFmpeg) Probe(context.Context, []string) (string, error) { return "", nil } - func (NoopFFmpeg) ProbeAudioStream(context.Context, string) (*ffmpeg.AudioProbeResult, error) { return nil, errors.New("noop ffmpeg: probe not supported") } diff --git a/tests/mock_ffmpeg.go b/tests/mock_ffmpeg.go index f9862767e..8e4d12f0b 100644 --- a/tests/mock_ffmpeg.go +++ b/tests/mock_ffmpeg.go @@ -57,12 +57,6 @@ func (ff *MockFFmpeg) ConvertAnimatedImage(_ context.Context, reader io.Reader, return io.NopCloser(bytes.NewReader(data)), nil } -func (ff *MockFFmpeg) Probe(context.Context, []string) (string, error) { - if ff.Error != nil { - return "", ff.Error - } - return "", nil -} func (ff *MockFFmpeg) ProbeAudioStream(context.Context, string) (*ffmpeg.AudioProbeResult, error) { if ff.Error != nil { return nil, ff.Error From c22ce9ebb27c5e80b104c28a82ac70b018f07895 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sat, 26 Sep 2026 15:27:23 -0400 Subject: [PATCH 17/26] feat(api): add the API v1 foundation behind DevAPIv1 (#6227) * feat(api): add OpenAPI v1 spec skeleton, lint ruleset and bundle tooling vacuum v0.30.6's `bundle --composed` mangles component names for this spec's multi-file layout (duplicates Problem as Problem__schemas etc.), so api-bundle uses the Redocly CLI (npx @redocly/cli bundle) instead. * fix(api): pin the Redocly CLI version Tried moving components out of the root document (per libopenapi's nested_files example) so vacuum's own bundler could produce clean names, but any component declared via $ref inside components.* still gets a __-suffixed twin regardless of collisions elsewhere, so vacuum's --composed bundler can't cleanly bundle this spec. Pin the already-working Redocly fallback to an exact version instead of @latest. * fix(api): bundle the OpenAPI spec with vacuum vacuum's --composed bundler suffixes any component reached via a $ref written directly inside the root document's own components.* block, regardless of collisions elsewhere. Dropping the root-level schemas/ parameters/responses declarations (keeping only securitySchemes, and leaving every component file under api/openapi/components/ untouched) lets vacuum bundle cleanly with no __ suffixes, going back to Go-only tooling. Components nothing references yet (ListMeta, offset, limit, BadRequest, Unauthorized, Forbidden, NotFound) are absent from the bundle until a later task's operation references them. * fix(api): make spec lint rules cover all schemas and error codes nd-schema-property-descriptions targeted $.components.schemas, but our schemas live in path/response files, not the root document, so it was dead code; switched to $..properties[*] to walk every resolved schema wherever it ends up. nd-error-responses-are-problems only checked a hardcoded status-code list; switched to a patternProperties schema matching the full 4xx/5xx range. Also: api-diff now diffs against the merge-base with API_DIFF_BASE (falling back to its tip with a notice if no merge-base exists), gen no longer depends on api-gen until Task 3 wires up oapi-codegen, and api-lint suppresses vacuum's banner. * feat(api): embed the bundled OpenAPI spec and expose its version * feat(api): generate the v1 server interface with oapi-codegen * feat(api): add RFC 9457 problem responses for API v1 * feat(api): add API v1 router with /server discovery and spec routes * fix(api): serve the OpenAPI document without range support * feat(api): mount API v1 behind the DevAPIv1 flag * chore(ci): lint, regenerate and diff the OpenAPI v1 spec * refactor(api): tighten spec version access, lint rules and test naming * refactor(api): simplify spec routes, tests and OpenAPI tooling Share one If-None-Match parser (utils/req) between the image and spec routes, declare the YAML spec response as an object so tests need no decoder override, and reuse ETag/304 spec components. Install the OpenAPI tools only when missing or at a different version, fail api-diff when its base ref does not exist, and in CI cache the tools, fold regeneration into the go generate check, and fetch only the PR base commit for the breaking-change gate. * refactor(api): raise the list limit maximum to 2000 and drop the flag test * feat(api): treat added enum values as non-breaking Enums in API v1 are open: clients must accept unknown values. api-diff now downgrades response-property-enum-value-added to INFO, while removing a value from a request enum stays breaking. * feat(api): gate breaking changes on x-stability-level Every operation declares x-stability-level (alpha, beta, stable). oasdiff ignores breaking changes to alpha operations and rejects lowering a level, so unreleased endpoints can evolve while beta and stable ones stay additive. All current operations start as alpha. * feat(api): declare loginMethods as an enum Prefix generated enum constants with their type name so enums sharing a value (for example password) cannot collide in package apiv1. * feat(api): send Allow on 405 and answer HEAD wherever GET is routed chi only sets Allow in its default 405 handler, so the problem-format handler now builds it by matching each method against the v1 router. HEAD requests fall back to the GET route, as RFC 9110 expects. * refactor(api): hash the spec ETag with xxh3 The bytes are compiled in, and the digest was truncated to 64 bits anyway, so this matches the artwork ETags instead of paying for cryptographic strength we discard. * docs(api): explain the about:blank problem type * feat(api): make code the problem identifier and omit a blank type RFC 9457 says clients switch on the type URI, but no adopter surveyed ships both a populated type and a separate code. Declare code as an enum, and send type only once a problem has semantics of its own. * fix(api): advertise the configured base path in the served OpenAPI spec With BaseURL=/music the API is mounted at /music/api/v1, but the spec told clients to call /api/v1 at the host root. The server now rewrites servers[0].url to BasePath + /api/v1 when it serves the document. Relative server URLs were tested first: "." and "../v1" work in openapi-generator, Swagger UI and Redoc, but Scalar resolves them against the page origin, so it breaks even without a base path. The committed bundle keeps /api/v1, and a test pins that it appears exactly once, which the rewrite relies on. --- .github/workflows/pipeline.yml | 23 +- Makefile | 45 +- api/.oasdiff-levels.txt | 1 + api/.vacuum.yaml | 155 +++++++ api/api_suite_test.go | 17 + api/bundled/openapi.json | 262 ++++++++++++ api/bundled/openapi.yaml | 194 +++++++++ api/embed.go | 35 ++ api/embed_test.go | 37 ++ api/openapi/components/headers/ETag.yaml | 3 + api/openapi/components/parameters/limit.yaml | 9 + api/openapi/components/parameters/offset.yaml | 8 + .../components/responses/BadRequest.yaml | 5 + .../components/responses/Forbidden.yaml | 5 + .../components/responses/InternalError.yaml | 5 + .../components/responses/NotFound.yaml | 5 + .../components/responses/NotModified.yaml | 4 + .../components/responses/Unauthorized.yaml | 5 + api/openapi/components/schemas/ListMeta.yaml | 13 + api/openapi/components/schemas/Problem.yaml | 35 ++ .../components/schemas/ServerInfo.yaml | 22 + .../components/schemas/ValidationError.yaml | 10 + api/openapi/openapi.yaml | 39 ++ api/openapi/paths/openapi.yaml | 42 ++ api/openapi/paths/server.yaml | 19 + cmd/root.go | 3 + cmd/wire_gen.go | 10 +- cmd/wire_injectors.go | 8 + conf/configuration.go | 2 + consts/consts.go | 1 + go.mod | 6 + go.sum | 18 +- server/apiv1/api.go | 98 +++++ server/apiv1/api_gen.go | 390 ++++++++++++++++++ server/apiv1/api_test.go | 75 ++++ server/apiv1/apiv1_suite_test.go | 66 +++ server/apiv1/oapi-codegen.yaml | 12 + server/apiv1/problem.go | 72 ++++ server/apiv1/problem_test.go | 114 +++++ server/apiv1/server_info.go | 23 ++ server/apiv1/server_info_test.go | 61 +++ server/apiv1/spec.go | 51 +++ server/apiv1/spec_test.go | 136 ++++++ server/imghttp/headers.go | 23 +- utils/req/req.go | 18 + utils/req/req_test.go | 18 + 46 files changed, 2177 insertions(+), 26 deletions(-) create mode 100644 api/.oasdiff-levels.txt create mode 100644 api/.vacuum.yaml create mode 100644 api/api_suite_test.go create mode 100644 api/bundled/openapi.json create mode 100644 api/bundled/openapi.yaml create mode 100644 api/embed.go create mode 100644 api/embed_test.go create mode 100644 api/openapi/components/headers/ETag.yaml create mode 100644 api/openapi/components/parameters/limit.yaml create mode 100644 api/openapi/components/parameters/offset.yaml create mode 100644 api/openapi/components/responses/BadRequest.yaml create mode 100644 api/openapi/components/responses/Forbidden.yaml create mode 100644 api/openapi/components/responses/InternalError.yaml create mode 100644 api/openapi/components/responses/NotFound.yaml create mode 100644 api/openapi/components/responses/NotModified.yaml create mode 100644 api/openapi/components/responses/Unauthorized.yaml create mode 100644 api/openapi/components/schemas/ListMeta.yaml create mode 100644 api/openapi/components/schemas/Problem.yaml create mode 100644 api/openapi/components/schemas/ServerInfo.yaml create mode 100644 api/openapi/components/schemas/ValidationError.yaml create mode 100644 api/openapi/openapi.yaml create mode 100644 api/openapi/paths/openapi.yaml create mode 100644 api/openapi/paths/server.yaml create mode 100644 server/apiv1/api.go create mode 100644 server/apiv1/api_gen.go create mode 100644 server/apiv1/api_test.go create mode 100644 server/apiv1/apiv1_suite_test.go create mode 100644 server/apiv1/oapi-codegen.yaml create mode 100644 server/apiv1/problem.go create mode 100644 server/apiv1/problem_test.go create mode 100644 server/apiv1/server_info.go create mode 100644 server/apiv1/server_info_test.go create mode 100644 server/apiv1/spec.go create mode 100644 server/apiv1/spec_test.go diff --git a/.github/workflows/pipeline.yml b/.github/workflows/pipeline.yml index 2702163d3..c1ba8714c 100644 --- a/.github/workflows/pipeline.yml +++ b/.github/workflows/pipeline.yml @@ -92,8 +92,23 @@ jobs: exit 1 fi + - name: Resolve OpenAPI tool versions + id: api-tools + run: echo "key=$(grep -E '^(VACUUM|OAPI_CODEGEN|OASDIFF)_VERSION' Makefile | tr -d ' \n')" >> "$GITHUB_OUTPUT" + + - name: Cache OpenAPI tools + uses: actions/cache@v6 + with: + path: bin + key: api-tools-${{ runner.os }}-${{ steps.api-tools.outputs.key }} + + - name: Lint OpenAPI spec + run: make api-lint + - name: Run go generate - run: go generate ./... + run: | + make api-gen + go generate ./... - name: Verify no changes from go generate run: | git status --porcelain @@ -102,6 +117,12 @@ jobs: exit 1 fi + - name: Check for breaking OpenAPI changes + if: github.event_name == 'pull_request' + run: | + git fetch --no-tags --depth=1 origin ${{ github.event.pull_request.base.sha }} + make api-diff API_DIFF_BASE=${{ github.event.pull_request.base.sha }} + validate-migrations: name: Validate DB migrations runs-on: ubuntu-latest diff --git a/Makefile b/Makefile index eccccbffb..f31d98a01 100644 --- a/Makefile +++ b/Makefile @@ -21,6 +21,10 @@ PLATFORMS ?= $(SUPPORTED_PLATFORMS) DOCKER_TAG ?= deluan/navidrome:develop GOLANGCI_LINT_VERSION ?= v2.14.0 +VACUUM_VERSION ?= v0.30.6 +OAPI_CODEGEN_VERSION ?= v2.8.0 +OASDIFF_VERSION ?= v1.32.1 +API_DIFF_BASE ?= origin/master UI_SRC_FILES := $(shell find ui -type f -not -path "ui/build/*" -not -path "ui/node_modules/*") @@ -92,6 +96,45 @@ install-golangci-lint: ##@Development Install golangci-lint if not present fi .PHONY: install-golangci-lint +install-api-tools: ##@Development Install OpenAPI tools (vacuum, oapi-codegen, oasdiff) into ./bin + @STAMP=bin/.api-tools-$(VACUUM_VERSION)-$(OAPI_CODEGEN_VERSION)-$(OASDIFF_VERSION); \ + if [ ! -f $$STAMP ] || [ ! -x bin/vacuum ] || [ ! -x bin/oapi-codegen ] || [ ! -x bin/oasdiff ]; then \ + echo "Installing OpenAPI tools..."; \ + GOBIN=$(CURDIR)/bin go install github.com/daveshanley/vacuum@$(VACUUM_VERSION) && \ + GOBIN=$(CURDIR)/bin go install github.com/oapi-codegen/oapi-codegen/v2/cmd/oapi-codegen@$(OAPI_CODEGEN_VERSION) && \ + GOBIN=$(CURDIR)/bin go install github.com/oasdiff/oasdiff@$(OASDIFF_VERSION) && \ + rm -f bin/.api-tools-* && touch $$STAMP; \ + fi +.PHONY: install-api-tools + +api-lint: install-api-tools ##@Development Lint the OpenAPI spec + ./bin/vacuum lint -r api/.vacuum.yaml -d -q -b --fail-severity error api/openapi/openapi.yaml +.PHONY: api-lint + +api-bundle: install-api-tools ##@Development Bundle the multi-file OpenAPI spec into api/bundled + ./bin/vacuum bundle -q --composed -p api/openapi api/openapi/openapi.yaml api/bundled/openapi.yaml + ./bin/vacuum bundle -q --composed --format json -p api/openapi api/openapi/openapi.yaml api/bundled/openapi.json +.PHONY: api-bundle + +api-gen: api-bundle ##@Development Generate the API v1 server code from the bundled spec + ./bin/oapi-codegen -config server/apiv1/oapi-codegen.yaml api/bundled/openapi.json +.PHONY: api-gen + +api-diff: api-bundle ##@Development Fail on breaking OpenAPI changes against the merge-base with $(API_DIFF_BASE) + @git rev-parse --verify --quiet $(API_DIFF_BASE)^{commit} >/dev/null || { echo "Base ref $(API_DIFF_BASE) not found; set API_DIFF_BASE"; exit 1; }; \ + BASE="$$(git merge-base HEAD $(API_DIFF_BASE) 2>/dev/null)"; \ + if [ -z "$$BASE" ]; then \ + echo "No merge-base with $(API_DIFF_BASE); falling back to its tip"; \ + BASE=$(API_DIFF_BASE); \ + fi; \ + if git cat-file -e $$BASE:api/bundled/openapi.json 2>/dev/null; then \ + git show $$BASE:api/bundled/openapi.json > $(CURDIR)/bin/api-base.json && \ + ./bin/oasdiff breaking $(CURDIR)/bin/api-base.json api/bundled/openapi.json --fail-on ERR --severity-levels api/.oasdiff-levels.txt; \ + else \ + echo "No bundled spec at $$BASE; skipping breaking-change check"; \ + fi +.PHONY: api-diff + lint: install-golangci-lint ##@Development Lint Go code PATH=./bin:$$PATH golangci-lint run --timeout 5m .PHONY: lint @@ -111,7 +154,7 @@ wire: check_go_env ##@Development Update Dependency Injection go tool wire gen -tags="$$(echo '$(GO_BUILD_TAGS)' | tr ',' ' ')" ./... .PHONY: wire -gen: check_go_env ##@Development Run go generate for code generation +gen: check_go_env api-gen ##@Development Run go generate for code generation go generate ./... cd plugins/cmd/ndpgen && go run . -shared-types -input=../../types -output=../../pdk -go -rust cd plugins/cmd/ndpgen && go run . -host-wrappers -input=../../host -package=host -shared=../../types diff --git a/api/.oasdiff-levels.txt b/api/.oasdiff-levels.txt new file mode 100644 index 000000000..1792adcb2 --- /dev/null +++ b/api/.oasdiff-levels.txt @@ -0,0 +1 @@ +response-property-enum-value-added INFO diff --git a/api/.vacuum.yaml b/api/.vacuum.yaml new file mode 100644 index 000000000..bd31e23f8 --- /dev/null +++ b/api/.vacuum.yaml @@ -0,0 +1,155 @@ +extends: [[spectral:oas, recommended]] +rules: + # vacuum's `enumeration` function mis-resolves hyphenated `then.field` names, + # so the value check below targets `x-module` via `given` instead. + nd-operation-x-module-required: + description: Every operation belongs to exactly one capability module. + severity: error + given: $.paths[*][get,put,post,delete,patch] + then: + field: x-module + function: truthy + nd-operation-x-module: + description: Every operation's capability module is one of the known values. + severity: error + given: $.paths[*][get,put,post,delete,patch]['x-module'] + then: + function: enumeration + functionOptions: + values: + - core + - streaming + - download + - artwork + - lyrics + - transcoding + - annotations + - playback + - queue + - custom-tags + - grouping + - playlists + - smart-playlists + - sync + - events + - jukebox + - sharing + - radio + - admin + nd-operation-stability-level-required: + description: Every operation declares its stability level, which the breaking-change gate relies on. + severity: error + given: $.paths[*][get,put,post,delete,patch] + then: + field: x-stability-level + function: truthy + nd-operation-stability-level: + description: Every operation's stability level is alpha, beta, or stable. + severity: error + given: $.paths[*][get,put,post,delete,patch]['x-stability-level'] + then: + function: enumeration + functionOptions: + values: + - alpha + - beta + - stable + nd-operation-required-fields: + description: Operations need a stable operationId, summary, description and tags. + severity: error + given: $.paths[*][get,put,post,delete,patch] + then: + - field: operationId + function: truthy + - field: summary + function: truthy + - field: description + function: truthy + - field: tags + function: truthy + # Our schemas live in path/response files, not root components, so this + # walks every resolved `properties` map in the document via `$..` instead. + nd-schema-property-descriptions: + description: Every schema property is documented. + severity: error + given: $..properties[*] + then: + field: description + function: truthy + # patternProperties covers the full 4xx/5xx range; needs an explicit + # `properties` entry too, or `additionalProperties: false` rejects it. + nd-error-responses-are-problems: + description: 4xx and 5xx responses use application/problem+json. + severity: error + given: $.paths[*][*].responses + then: + function: schema + functionOptions: + forceValidationOnCurrentNode: true + schema: + type: object + patternProperties: + "^[45][0-9][0-9]$": + type: object + required: [content] + properties: + content: + type: object + properties: + application/problem+json: {} + required: [application/problem+json] + additionalProperties: false + # Same filter limitation applies here: "is this a list endpoint" is expressed + # as a JSON Schema if/then on the operation object instead of a `given` filter. + nd-list-endpoints-paginate: + description: List endpoints declare the shared offset and limit parameters. + severity: error + given: $.paths[*].get + then: + function: schema + functionOptions: + forceValidationOnCurrentNode: true + schema: + type: object + if: + required: [responses] + properties: + responses: + type: object + required: ['200'] + properties: + '200': + type: object + required: [content] + properties: + content: + type: object + required: [application/json] + properties: + application/json: + type: object + required: [schema] + properties: + schema: + type: object + required: [properties] + properties: + properties: + type: object + required: [items] + then: + required: [parameters] + properties: + parameters: + type: array + allOf: + - contains: + type: object + properties: + name: + const: offset + - contains: + type: object + properties: + name: + const: limit diff --git a/api/api_suite_test.go b/api/api_suite_test.go new file mode 100644 index 000000000..62a547b7c --- /dev/null +++ b/api/api_suite_test.go @@ -0,0 +1,17 @@ +package api_test + +import ( + "testing" + + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestAPI(t *testing.T) { + tests.Init(t, false) + log.SetLevel(log.LevelFatal) + RegisterFailHandler(Fail) + RunSpecs(t, "API Spec Suite") +} diff --git a/api/bundled/openapi.json b/api/bundled/openapi.json new file mode 100644 index 000000000..bfe794c71 --- /dev/null +++ b/api/bundled/openapi.json @@ -0,0 +1,262 @@ +{ + "openapi": "3.0.3", + "info": { + "title": "Navidrome API", + "version": "1.0.0", + "description": "Navidrome API v1. Spec-first, additive within v1. Clients discover implemented\ncapability modules through `GET /server` and never sniff versions.\n\nEnums are open: new values may be added to any enum within v1. Clients must\naccept values they do not recognise instead of failing.\n\nEvery operation declares `x-stability-level`: `alpha` operations may change or\ndisappear without notice, `beta` and `stable` operations only change additively.\nA level is only ever raised, never lowered.\n\n`HEAD` is accepted wherever `GET` is. A `405` response lists the allowed methods\nin its `Allow` header.\n", + "license": { + "name": "GPL-3.0", + "url": "https://www.gnu.org/licenses/gpl-3.0.html" + } + }, + "servers": [ + { + "url": "/api/v1" + } + ], + "tags": [ + { + "name": "server", + "description": "Server discovery and the published OpenAPI document." + } + ], + "paths": { + "/server": { + "get": { + "operationId": "getServerInfo", + "x-module": "core", + "x-stability-level": "alpha", + "tags": [ + "server" + ], + "summary": "Describe the server", + "description": "Returns the public server description. No authentication required.\nAuthenticated requests will additionally receive the implemented capability modules\nonce authentication is available.\n", + "responses": { + "200": { + "description": "Server description.", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ServerInfo" + } + } + } + }, + "500": { + "$ref": "#/components/responses/InternalError" + } + } + } + }, + "/openapi.json": { + "get": { + "operationId": "getOpenAPISpecJSON", + "x-module": "core", + "x-stability-level": "alpha", + "tags": [ + "server" + ], + "summary": "Get the OpenAPI document (JSON)", + "description": "The bundled OpenAPI document of the running server version. Supports ETag revalidation.", + "responses": { + "200": { + "description": "The OpenAPI document.", + "headers": { + "ETag": { + "$ref": "#/components/headers/ETag" + } + }, + "content": { + "application/json": { + "schema": { + "type": "object", + "description": "OpenAPI 3.0 document." + } + } + } + }, + "304": { + "$ref": "#/components/responses/NotModified" + } + } + } + }, + "/openapi.yaml": { + "get": { + "operationId": "getOpenAPISpecYAML", + "x-module": "core", + "x-stability-level": "alpha", + "tags": [ + "server" + ], + "summary": "Get the OpenAPI document (YAML)", + "description": "The bundled OpenAPI document of the running server version. Supports ETag revalidation.", + "responses": { + "200": { + "description": "The OpenAPI document.", + "headers": { + "ETag": { + "$ref": "#/components/headers/ETag" + } + }, + "content": { + "application/yaml": { + "schema": { + "type": "object", + "description": "OpenAPI 3.0 document." + } + } + } + }, + "304": { + "$ref": "#/components/responses/NotModified" + } + } + } + } + }, + "components": { + "securitySchemes": { + "bearerAuth": { + "type": "http", + "scheme": "bearer", + "bearerFormat": "JWT", + "description": "Short-lived access token minted from a device grant. Not yet applied to any operation." + } + }, + "schemas": { + "ServerInfo": { + "type": "object", + "description": "Public server description. Everything an add-server screen needs before login.", + "required": [ + "name", + "serverVersion", + "specVersion", + "setupRequired", + "loginMethods" + ], + "properties": { + "name": { + "type": "string", + "description": "Human-readable server product name." + }, + "serverVersion": { + "type": "string", + "description": "Version of the running server build." + }, + "specVersion": { + "type": "string", + "description": "Version of the OpenAPI document this server implements." + }, + "setupRequired": { + "type": "boolean", + "description": "True until the first admin user has been created." + }, + "loginMethods": { + "type": "array", + "description": "Login methods this server accepts. New methods may be added; clients ignore values they do not recognise.", + "items": { + "type": "string", + "enum": [ + "password" + ] + } + } + } + }, + "Problem": { + "type": "object", + "description": "RFC 9457 problem details, returned for every 4xx and 5xx response.", + "required": [ + "title", + "status", + "code" + ], + "properties": { + "type": { + "type": "string", + "description": "URI reference identifying the problem type. Omitted while the problem carries no semantics\nbeyond its HTTP status code, which RFC 9457 defines as `about:blank`. Problems with their\nown semantics get their own URI; switch on `code` instead.\n" + }, + "title": { + "type": "string", + "description": "Short human-readable summary, the same for all occurrences of this problem type." + }, + "status": { + "type": "integer", + "description": "HTTP status code of this response." + }, + "detail": { + "type": "string", + "description": "Human-readable explanation specific to this occurrence. Omitted for internal errors." + }, + "code": { + "type": "string", + "description": "Machine-readable error code, and the value clients switch on. New codes may be added.", + "enum": [ + "validation", + "unauthorized", + "forbidden", + "not_found", + "method_not_allowed", + "unavailable", + "internal" + ] + }, + "errors": { + "type": "array", + "description": "Per-field failures. Present only when `code` is `validation`.", + "items": { + "$ref": "#/components/schemas/ValidationError" + } + } + } + }, + "ValidationError": { + "type": "object", + "description": "One field-level validation failure.", + "required": [ + "field", + "message" + ], + "properties": { + "field": { + "type": "string", + "description": "Name of the offending query parameter, path parameter, or body field (dotted for nested)." + }, + "message": { + "type": "string", + "description": "Why the value was rejected." + } + } + } + }, + "responses": { + "InternalError": { + "description": "Unexpected server failure. Details are in the server log.", + "content": { + "application/problem+json": { + "schema": { + "$ref": "#/components/schemas/Problem" + } + } + } + }, + "NotModified": { + "description": "Not modified.", + "headers": { + "ETag": { + "$ref": "#/components/headers/ETag" + } + } + } + }, + "headers": { + "ETag": { + "description": "Entity tag for `If-None-Match` revalidation.", + "schema": { + "type": "string" + } + } + } + } +} \ No newline at end of file diff --git a/api/bundled/openapi.yaml b/api/bundled/openapi.yaml new file mode 100644 index 000000000..f1ae95b76 --- /dev/null +++ b/api/bundled/openapi.yaml @@ -0,0 +1,194 @@ +openapi: 3.0.3 +info: + title: Navidrome API + version: 1.0.0 + description: | + Navidrome API v1. Spec-first, additive within v1. Clients discover implemented + capability modules through `GET /server` and never sniff versions. + + Enums are open: new values may be added to any enum within v1. Clients must + accept values they do not recognise instead of failing. + + Every operation declares `x-stability-level`: `alpha` operations may change or + disappear without notice, `beta` and `stable` operations only change additively. + A level is only ever raised, never lowered. + + `HEAD` is accepted wherever `GET` is. A `405` response lists the allowed methods + in its `Allow` header. + license: + name: GPL-3.0 + url: https://www.gnu.org/licenses/gpl-3.0.html +servers: + - url: /api/v1 +tags: + - name: server + description: Server discovery and the published OpenAPI document. +paths: + /server: + get: + operationId: getServerInfo + x-module: core + x-stability-level: alpha + tags: [server] + summary: Describe the server + description: | + Returns the public server description. No authentication required. + Authenticated requests will additionally receive the implemented capability modules + once authentication is available. + responses: + '200': + description: Server description. + content: + application/json: + schema: + $ref: '#/components/schemas/ServerInfo' + '500': + $ref: '#/components/responses/InternalError' + /openapi.json: + get: + operationId: getOpenAPISpecJSON + x-module: core + x-stability-level: alpha + tags: [server] + summary: Get the OpenAPI document (JSON) + description: The bundled OpenAPI document of the running server version. Supports ETag revalidation. + responses: + '200': + description: The OpenAPI document. + headers: + ETag: + $ref: '#/components/headers/ETag' + content: + application/json: + schema: + type: object + description: OpenAPI 3.0 document. + '304': + $ref: '#/components/responses/NotModified' + /openapi.yaml: + get: + operationId: getOpenAPISpecYAML + x-module: core + x-stability-level: alpha + tags: [server] + summary: Get the OpenAPI document (YAML) + description: The bundled OpenAPI document of the running server version. Supports ETag revalidation. + responses: + '200': + description: The OpenAPI document. + headers: + ETag: + $ref: '#/components/headers/ETag' + content: + application/yaml: + schema: + type: object + description: OpenAPI 3.0 document. + '304': + $ref: '#/components/responses/NotModified' +components: + securitySchemes: + bearerAuth: + type: http + scheme: bearer + bearerFormat: JWT + description: Short-lived access token minted from a device grant. Not yet applied to any operation. + schemas: + ServerInfo: + type: object + description: Public server description. Everything an add-server screen needs before login. + required: + - name + - serverVersion + - specVersion + - setupRequired + - loginMethods + properties: + name: + type: string + description: Human-readable server product name. + serverVersion: + type: string + description: Version of the running server build. + specVersion: + type: string + description: Version of the OpenAPI document this server implements. + setupRequired: + type: boolean + description: True until the first admin user has been created. + loginMethods: + type: array + description: Login methods this server accepts. New methods may be added; clients ignore values they do not recognise. + items: + type: string + enum: + - password + Problem: + type: object + description: RFC 9457 problem details, returned for every 4xx and 5xx response. + required: + - title + - status + - code + properties: + type: + type: string + description: | + URI reference identifying the problem type. Omitted while the problem carries no semantics + beyond its HTTP status code, which RFC 9457 defines as `about:blank`. Problems with their + own semantics get their own URI; switch on `code` instead. + title: + type: string + description: Short human-readable summary, the same for all occurrences of this problem type. + status: + type: integer + description: HTTP status code of this response. + detail: + type: string + description: Human-readable explanation specific to this occurrence. Omitted for internal errors. + code: + type: string + description: Machine-readable error code, and the value clients switch on. New codes may be added. + enum: + - validation + - unauthorized + - forbidden + - not_found + - method_not_allowed + - unavailable + - internal + errors: + type: array + description: Per-field failures. Present only when `code` is `validation`. + items: + $ref: '#/components/schemas/ValidationError' + ValidationError: + type: object + description: One field-level validation failure. + required: + - field + - message + properties: + field: + type: string + description: Name of the offending query parameter, path parameter, or body field (dotted for nested). + message: + type: string + description: Why the value was rejected. + responses: + InternalError: + description: Unexpected server failure. Details are in the server log. + content: + application/problem+json: + schema: + $ref: '#/components/schemas/Problem' + NotModified: + description: Not modified. + headers: + ETag: + $ref: '#/components/headers/ETag' + headers: + ETag: + description: Entity tag for `If-None-Match` revalidation. + schema: + type: string diff --git a/api/embed.go b/api/embed.go new file mode 100644 index 000000000..7e94c2362 --- /dev/null +++ b/api/embed.go @@ -0,0 +1,35 @@ +package api + +import ( + _ "embed" + "encoding/json" + "sync" +) + +//go:embed bundled/openapi.json +var specJSON []byte + +//go:embed bundled/openapi.yaml +var specYAML []byte + +func SpecJSON() []byte { + return specJSON +} + +func SpecYAML() []byte { + return specYAML +} + +var specVersion = sync.OnceValue(func() string { + var doc struct { + Info struct { + Version string `json:"version"` + } `json:"info"` + } + _ = json.Unmarshal(SpecJSON(), &doc) + return doc.Info.Version +}) + +func SpecVersion() string { + return specVersion() +} diff --git a/api/embed_test.go b/api/embed_test.go new file mode 100644 index 000000000..ebd5a7d1d --- /dev/null +++ b/api/embed_test.go @@ -0,0 +1,37 @@ +package api_test + +import ( + "os" + + "github.com/getkin/kin-openapi/openapi3" + "github.com/navidrome/navidrome/api" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gopkg.in/yaml.v3" +) + +var _ = Describe("Bundled spec", func() { + It("embeds a valid OpenAPI 3 document", func() { + doc, err := openapi3.NewLoader().LoadFromData(api.SpecJSON()) + Expect(err).ToNot(HaveOccurred()) + Expect(doc.Validate(GinkgoT().Context())).To(Succeed()) + Expect(doc.Paths.Find("/server")).ToNot(BeNil()) + }) + + It("embeds the YAML variant", func() { + var doc map[string]any + Expect(yaml.Unmarshal(api.SpecYAML(), &doc)).To(Succeed()) + Expect(doc).To(HaveKey("paths")) + }) + + It("reports the version from the bundle, matching the source root document", func() { + src, err := os.ReadFile("api/openapi/openapi.yaml") + Expect(err).ToNot(HaveOccurred()) + var root struct { + Info struct{ Version string } `yaml:"info"` + } + Expect(yaml.Unmarshal(src, &root)).To(Succeed()) + Expect(api.SpecVersion()).To(Equal(root.Info.Version)) + Expect(api.SpecVersion()).ToNot(BeEmpty()) + }) +}) diff --git a/api/openapi/components/headers/ETag.yaml b/api/openapi/components/headers/ETag.yaml new file mode 100644 index 000000000..0f64e6792 --- /dev/null +++ b/api/openapi/components/headers/ETag.yaml @@ -0,0 +1,3 @@ +description: Entity tag for `If-None-Match` revalidation. +schema: + type: string diff --git a/api/openapi/components/parameters/limit.yaml b/api/openapi/components/parameters/limit.yaml new file mode 100644 index 000000000..9bef4fc8b --- /dev/null +++ b/api/openapi/components/parameters/limit.yaml @@ -0,0 +1,9 @@ +name: limit +in: query +description: Maximum number of items to return. +required: false +schema: + type: integer + minimum: 1 + maximum: 2000 + default: 100 diff --git a/api/openapi/components/parameters/offset.yaml b/api/openapi/components/parameters/offset.yaml new file mode 100644 index 000000000..9145d6eaa --- /dev/null +++ b/api/openapi/components/parameters/offset.yaml @@ -0,0 +1,8 @@ +name: offset +in: query +description: Zero-based index of the first item to return. +required: false +schema: + type: integer + minimum: 0 + default: 0 diff --git a/api/openapi/components/responses/BadRequest.yaml b/api/openapi/components/responses/BadRequest.yaml new file mode 100644 index 000000000..1a6e0b657 --- /dev/null +++ b/api/openapi/components/responses/BadRequest.yaml @@ -0,0 +1,5 @@ +description: The request is malformed or fails validation. +content: + application/problem+json: + schema: + $ref: ../schemas/Problem.yaml diff --git a/api/openapi/components/responses/Forbidden.yaml b/api/openapi/components/responses/Forbidden.yaml new file mode 100644 index 000000000..6259185ea --- /dev/null +++ b/api/openapi/components/responses/Forbidden.yaml @@ -0,0 +1,5 @@ +description: The caller is authenticated but not allowed to do this. +content: + application/problem+json: + schema: + $ref: ../schemas/Problem.yaml diff --git a/api/openapi/components/responses/InternalError.yaml b/api/openapi/components/responses/InternalError.yaml new file mode 100644 index 000000000..20e654d44 --- /dev/null +++ b/api/openapi/components/responses/InternalError.yaml @@ -0,0 +1,5 @@ +description: Unexpected server failure. Details are in the server log. +content: + application/problem+json: + schema: + $ref: ../schemas/Problem.yaml diff --git a/api/openapi/components/responses/NotFound.yaml b/api/openapi/components/responses/NotFound.yaml new file mode 100644 index 000000000..6083a5cd1 --- /dev/null +++ b/api/openapi/components/responses/NotFound.yaml @@ -0,0 +1,5 @@ +description: No such resource or endpoint. +content: + application/problem+json: + schema: + $ref: ../schemas/Problem.yaml diff --git a/api/openapi/components/responses/NotModified.yaml b/api/openapi/components/responses/NotModified.yaml new file mode 100644 index 000000000..36bc1a61c --- /dev/null +++ b/api/openapi/components/responses/NotModified.yaml @@ -0,0 +1,4 @@ +description: Not modified. +headers: + ETag: + $ref: ../headers/ETag.yaml diff --git a/api/openapi/components/responses/Unauthorized.yaml b/api/openapi/components/responses/Unauthorized.yaml new file mode 100644 index 000000000..0209f4dd9 --- /dev/null +++ b/api/openapi/components/responses/Unauthorized.yaml @@ -0,0 +1,5 @@ +description: Missing, invalid, or expired credentials. +content: + application/problem+json: + schema: + $ref: ../schemas/Problem.yaml diff --git a/api/openapi/components/schemas/ListMeta.yaml b/api/openapi/components/schemas/ListMeta.yaml new file mode 100644 index 000000000..6321e0588 --- /dev/null +++ b/api/openapi/components/schemas/ListMeta.yaml @@ -0,0 +1,13 @@ +type: object +description: Pagination metadata carried by every list response. +required: [total, offset, limit] +properties: + total: + type: integer + description: Total number of items matching the request, ignoring pagination. + offset: + type: integer + description: Zero-based index of the first returned item. + limit: + type: integer + description: Maximum number of items in this page. diff --git a/api/openapi/components/schemas/Problem.yaml b/api/openapi/components/schemas/Problem.yaml new file mode 100644 index 000000000..0fd36d4b1 --- /dev/null +++ b/api/openapi/components/schemas/Problem.yaml @@ -0,0 +1,35 @@ +type: object +description: RFC 9457 problem details, returned for every 4xx and 5xx response. +required: [title, status, code] +properties: + type: + type: string + description: | + URI reference identifying the problem type. Omitted while the problem carries no semantics + beyond its HTTP status code, which RFC 9457 defines as `about:blank`. Problems with their + own semantics get their own URI; switch on `code` instead. + title: + type: string + description: Short human-readable summary, the same for all occurrences of this problem type. + status: + type: integer + description: HTTP status code of this response. + detail: + type: string + description: Human-readable explanation specific to this occurrence. Omitted for internal errors. + code: + type: string + description: Machine-readable error code, and the value clients switch on. New codes may be added. + enum: + - validation + - unauthorized + - forbidden + - not_found + - method_not_allowed + - unavailable + - internal + errors: + type: array + description: Per-field failures. Present only when `code` is `validation`. + items: + $ref: ./ValidationError.yaml diff --git a/api/openapi/components/schemas/ServerInfo.yaml b/api/openapi/components/schemas/ServerInfo.yaml new file mode 100644 index 000000000..8906924eb --- /dev/null +++ b/api/openapi/components/schemas/ServerInfo.yaml @@ -0,0 +1,22 @@ +type: object +description: Public server description. Everything an add-server screen needs before login. +required: [name, serverVersion, specVersion, setupRequired, loginMethods] +properties: + name: + type: string + description: Human-readable server product name. + serverVersion: + type: string + description: Version of the running server build. + specVersion: + type: string + description: Version of the OpenAPI document this server implements. + setupRequired: + type: boolean + description: True until the first admin user has been created. + loginMethods: + type: array + description: Login methods this server accepts. New methods may be added; clients ignore values they do not recognise. + items: + type: string + enum: [password] diff --git a/api/openapi/components/schemas/ValidationError.yaml b/api/openapi/components/schemas/ValidationError.yaml new file mode 100644 index 000000000..8a1cbc4f9 --- /dev/null +++ b/api/openapi/components/schemas/ValidationError.yaml @@ -0,0 +1,10 @@ +type: object +description: One field-level validation failure. +required: [field, message] +properties: + field: + type: string + description: Name of the offending query parameter, path parameter, or body field (dotted for nested). + message: + type: string + description: Why the value was rejected. diff --git a/api/openapi/openapi.yaml b/api/openapi/openapi.yaml new file mode 100644 index 000000000..73cbb7b27 --- /dev/null +++ b/api/openapi/openapi.yaml @@ -0,0 +1,39 @@ +openapi: 3.0.3 +info: + title: Navidrome API + version: 1.0.0 + description: | + Navidrome API v1. Spec-first, additive within v1. Clients discover implemented + capability modules through `GET /server` and never sniff versions. + + Enums are open: new values may be added to any enum within v1. Clients must + accept values they do not recognise instead of failing. + + Every operation declares `x-stability-level`: `alpha` operations may change or + disappear without notice, `beta` and `stable` operations only change additively. + A level is only ever raised, never lowered. + + `HEAD` is accepted wherever `GET` is. A `405` response lists the allowed methods + in its `Allow` header. + license: + name: GPL-3.0 + url: https://www.gnu.org/licenses/gpl-3.0.html +servers: + - url: /api/v1 +tags: + - name: server + description: Server discovery and the published OpenAPI document. +paths: + /server: + $ref: ./paths/server.yaml + /openapi.json: + $ref: ./paths/openapi.yaml#/json + /openapi.yaml: + $ref: ./paths/openapi.yaml#/yaml +components: + securitySchemes: + bearerAuth: + type: http + scheme: bearer + bearerFormat: JWT + description: Short-lived access token minted from a device grant. Not yet applied to any operation. diff --git a/api/openapi/paths/openapi.yaml b/api/openapi/paths/openapi.yaml new file mode 100644 index 000000000..3c25dc8b7 --- /dev/null +++ b/api/openapi/paths/openapi.yaml @@ -0,0 +1,42 @@ +json: + get: + operationId: getOpenAPISpecJSON + x-module: core + x-stability-level: alpha + tags: [server] + summary: Get the OpenAPI document (JSON) + description: The bundled OpenAPI document of the running server version. Supports ETag revalidation. + responses: + '200': + description: The OpenAPI document. + headers: + ETag: + $ref: ../components/headers/ETag.yaml + content: + application/json: + schema: + type: object + description: OpenAPI 3.0 document. + '304': + $ref: ../components/responses/NotModified.yaml +yaml: + get: + operationId: getOpenAPISpecYAML + x-module: core + x-stability-level: alpha + tags: [server] + summary: Get the OpenAPI document (YAML) + description: The bundled OpenAPI document of the running server version. Supports ETag revalidation. + responses: + '200': + description: The OpenAPI document. + headers: + ETag: + $ref: ../components/headers/ETag.yaml + content: + application/yaml: + schema: + type: object + description: OpenAPI 3.0 document. + '304': + $ref: ../components/responses/NotModified.yaml diff --git a/api/openapi/paths/server.yaml b/api/openapi/paths/server.yaml new file mode 100644 index 000000000..1f881dbb1 --- /dev/null +++ b/api/openapi/paths/server.yaml @@ -0,0 +1,19 @@ +get: + operationId: getServerInfo + x-module: core + x-stability-level: alpha + tags: [server] + summary: Describe the server + description: | + Returns the public server description. No authentication required. + Authenticated requests will additionally receive the implemented capability modules + once authentication is available. + responses: + '200': + description: Server description. + content: + application/json: + schema: + $ref: ../components/schemas/ServerInfo.yaml + '500': + $ref: ../components/responses/InternalError.yaml diff --git a/cmd/root.go b/cmd/root.go index b23674441..edfcbe69c 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -133,6 +133,9 @@ func startServer(ctx context.Context) func() error { if conf.Server.Jellyfin.Enabled { a.MountRouter("Jellyfin API", consts.URLPathJellyfinAPI, CreateJellyfinAPIRouter(ctx)) } + if conf.Server.DevAPIv1 { + a.MountRouter("API v1", consts.URLPathAPIv1, CreateAPIv1Router(ctx)) + } if conf.Server.Prometheus.Enabled { p := CreatePrometheus() // blocking call because takes <100ms but useful if fails diff --git a/cmd/wire_gen.go b/cmd/wire_gen.go index cb06cc047..19f92d9d5 100644 --- a/cmd/wire_gen.go +++ b/cmd/wire_gen.go @@ -31,6 +31,7 @@ import ( "github.com/navidrome/navidrome/plugins" "github.com/navidrome/navidrome/scanner" "github.com/navidrome/navidrome/server" + "github.com/navidrome/navidrome/server/apiv1" "github.com/navidrome/navidrome/server/events" "github.com/navidrome/navidrome/server/jellyfin" "github.com/navidrome/navidrome/server/nativeapi" @@ -142,6 +143,13 @@ func CreateJellyfinAPIRouter(ctx context.Context) *jellyfin.Router { return router } +func CreateAPIv1Router(ctx context.Context) *apiv1.Router { + sqlDB := db.Db() + dataStore := persistence.New(sqlDB) + router := apiv1.New(dataStore) + return router +} + func CreatePublicRouter() *public.Router { sqlDB := db.Db() dataStore := persistence.New(sqlDB) @@ -259,7 +267,7 @@ func getPluginManager() *plugins.Manager { // wire_injectors.go: -var allProviders = wire.NewSet(core.Set, artwork.Set, server.New, subsonic.New, jellyfin.New, jellyfin.NewDiscovery, nativeapi.New, public.New, persistence.New, lastfm.NewRouter, listenbrainz.NewRouter, events.GetBroker, scanner.GetInstance, scanner.GetWatcher, metrics.GetPrometheusInstance, db.Db, plugins.GetManager, sonic.New, wire.Bind(new(agents.PluginLoader), new(*plugins.Manager)), wire.Bind(new(scrobbler.PluginLoader), new(*plugins.Manager)), wire.Bind(new(lyrics.PluginLoader), new(*plugins.Manager)), wire.Bind(new(sonic.PluginLoader), new(*plugins.Manager)), wire.Bind(new(sonic.Engine), new(*sonic.Sonic)), wire.Bind(new(nativeapi.PluginManager), new(*plugins.Manager)), wire.Bind(new(core.PluginUnloader), new(*plugins.Manager)), wire.Bind(new(plugins.PluginMetricsRecorder), new(metrics.Metrics)), wire.Bind(new(core.Watcher), new(scanner.Watcher)), wire.Bind(new(playlists.ImageUploadService), new(artwork.Uploader))) +var allProviders = wire.NewSet(core.Set, artwork.Set, server.New, subsonic.New, jellyfin.New, jellyfin.NewDiscovery, apiv1.New, nativeapi.New, public.New, persistence.New, lastfm.NewRouter, listenbrainz.NewRouter, events.GetBroker, scanner.GetInstance, scanner.GetWatcher, metrics.GetPrometheusInstance, db.Db, plugins.GetManager, sonic.New, wire.Bind(new(agents.PluginLoader), new(*plugins.Manager)), wire.Bind(new(scrobbler.PluginLoader), new(*plugins.Manager)), wire.Bind(new(lyrics.PluginLoader), new(*plugins.Manager)), wire.Bind(new(sonic.PluginLoader), new(*plugins.Manager)), wire.Bind(new(sonic.Engine), new(*sonic.Sonic)), wire.Bind(new(nativeapi.PluginManager), new(*plugins.Manager)), wire.Bind(new(core.PluginUnloader), new(*plugins.Manager)), wire.Bind(new(plugins.PluginMetricsRecorder), new(metrics.Metrics)), wire.Bind(new(core.Watcher), new(scanner.Watcher)), wire.Bind(new(playlists.ImageUploadService), new(artwork.Uploader))) func GetPluginManager(ctx context.Context) *plugins.Manager { manager := getPluginManager() diff --git a/cmd/wire_injectors.go b/cmd/wire_injectors.go index 527617959..c5006e07f 100644 --- a/cmd/wire_injectors.go +++ b/cmd/wire_injectors.go @@ -23,6 +23,7 @@ import ( "github.com/navidrome/navidrome/plugins" "github.com/navidrome/navidrome/scanner" "github.com/navidrome/navidrome/server" + "github.com/navidrome/navidrome/server/apiv1" "github.com/navidrome/navidrome/server/events" "github.com/navidrome/navidrome/server/jellyfin" "github.com/navidrome/navidrome/server/nativeapi" @@ -37,6 +38,7 @@ var allProviders = wire.NewSet( subsonic.New, jellyfin.New, jellyfin.NewDiscovery, + apiv1.New, nativeapi.New, public.New, persistence.New, @@ -91,6 +93,12 @@ func CreateJellyfinAPIRouter(ctx context.Context) *jellyfin.Router { )) } +func CreateAPIv1Router(ctx context.Context) *apiv1.Router { + panic(wire.Build( + allProviders, + )) +} + func CreatePublicRouter() *public.Router { panic(wire.Build( allProviders, diff --git a/conf/configuration.go b/conf/configuration.go index 0173f5628..efa9cbf9a 100644 --- a/conf/configuration.go +++ b/conf/configuration.go @@ -161,6 +161,7 @@ type configOptions struct { DevExternalArtistFetchMultiplier float64 DevPreserveUnicodeInExternalCalls bool DevEnableMediaFileProbe bool + DevAPIv1 bool } type scannerOptions struct { @@ -1124,6 +1125,7 @@ func setViperDefaults() { viper.SetDefault("devshowartistpage", true) viper.SetDefault("devuishowconfig", true) viper.SetDefault("devneweventstream", true) + viper.SetDefault("devapiv1", false) viper.SetDefault("devoffsetoptimize", 50000) // Half the pool: streams may take up to this many connections, leaving the rest for the scanner, // scrobbles and the UI. See MaxOpenConns. diff --git a/consts/consts.go b/consts/consts.go index 486ea66bc..7228299c6 100644 --- a/consts/consts.go +++ b/consts/consts.go @@ -59,6 +59,7 @@ const ( URLPathPublic = "/share" URLPathPublicImages = URLPathPublic + "/img" URLPathJellyfinAPI = "/jellyfin" + URLPathAPIv1 = "/api/v1" // JellyfinServerIDKey is the Property key for the stable, persisted server Id reported by the // Jellyfin API. Jellyfin clients cache this value, so it must survive process restarts. diff --git a/go.mod b/go.mod index 464038d7e..e96b8c8b3 100644 --- a/go.mod +++ b/go.mod @@ -20,6 +20,7 @@ require ( github.com/extism/go-sdk v1.7.1 github.com/fatih/structs v1.1.0 github.com/gen2brain/webp v0.6.4 + github.com/getkin/kin-openapi v0.149.0 github.com/go-chi/chi/v5 v5.3.2 github.com/go-chi/cors v1.2.2 github.com/go-chi/httprate v0.16.0 @@ -84,6 +85,8 @@ require ( github.com/ebitengine/purego v0.11.1 // indirect github.com/fsnotify/fsnotify v1.10.1 // indirect github.com/go-logr/logr v1.4.4 // indirect + github.com/go-openapi/jsonpointer v0.22.5 // indirect + github.com/go-openapi/swag/jsonname v0.25.5 // indirect github.com/go-task/slim-sprig/v3 v3.0.0 // indirect github.com/gobwas/glob v1.0.0 // indirect github.com/goccy/go-json v0.10.6 // indirect @@ -92,6 +95,7 @@ require ( github.com/google/pprof v0.0.0-20260906184651-6331bc6350fe // indirect github.com/google/subcommands v1.2.0 // indirect github.com/gorilla/css v1.0.1 // indirect + github.com/gorilla/mux v1.8.0 // indirect github.com/hashicorp/errwrap v1.1.0 // indirect github.com/ianlancetaylor/demangle v0.0.0-20260724033716-83e58baca724 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect @@ -110,6 +114,8 @@ require ( github.com/mfridman/interpolate v0.0.2 // indirect github.com/mitchellh/go-wordwrap v1.0.1 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect + github.com/oasdiff/yaml v0.1.1 // indirect + github.com/oasdiff/yaml3 v0.0.14 // indirect github.com/ogier/pflag v0.0.1 // indirect github.com/pkg/errors v0.9.1 // indirect github.com/prometheus/client_model v0.6.2 // indirect diff --git a/go.sum b/go.sum index 929366a76..fe6dbbc3f 100644 --- a/go.sum +++ b/go.sum @@ -63,6 +63,8 @@ github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx5 github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= github.com/gen2brain/webp v0.6.4 h1:SUDdmxADOAiPQ+5ylNmuHhuYf2dOi0KgKZHL5vpVCNU= github.com/gen2brain/webp v0.6.4/go.mod h1:iGWMaCSw7t3I/Cv9llzEKmpnR36S8lS8VL/ZVjxU0JE= +github.com/getkin/kin-openapi v0.149.0 h1:ZbhmVJ4yq5RZDUsyP8lcBcGMsjsaTqXEFt6isdtMDfA= +github.com/getkin/kin-openapi v0.149.0/go.mod h1:1+BHDzstro+P5CKtPy1X4PfofnFgmRe6uvMy9+r9fKY= github.com/gkampitakis/ciinfo v0.3.2 h1:JcuOPk8ZU7nZQjdUhctuhQofk7BGHuIy0c9Ez8BNhXs= github.com/gkampitakis/ciinfo v0.3.2/go.mod h1:1NIwaOcFChN4fa/B0hEBdAb6npDlFL8Bwx4dfRLRqAo= github.com/gkampitakis/go-diff v1.3.2 h1:Qyn0J9XJSDTgnsgHRdz9Zp24RaJeKMUHg2+PDZZdC4M= @@ -79,6 +81,12 @@ github.com/go-chi/jwtauth/v5 v5.4.0 h1:Ieh0xMJsFvqylqJ02/mQHKzbbKO9DYNBh4DPKCwTw github.com/go-chi/jwtauth/v5 v5.4.0/go.mod h1:w6yjqUUXz1b8+oiJel64Sz1KJwduQM6qUA5QNzO5+bQ= github.com/go-logr/logr v1.4.4 h1:tG4xh9yMsRCAiodLVTxyrkzSZ9+o0L1Kg/+cPVcbP/8= github.com/go-logr/logr v1.4.4/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-openapi/jsonpointer v0.22.5 h1:8on/0Yp4uTb9f4XvTrM2+1CPrV05QPZXu+rvu2o9jcA= +github.com/go-openapi/jsonpointer v0.22.5/go.mod h1:gyUR3sCvGSWchA2sUBJGluYMbe1zazrYWIkWPjjMUY0= +github.com/go-openapi/swag/jsonname v0.25.5 h1:8p150i44rv/Drip4vWI3kGi9+4W9TdI3US3uUYSFhSo= +github.com/go-openapi/swag/jsonname v0.25.5/go.mod h1:jNqqikyiAK56uS7n8sLkdaNY/uq6+D2m2LANat09pKU= +github.com/go-openapi/testify/v2 v2.4.0 h1:8nsPrHVCWkQ4p8h1EsRVymA2XABB4OT40gcvAu+voFM= +github.com/go-openapi/testify/v2 v2.4.0/go.mod h1:HCPmvFFnheKK2BuwSA0TbbdxJ3I16pjwMkYkP4Ywn54= github.com/go-sql-driver/mysql v1.4.1/go.mod h1:zAC/RDZ24gD3HViQzih4MyKcchzm+sOG5ZlKdlhCg5w= github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw= github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= @@ -111,6 +119,8 @@ 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/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= +github.com/gorilla/mux v1.8.0 h1:i40aqfkR1h2SlN9hojwV5ZA91wcXFOvkdNIeFDP5koI= +github.com/gorilla/mux v1.8.0/go.mod h1:DVbg23sWSpFRCP0SfiEN6jmj59UnW/n46BH5rLB71So= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= @@ -178,6 +188,10 @@ github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/oasdiff/yaml v0.1.1 h1:6nHx+pn9gBRM6YpBlFZFQGCCd1nuvqOBtTD3KKTgGxY= +github.com/oasdiff/yaml v0.1.1/go.mod h1:EYJNoyktvWMJ0Hmhx+6qTaqMOsalUaRGT8Sj1hNcegU= +github.com/oasdiff/yaml3 v0.0.14 h1:aLJee3hxBK2H5wdXd9iPcIXb93Nty1Ge0pT171eHtkw= +github.com/oasdiff/yaml3 v0.0.14/go.mod h1:csto2xfDjYccdUn/yw/bPjj/cYTdp6HtFA0J4TWG+gg= github.com/ogier/pflag v0.0.1 h1:RW6JSWSu/RkSatfcLtogGfFgpim5p7ARQ10ECk5O750= github.com/ogier/pflag v0.0.1/go.mod h1:zkFki7tvTa0tafRvTBIZTvzYyAu6kQhPZFnshFFPE+g= github.com/onsi/ginkgo/v2 v2.33.0 h1:C8gBA6Uc2ZEubiV+SXiu5tZnMTwEmXHgkJwGozKtZf8= @@ -326,8 +340,8 @@ google.golang.org/appengine v1.6.5/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCID google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127 h1:qIbj1fsPNlZgppZ+VLlY7N33q108Sa+fhmuc+sWQYwY= -gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/ini.v1 v1.67.3 h1:iM9Lhz5MRSGhHVGGwCuzG9KO8PoirCXj/m/qTmOJJQw= gopkg.in/ini.v1 v1.67.3/go.mod h1:x/cyOwCgZqOkJoDIJ3c1KNHMo10+nLGAhh+kn3Zizss= gopkg.in/natefinch/npipe.v2 v2.0.0-20160621034901-c1b8fa8bdcce h1:+JknDZhAj8YMt7GC73Ei8pv4MzjDUNPHgQWJdtMAaDU= diff --git a/server/apiv1/api.go b/server/apiv1/api.go new file mode 100644 index 000000000..a75de3d8d --- /dev/null +++ b/server/apiv1/api.go @@ -0,0 +1,98 @@ +package apiv1 + +import ( + "errors" + "net/http" + "runtime/debug" + "slices" + "strings" + + "github.com/go-chi/chi/v5" + "github.com/navidrome/navidrome/api" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/model" +) + +type Router struct { + http.Handler + ds model.DataStore +} + +func New(ds model.DataStore) *Router { + rt := &Router{ds: ds} + rt.Handler = rt.routes() + return rt +} + +func (rt *Router) routes() http.Handler { + r := chi.NewRouter() + r.Use(problemRecoverer, headAsGet(r)) + r.NotFound(func(w http.ResponseWriter, req *http.Request) { + writeProblemStatus(w, req, http.StatusNotFound, ProblemCodeNotFound, "no such endpoint") + }) + r.MethodNotAllowed(func(w http.ResponseWriter, req *http.Request) { + w.Header().Set("Allow", strings.Join(allowedMethods(r, req), ", ")) + writeProblemStatus(w, req, http.StatusMethodNotAllowed, ProblemCodeMethodNotAllowed, "") + }) + + r.Get("/openapi.json", specHandler(withBasePath(api.SpecJSON(), `"url": `, true), "application/json")) + r.Get("/openapi.yaml", specHandler(withBasePath(api.SpecYAML(), "url: ", false), "application/yaml")) + + strict := NewStrictHandlerWithOptions(rt, nil, StrictHTTPServerOptions{ + RequestErrorHandlerFunc: func(w http.ResponseWriter, req *http.Request, err error) { + writeProblemStatus(w, req, http.StatusBadRequest, "validation", err.Error()) + }, + ResponseErrorHandlerFunc: writeProblem, + }) + HandlerWithOptions(strict, ChiServerOptions{BaseRouter: r, ErrorHandlerFunc: bindingErrorHandler}) + return r +} + +func problemRecoverer(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer func() { + rec := recover() + if rec == nil { + return + } + if err, ok := rec.(error); ok && errors.Is(err, http.ErrAbortHandler) { + panic(rec) + } + log.Error(r.Context(), "API v1: panic in handler", "panic", rec, "stack", string(debug.Stack())) + writeProblemStatus(w, r, http.StatusInternalServerError, ProblemCodeInternal, "") + }() + next.ServeHTTP(w, r) + }) +} + +var routableMethods = []string{http.MethodGet, http.MethodHead, http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete} + +// Looks routes up on mux itself: chi's RouteContext().Routes points at the parent router when mounted. +func allowedMethods(mux chi.Routes, req *http.Request) []string { + path := routePath(req) + var allowed []string + for _, m := range routableMethods { + if mux.Match(chi.NewRouteContext(), m, path) || (m == http.MethodHead && slices.Contains(allowed, http.MethodGet)) { + allowed = append(allowed, m) + } + } + return allowed +} + +func headAsGet(mux chi.Routes) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + if req.Method == http.MethodHead && !mux.Match(chi.NewRouteContext(), http.MethodHead, routePath(req)) { + chi.RouteContext(req.Context()).RouteMethod = http.MethodGet + } + next.ServeHTTP(w, req) + }) + } +} + +func routePath(req *http.Request) string { + if rctx := chi.RouteContext(req.Context()); rctx != nil && rctx.RoutePath != "" { + return rctx.RoutePath + } + return req.URL.Path +} diff --git a/server/apiv1/api_gen.go b/server/apiv1/api_gen.go new file mode 100644 index 000000000..4bffd47fc --- /dev/null +++ b/server/apiv1/api_gen.go @@ -0,0 +1,390 @@ +// Package apiv1 provides primitives to interact with the openapi HTTP API. +// +// Code generated by github.com/oapi-codegen/oapi-codegen/v2 version v2.8.0 DO NOT EDIT. +package apiv1 + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "net/http" + + "github.com/go-chi/chi/v5" +) + +// Defines values for ProblemCode. +const ( + ProblemCodeForbidden ProblemCode = "forbidden" + ProblemCodeInternal ProblemCode = "internal" + ProblemCodeMethodNotAllowed ProblemCode = "method_not_allowed" + ProblemCodeNotFound ProblemCode = "not_found" + ProblemCodeUnauthorized ProblemCode = "unauthorized" + ProblemCodeUnavailable ProblemCode = "unavailable" + ProblemCodeValidation ProblemCode = "validation" +) + +// Valid indicates whether the value is a known member of the ProblemCode enum. +func (e ProblemCode) Valid() bool { + switch e { + case ProblemCodeForbidden: + return true + case ProblemCodeInternal: + return true + case ProblemCodeMethodNotAllowed: + return true + case ProblemCodeNotFound: + return true + case ProblemCodeUnauthorized: + return true + case ProblemCodeUnavailable: + return true + case ProblemCodeValidation: + return true + default: + return false + } +} + +// Defines values for ServerInfoLoginMethods. +const ( + ServerInfoLoginMethodsPassword ServerInfoLoginMethods = "password" +) + +// Valid indicates whether the value is a known member of the ServerInfoLoginMethods enum. +func (e ServerInfoLoginMethods) Valid() bool { + switch e { + case ServerInfoLoginMethodsPassword: + return true + default: + return false + } +} + +// Problem RFC 9457 problem details, returned for every 4xx and 5xx response. +type Problem struct { + // Code Machine-readable error code, and the value clients switch on. New codes may be added. + Code ProblemCode `json:"code"` + + // Detail Human-readable explanation specific to this occurrence. Omitted for internal errors. + Detail *string `json:"detail,omitempty"` + + // Errors Per-field failures. Present only when `code` is `validation`. + Errors *[]ValidationError `json:"errors,omitempty"` + + // Status HTTP status code of this response. + Status int `json:"status"` + + // Title Short human-readable summary, the same for all occurrences of this problem type. + Title string `json:"title"` + + // Type URI reference identifying the problem type. Omitted while the problem carries no semantics + // beyond its HTTP status code, which RFC 9457 defines as `about:blank`. Problems with their + // own semantics get their own URI; switch on `code` instead. + Type *string `json:"type,omitempty"` +} + +// ProblemCode Machine-readable error code, and the value clients switch on. New codes may be added. +type ProblemCode string + +// ServerInfo Public server description. Everything an add-server screen needs before login. +type ServerInfo struct { + // LoginMethods Login methods this server accepts. New methods may be added; clients ignore values they do not recognise. + LoginMethods []ServerInfoLoginMethods `json:"loginMethods"` + + // Name Human-readable server product name. + Name string `json:"name"` + + // ServerVersion Version of the running server build. + ServerVersion string `json:"serverVersion"` + + // SetupRequired True until the first admin user has been created. + SetupRequired bool `json:"setupRequired"` + + // SpecVersion Version of the OpenAPI document this server implements. + SpecVersion string `json:"specVersion"` +} + +// ServerInfoLoginMethods defines model for ServerInfo.LoginMethods. +type ServerInfoLoginMethods string + +// ValidationError One field-level validation failure. +type ValidationError struct { + // Field Name of the offending query parameter, path parameter, or body field (dotted for nested). + Field string `json:"field"` + + // Message Why the value was rejected. + Message string `json:"message"` +} + +// InternalError RFC 9457 problem details, returned for every 4xx and 5xx response. +type InternalError = Problem + +// ServerInterface represents all server handlers. +type ServerInterface interface { + // GetServerInfo Describe the server + // (GET /server) + GetServerInfo(w http.ResponseWriter, r *http.Request) +} + +// Unimplemented server implementation that returns http.StatusNotImplemented for each endpoint. + +type Unimplemented struct{} + +// GetServerInfo Describe the server +// (GET /server) +func (_ Unimplemented) GetServerInfo(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotImplemented) +} + +// ServerInterfaceWrapper converts contexts to parameters. +type ServerInterfaceWrapper struct { + Handler ServerInterface + HandlerMiddlewares []MiddlewareFunc + ErrorHandlerFunc func(w http.ResponseWriter, r *http.Request, err error) +} + +type MiddlewareFunc func(http.Handler) http.Handler + +// GetServerInfo operation middleware +func (siw *ServerInterfaceWrapper) GetServerInfo(w http.ResponseWriter, r *http.Request) { + + handler := http.Handler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + siw.Handler.GetServerInfo(w, r) + })) + + for _, middleware := range siw.HandlerMiddlewares { + handler = middleware(handler) + } + + handler.ServeHTTP(w, r) +} + +type UnescapedCookieParamError struct { + ParamName string + Err error +} + +func (e *UnescapedCookieParamError) Error() string { + return fmt.Sprintf("error unescaping cookie parameter '%s'", e.ParamName) +} + +func (e *UnescapedCookieParamError) Unwrap() error { + return e.Err +} + +type UnmarshalingParamError struct { + ParamName string + Err error +} + +func (e *UnmarshalingParamError) Error() string { + return fmt.Sprintf("Error unmarshaling parameter %s as JSON: %s", e.ParamName, e.Err.Error()) +} + +func (e *UnmarshalingParamError) Unwrap() error { + return e.Err +} + +type RequiredParamError struct { + ParamName string +} + +func (e *RequiredParamError) Error() string { + return fmt.Sprintf("Query argument %s is required, but not found", e.ParamName) +} + +type RequiredHeaderError struct { + ParamName string + Err error +} + +func (e *RequiredHeaderError) Error() string { + return fmt.Sprintf("Header parameter %s is required, but not found", e.ParamName) +} + +func (e *RequiredHeaderError) Unwrap() error { + return e.Err +} + +type InvalidParamFormatError struct { + ParamName string + Err error +} + +func (e *InvalidParamFormatError) Error() string { + return fmt.Sprintf("Invalid format for parameter %s: %s", e.ParamName, e.Err.Error()) +} + +func (e *InvalidParamFormatError) Unwrap() error { + return e.Err +} + +type TooManyValuesForParamError struct { + ParamName string + Count int +} + +func (e *TooManyValuesForParamError) Error() string { + return fmt.Sprintf("Expected one value for %s, got %d", e.ParamName, e.Count) +} + +// Handler creates http.Handler with routing matching OpenAPI spec. +func Handler(si ServerInterface) http.Handler { + return HandlerWithOptions(si, ChiServerOptions{}) +} + +type ChiServerOptions struct { + BaseURL string + BaseRouter chi.Router + Middlewares []MiddlewareFunc + ErrorHandlerFunc func(w http.ResponseWriter, r *http.Request, err error) +} + +// HandlerFromMux creates http.Handler with routing matching OpenAPI spec based on the provided mux. +func HandlerFromMux(si ServerInterface, r chi.Router) http.Handler { + return HandlerWithOptions(si, ChiServerOptions{ + BaseRouter: r, + }) +} + +func HandlerFromMuxWithBaseURL(si ServerInterface, r chi.Router, baseURL string) http.Handler { + return HandlerWithOptions(si, ChiServerOptions{ + BaseURL: baseURL, + BaseRouter: r, + }) +} + +// HandlerWithOptions creates http.Handler with additional options +func HandlerWithOptions(si ServerInterface, options ChiServerOptions) http.Handler { + r := options.BaseRouter + + if r == nil { + r = chi.NewRouter() + } + if options.ErrorHandlerFunc == nil { + options.ErrorHandlerFunc = func(w http.ResponseWriter, r *http.Request, err error) { + http.Error(w, err.Error(), http.StatusBadRequest) + } + } + wrapper := ServerInterfaceWrapper{ + Handler: si, + HandlerMiddlewares: options.Middlewares, + ErrorHandlerFunc: options.ErrorHandlerFunc, + } + + r.Group(func(r chi.Router) { + r.Get(options.BaseURL+"/server", wrapper.GetServerInfo) + }) + + return r +} + +type InternalErrorApplicationProblemPlusJSONResponse Problem + +type GetServerInfoRequestObject struct { +} + +type GetServerInfoResponseObject interface { + VisitGetServerInfoResponse(w http.ResponseWriter) error +} + +type GetServerInfo200JSONResponse ServerInfo + +func (response GetServerInfo200JSONResponse) VisitGetServerInfoResponse(w http.ResponseWriter) error { + + var buf bytes.Buffer + if err := json.NewEncoder(&buf).Encode(response); err != nil { + return err + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(200) + _, err := buf.WriteTo(w) + return err +} + +type GetServerInfo500ApplicationProblemPlusJSONResponse struct { + InternalErrorApplicationProblemPlusJSONResponse +} + +func (response GetServerInfo500ApplicationProblemPlusJSONResponse) VisitGetServerInfoResponse(w http.ResponseWriter) error { + + var buf bytes.Buffer + if err := json.NewEncoder(&buf).Encode(response); err != nil { + return err + } + w.Header().Set("Content-Type", "application/problem+json") + w.WriteHeader(500) + _, err := buf.WriteTo(w) + return err +} + +// StrictServerInterface represents all server handlers. +type StrictServerInterface interface { + // GetServerInfo Describe the server + // (GET /server) + GetServerInfo(ctx context.Context, request GetServerInfoRequestObject) (GetServerInfoResponseObject, error) +} + +type StrictHandlerFunc func(ctx context.Context, w http.ResponseWriter, r *http.Request, request any) (any, error) +type StrictMiddlewareFunc func(f StrictHandlerFunc, operationID string) StrictHandlerFunc + +type StrictHTTPServerOptions struct { + RequestErrorHandlerFunc func(w http.ResponseWriter, r *http.Request, err error) + ResponseErrorHandlerFunc func(w http.ResponseWriter, r *http.Request, err error) +} + +func NewStrictHandler(ssi StrictServerInterface, middlewares []StrictMiddlewareFunc) ServerInterface { + return &strictHandler{ssi: ssi, middlewares: middlewares, options: StrictHTTPServerOptions{ + RequestErrorHandlerFunc: func(w http.ResponseWriter, r *http.Request, err error) { + http.Error(w, err.Error(), http.StatusBadRequest) + }, + ResponseErrorHandlerFunc: func(w http.ResponseWriter, r *http.Request, err error) { + http.Error(w, err.Error(), http.StatusInternalServerError) + }, + }} +} + +func NewStrictHandlerWithOptions(ssi StrictServerInterface, middlewares []StrictMiddlewareFunc, options StrictHTTPServerOptions) ServerInterface { + if options.RequestErrorHandlerFunc == nil { + options.RequestErrorHandlerFunc = func(w http.ResponseWriter, r *http.Request, err error) { + http.Error(w, err.Error(), http.StatusBadRequest) + } + } + if options.ResponseErrorHandlerFunc == nil { + options.ResponseErrorHandlerFunc = func(w http.ResponseWriter, r *http.Request, err error) { + http.Error(w, err.Error(), http.StatusInternalServerError) + } + } + return &strictHandler{ssi: ssi, middlewares: middlewares, options: options} +} + +type strictHandler struct { + ssi StrictServerInterface + middlewares []StrictMiddlewareFunc + options StrictHTTPServerOptions +} + +// GetServerInfo operation middleware +func (sh *strictHandler) GetServerInfo(w http.ResponseWriter, r *http.Request) { + var request GetServerInfoRequestObject + + handler := func(ctx context.Context, w http.ResponseWriter, r *http.Request, request interface{}) (interface{}, error) { + return sh.ssi.GetServerInfo(ctx, request.(GetServerInfoRequestObject)) + } + for _, middleware := range sh.middlewares { + handler = middleware(handler, "GetServerInfo") + } + + response, err := handler(r.Context(), w, r, request) + + if err != nil { + sh.options.ResponseErrorHandlerFunc(w, r, err) + } else if validResponse, ok := response.(GetServerInfoResponseObject); ok { + if err := validResponse.VisitGetServerInfoResponse(w); err != nil { + sh.options.ResponseErrorHandlerFunc(w, r, err) + } + } else if response != nil { + sh.options.ResponseErrorHandlerFunc(w, r, fmt.Errorf("unexpected response type: %T", response)) + } +} diff --git a/server/apiv1/api_test.go b/server/apiv1/api_test.go new file mode 100644 index 000000000..2552e40b0 --- /dev/null +++ b/server/apiv1/api_test.go @@ -0,0 +1,75 @@ +package apiv1 + +import ( + "net/http" + "net/http/httptest" + + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Router", func() { + var router *Router + + BeforeEach(func() { + router = New(&tests.MockDataStore{}) + }) + + It("returns a 404 problem for unknown paths", func() { + w := serve(router, httptest.NewRequest(http.MethodGet, "/api/v1/nope", nil)) + Expect(w.Code).To(Equal(http.StatusNotFound)) + Expect(w.Header().Get("Content-Type")).To(Equal(problemContentType)) + Expect(decodeProblem(w).Code).To(Equal(ProblemCodeNotFound)) + }) + + It("returns a 405 problem listing the allowed methods for a wrong method on a known path", func() { + w := serve(router, httptest.NewRequest(http.MethodPost, "/api/v1/server", nil)) + Expect(w.Code).To(Equal(http.StatusMethodNotAllowed)) + Expect(w.Header().Get("Allow")).To(Equal("GET, HEAD")) + Expect(decodeProblem(w).Code).To(Equal(ProblemCodeMethodNotAllowed)) + }) + + DescribeTable("answers HEAD wherever GET is routed", + func(path, contentType string) { + w := serve(router, httptest.NewRequest(http.MethodHead, path, nil)) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).To(Equal(contentType)) + }, + Entry("server info", "/api/v1/server", "application/json"), + Entry("JSON spec", "/api/v1/openapi.json", "application/json"), + Entry("YAML spec", "/api/v1/openapi.yaml", "application/yaml"), + ) + + It("revalidates HEAD requests with If-None-Match", func() { + etag := serve(router, httptest.NewRequest(http.MethodHead, "/api/v1/openapi.json", nil)).Header().Get("ETag") + req := httptest.NewRequest(http.MethodHead, "/api/v1/openapi.json", nil) + req.Header.Set("If-None-Match", etag) + Expect(serve(router, req).Code).To(Equal(http.StatusNotModified)) + }) + + It("returns a 404 problem for HEAD on unknown paths", func() { + w := serve(router, httptest.NewRequest(http.MethodHead, "/api/v1/nope", nil)) + Expect(w.Code).To(Equal(http.StatusNotFound)) + Expect(w.Header().Get("Allow")).To(BeEmpty()) + }) + + panicking := func(v any) http.Handler { + return problemRecoverer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { panic(v) })) + } + + It("turns a handler panic into a 500 problem", func() { + w := httptest.NewRecorder() + panicking("kaboom").ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/boom", nil)) + Expect(w.Code).To(Equal(http.StatusInternalServerError)) + p := decodeProblem(w) + Expect(p.Code).To(Equal(ProblemCodeInternal)) + Expect(p.Detail).To(BeNil()) + }) + + It("re-panics http.ErrAbortHandler so the server can drop the connection", func() { + Expect(func() { + panicking(http.ErrAbortHandler).ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/abort", nil)) + }).To(PanicWith(http.ErrAbortHandler)) + }) +}) diff --git a/server/apiv1/apiv1_suite_test.go b/server/apiv1/apiv1_suite_test.go new file mode 100644 index 000000000..f89244e38 --- /dev/null +++ b/server/apiv1/apiv1_suite_test.go @@ -0,0 +1,66 @@ +package apiv1 + +import ( + "bytes" + "errors" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/getkin/kin-openapi/openapi3" + "github.com/getkin/kin-openapi/openapi3filter" + "github.com/getkin/kin-openapi/routers" + "github.com/getkin/kin-openapi/routers/gorillamux" + "github.com/go-chi/chi/v5" + "github.com/navidrome/navidrome/api" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestAPIv1(t *testing.T) { + tests.Init(t, false) + log.SetLevel(log.LevelFatal) + RegisterFailHandler(Fail) + RunSpecs(t, "API v1 Suite") +} + +var specRouter routers.Router + +var _ = BeforeSuite(func() { + doc, err := openapi3.NewLoader().LoadFromData(api.SpecJSON()) + Expect(err).ToNot(HaveOccurred()) + specRouter, err = gorillamux.NewRouter(doc) + Expect(err).ToNot(HaveOccurred()) +}) + +// serve routes req through h mounted at /api/v1 and asserts the response conforms to the spec. +func serve(h http.Handler, req *http.Request) *httptest.ResponseRecorder { + root := chi.NewRouter() + root.Mount("/api/v1", h) + w := httptest.NewRecorder() + root.ServeHTTP(w, req) + validateAgainstSpec(req, w) + return w +} + +func validateAgainstSpec(req *http.Request, w *httptest.ResponseRecorder) { + route, pathParams, err := specRouter.FindRoute(req) + if errors.Is(err, routers.ErrPathNotFound) || errors.Is(err, routers.ErrMethodNotAllowed) { + return + } + ExpectWithOffset(2, err).ToNot(HaveOccurred()) + input := &openapi3filter.ResponseValidationInput{ + RequestValidationInput: &openapi3filter.RequestValidationInput{ + Request: req, PathParams: pathParams, Route: route, + }, + Status: w.Code, + Header: w.Header(), + Body: io.NopCloser(bytes.NewReader(w.Body.Bytes())), + Options: &openapi3filter.Options{IncludeResponseStatus: true}, + } + ExpectWithOffset(2, openapi3filter.ValidateResponse(req.Context(), input)).To(Succeed(), + "response for %s %s does not conform to the spec", req.Method, req.URL.Path) +} diff --git a/server/apiv1/oapi-codegen.yaml b/server/apiv1/oapi-codegen.yaml new file mode 100644 index 000000000..b9236de1a --- /dev/null +++ b/server/apiv1/oapi-codegen.yaml @@ -0,0 +1,12 @@ +package: apiv1 +output: server/apiv1/api_gen.go +generate: + chi-server: true + strict-server: true + models: true +output-options: + exclude-operation-ids: + - getOpenAPISpecJSON + - getOpenAPISpecYAML +compatibility: + always-prefix-enum-values: true diff --git a/server/apiv1/problem.go b/server/apiv1/problem.go new file mode 100644 index 000000000..aa6389357 --- /dev/null +++ b/server/apiv1/problem.go @@ -0,0 +1,72 @@ +package apiv1 + +import ( + "encoding/json" + "errors" + "net/http" + + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/model" +) + +const problemContentType = "application/problem+json" + +func writeProblem(w http.ResponseWriter, r *http.Request, err error) { + status, code := classifyError(err) + detail := err.Error() + if status == http.StatusInternalServerError { + log.Error(r.Context(), "API v1: unexpected error", "path", r.URL.Path, err) + detail = "" + } + writeProblemStatus(w, r, status, code, detail) +} + +func classifyError(err error) (int, ProblemCode) { + switch { + case errors.Is(err, model.ErrNotFound): + return http.StatusNotFound, ProblemCodeNotFound + case errors.Is(err, model.ErrNotAuthorized): + return http.StatusForbidden, ProblemCodeForbidden + case errors.Is(err, model.ErrInvalidAuth), errors.Is(err, model.ErrExpired): + return http.StatusUnauthorized, ProblemCodeUnauthorized + case errors.Is(err, model.ErrValidation): + return http.StatusBadRequest, ProblemCodeValidation + case errors.Is(err, model.ErrNotAvailable): + return http.StatusServiceUnavailable, ProblemCodeUnavailable + } + return http.StatusInternalServerError, ProblemCodeInternal +} + +func writeProblemStatus(w http.ResponseWriter, r *http.Request, status int, code ProblemCode, detail string, fieldErrors ...ValidationError) { + p := Problem{Title: http.StatusText(status), Status: status, Code: code} + if detail != "" { + p.Detail = &detail + } + if len(fieldErrors) > 0 { + p.Errors = &fieldErrors + } + w.Header().Set("Content-Type", problemContentType) + w.WriteHeader(status) + if err := json.NewEncoder(w).Encode(p); err != nil { + log.Warn(r.Context(), "API v1: could not write problem response", err) + } +} + +func bindingErrorHandler(w http.ResponseWriter, r *http.Request, err error) { + var fieldErrors []ValidationError + var required *RequiredParamError + var invalid *InvalidParamFormatError + var tooMany *TooManyValuesForParamError + var unmarshal *UnmarshalingParamError + switch { + case errors.As(err, &required): + fieldErrors = append(fieldErrors, ValidationError{Field: required.ParamName, Message: "is required"}) + case errors.As(err, &invalid): + fieldErrors = append(fieldErrors, ValidationError{Field: invalid.ParamName, Message: invalid.Err.Error()}) + case errors.As(err, &tooMany): + fieldErrors = append(fieldErrors, ValidationError{Field: tooMany.ParamName, Message: "expected a single value"}) + case errors.As(err, &unmarshal): + fieldErrors = append(fieldErrors, ValidationError{Field: unmarshal.ParamName, Message: unmarshal.Err.Error()}) + } + writeProblemStatus(w, r, http.StatusBadRequest, ProblemCodeValidation, err.Error(), fieldErrors...) +} diff --git a/server/apiv1/problem_test.go b/server/apiv1/problem_test.go new file mode 100644 index 000000000..256296b3c --- /dev/null +++ b/server/apiv1/problem_test.go @@ -0,0 +1,114 @@ +package apiv1 + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func decodeProblem(w *httptest.ResponseRecorder) Problem { + var p Problem + ExpectWithOffset(1, json.Unmarshal(w.Body.Bytes(), &p)).To(Succeed()) + return p +} + +var _ = Describe("problem", func() { + var w *httptest.ResponseRecorder + var r *http.Request + + BeforeEach(func() { + w = httptest.NewRecorder() + r = httptest.NewRequest(http.MethodGet, "/api/v1/server", nil) + }) + + Describe("writeProblem", func() { + DescribeTable("maps domain errors to status and code", + func(err error, status int, code ProblemCode) { + writeProblem(w, r, err) + Expect(w.Code).To(Equal(status)) + Expect(w.Header().Get("Content-Type")).To(Equal(problemContentType)) + p := decodeProblem(w) + Expect(p.Status).To(Equal(status)) + Expect(p.Code).To(Equal(code)) + Expect(p.Title).To(Equal(http.StatusText(status))) + Expect(p.Type).To(BeNil()) + Expect(w.Body.String()).ToNot(ContainSubstring(`"type"`)) + }, + Entry("not found", model.ErrNotFound, http.StatusNotFound, ProblemCodeNotFound), + Entry("not authorized", model.ErrNotAuthorized, http.StatusForbidden, ProblemCodeForbidden), + Entry("invalid auth", model.ErrInvalidAuth, http.StatusUnauthorized, ProblemCodeUnauthorized), + Entry("expired", model.ErrExpired, http.StatusUnauthorized, ProblemCodeUnauthorized), + Entry("validation", model.ErrValidation, http.StatusBadRequest, ProblemCodeValidation), + Entry("not available", model.ErrNotAvailable, http.StatusServiceUnavailable, ProblemCodeUnavailable), + Entry("unknown", errors.New("boom"), http.StatusInternalServerError, ProblemCodeInternal), + ) + + DescribeTable("keeps the wrapping context as detail for client errors", + func(err error) { + writeProblem(w, r, err) + p := decodeProblem(w) + Expect(p.Status).To(Equal(http.StatusNotFound)) + Expect(p.Detail).ToNot(BeNil()) + Expect(*p.Detail).To(ContainSubstring("album 123")) + }, + Entry("fmt.Errorf %w", fmt.Errorf("album 123: %w", model.ErrNotFound)), + Entry("errors.Join", errors.Join(errors.New("album 123"), model.ErrNotFound)), + ) + + It("hides details for internal errors", func() { + writeProblem(w, r, errors.New("db password is hunter2")) + p := decodeProblem(w) + Expect(p.Detail).To(BeNil()) + }) + }) + + Describe("writeProblemStatus", func() { + It("writes field errors only when provided", func() { + writeProblemStatus(w, r, http.StatusBadRequest, ProblemCodeValidation, "bad input", + ValidationError{Field: "limit", Message: "must be <= 2000"}) + p := decodeProblem(w) + Expect(p.Errors).ToNot(BeNil()) + Expect(*p.Errors).To(HaveLen(1)) + Expect((*p.Errors)[0].Field).To(Equal("limit")) + }) + + It("omits detail when empty", func() { + writeProblemStatus(w, r, http.StatusMethodNotAllowed, ProblemCodeMethodNotAllowed, "") + Expect(w.Body.String()).ToNot(ContainSubstring(`"detail"`)) + Expect(w.Body.String()).ToNot(ContainSubstring(`"errors"`)) + }) + }) + + Describe("bindingErrorHandler", func() { + DescribeTable("maps parameter binding errors to a validation problem with the field", + func(err error, field, message string) { + bindingErrorHandler(w, r, err) + Expect(w.Code).To(Equal(http.StatusBadRequest)) + p := decodeProblem(w) + Expect(p.Code).To(Equal(ProblemCodeValidation)) + Expect(p.Errors).ToNot(BeNil()) + Expect(*p.Errors).To(HaveLen(1)) + Expect((*p.Errors)[0].Field).To(Equal(field)) + Expect((*p.Errors)[0].Message).To(ContainSubstring(message)) + }, + Entry("required", &RequiredParamError{ParamName: "limit"}, "limit", "is required"), + Entry("invalid format", &InvalidParamFormatError{ParamName: "offset", Err: errors.New("not a number")}, "offset", "not a number"), + Entry("too many values", &TooManyValuesForParamError{ParamName: "sort", Count: 2}, "sort", "single value"), + Entry("unmarshaling", &UnmarshalingParamError{ParamName: "ids", Err: errors.New("bad json")}, "ids", "bad json"), + ) + + It("still returns a validation problem for unknown binding errors", func() { + bindingErrorHandler(w, r, errors.New("weird")) + p := decodeProblem(w) + Expect(p.Status).To(Equal(http.StatusBadRequest)) + Expect(p.Code).To(Equal(ProblemCodeValidation)) + Expect(p.Errors).To(BeNil()) + }) + }) +}) diff --git a/server/apiv1/server_info.go b/server/apiv1/server_info.go new file mode 100644 index 000000000..458363efe --- /dev/null +++ b/server/apiv1/server_info.go @@ -0,0 +1,23 @@ +package apiv1 + +import ( + "context" + "fmt" + + "github.com/navidrome/navidrome/api" + "github.com/navidrome/navidrome/consts" +) + +func (rt *Router) GetServerInfo(ctx context.Context, _ GetServerInfoRequestObject) (GetServerInfoResponseObject, error) { + count, err := rt.ds.User().CountAll(ctx) + if err != nil { + return nil, fmt.Errorf("counting users: %w", err) + } + return GetServerInfo200JSONResponse{ + Name: "Navidrome", + ServerVersion: consts.Version, + SpecVersion: api.SpecVersion(), + SetupRequired: count == 0, + LoginMethods: []ServerInfoLoginMethods{ServerInfoLoginMethodsPassword}, + }, nil +} diff --git a/server/apiv1/server_info_test.go b/server/apiv1/server_info_test.go new file mode 100644 index 000000000..b9dae6f8d --- /dev/null +++ b/server/apiv1/server_info_test.go @@ -0,0 +1,61 @@ +package apiv1 + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + + "github.com/navidrome/navidrome/api" + "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("GET /server", func() { + var ctx context.Context + var ds *tests.MockDataStore + var users *tests.MockedUserRepo + + BeforeEach(func() { + ctx = GinkgoT().Context() + users = tests.CreateMockUserRepo() + ds = &tests.MockDataStore{MockedUser: users} + }) + + get := func() (*httptest.ResponseRecorder, ServerInfo) { + w := serve(New(ds), httptest.NewRequest(http.MethodGet, "/api/v1/server", nil)) + var info ServerInfo + if w.Code == http.StatusOK { + ExpectWithOffset(1, json.Unmarshal(w.Body.Bytes(), &info)).To(Succeed()) + } + return w, info + } + + It("describes the server with setupRequired when there are no users", func() { + w, info := get() + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).To(HavePrefix("application/json")) + Expect(info.Name).To(Equal("Navidrome")) + Expect(info.ServerVersion).To(Equal(consts.Version)) + Expect(info.SpecVersion).To(Equal(api.SpecVersion())) + Expect(info.SetupRequired).To(BeTrue()) + Expect(info.LoginMethods).To(ConsistOf(ServerInfoLoginMethodsPassword)) + }) + + It("reports setupRequired=false once a user exists", func() { + Expect(users.Put(ctx, &model.User{ID: "u1", UserName: "admin", IsAdmin: true})).To(Succeed()) + _, info := get() + Expect(info.SetupRequired).To(BeFalse()) + }) + + It("returns a 500 problem when the user count fails", func() { + users.Error = errors.New("db down") + w, _ := get() + Expect(w.Code).To(Equal(http.StatusInternalServerError)) + Expect(decodeProblem(w).Code).To(Equal(ProblemCodeInternal)) + }) +}) diff --git a/server/apiv1/spec.go b/server/apiv1/spec.go new file mode 100644 index 000000000..bdca19693 --- /dev/null +++ b/server/apiv1/spec.go @@ -0,0 +1,51 @@ +package apiv1 + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "path" + "strconv" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/utils/req" + "github.com/zeebo/xxh3" +) + +func specHandler(body []byte, contentType string) http.HandlerFunc { + etag := fmt.Sprintf("%016x", xxh3.Hash(body)) + return func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("ETag", `"`+etag+`"`) + w.Header().Set("Cache-Control", "no-cache") + if req.IfNoneMatch(r, etag) { + w.WriteHeader(http.StatusNotModified) + return + } + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusOK) + _, _ = w.Write(body) + } +} + +// withBasePath adds BasePath to the advertised server URL, since only the running server knows it. +func withBasePath(body []byte, key string, quotedInBundle bool) []byte { + serverURL := path.Join(conf.Server.BasePath, consts.URLPathAPIv1) + if serverURL == consts.URLPathAPIv1 { + return body + } + oldURL := consts.URLPathAPIv1 + if quotedInBundle { + oldURL = strconv.Quote(oldURL) + } + old := []byte(key + oldURL) + if bytes.Count(body, old) != 1 { + log.Error("API v1: server URL not found in the bundled spec, serving it without the base path", "key", key) + return body + } + // A JSON string is also a valid YAML double-quoted scalar, so one encoding escapes both formats. + newURL, _ := json.Marshal(serverURL) + return bytes.Replace(body, old, append([]byte(key), newURL...), 1) +} diff --git a/server/apiv1/spec_test.go b/server/apiv1/spec_test.go new file mode 100644 index 000000000..a2368f899 --- /dev/null +++ b/server/apiv1/spec_test.go @@ -0,0 +1,136 @@ +package apiv1 + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + + "github.com/navidrome/navidrome/api" + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gopkg.in/yaml.v3" +) + +var _ = Describe("OpenAPI document routes", func() { + var router *Router + + BeforeEach(func() { + router = New(&tests.MockDataStore{}) + }) + + get := func(path string, headers map[string]string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodGet, path, nil) + for k, v := range headers { + req.Header.Set(k, v) + } + return serve(router, req) + } + + DescribeTable("serves the embedded bundle", + func(path, contentType string, body []byte) { + w := get(path, nil) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).To(Equal(contentType)) + Expect(w.Header().Get("ETag")).ToNot(BeEmpty()) + Expect(w.Header().Get("Cache-Control")).To(Equal("no-cache")) + Expect(w.Body.Bytes()).To(Equal(body)) + }, + Entry("JSON", "/api/v1/openapi.json", "application/json", api.SpecJSON()), + Entry("YAML", "/api/v1/openapi.yaml", "application/yaml", api.SpecYAML()), + ) + + DescribeTable("revalidates with If-None-Match", + func(ifNoneMatch func(etag string) string, expected int) { + etag := get("/api/v1/openapi.json", nil).Header().Get("ETag") + w := get("/api/v1/openapi.json", map[string]string{"If-None-Match": ifNoneMatch(etag)}) + Expect(w.Code).To(Equal(expected)) + if expected == http.StatusNotModified { + Expect(w.Body.Len()).To(BeZero()) + Expect(w.Header().Get("ETag")).To(Equal(etag)) + } else { + Expect(w.Body.Bytes()).To(Equal(api.SpecJSON())) + } + }, + Entry("exact ETag", func(etag string) string { return etag }, http.StatusNotModified), + Entry("ETag in a list", func(etag string) string { return `"other", ` + etag }, http.StatusNotModified), + Entry("weak ETag", func(etag string) string { return "W/" + etag }, http.StatusNotModified), + Entry("wildcard", func(string) string { return "*" }, http.StatusNotModified), + Entry("stale ETag", func(string) string { return `"stale"` }, http.StatusOK), + ) + + It("ignores Range and returns the full document", func() { + w := get("/api/v1/openapi.json", map[string]string{"Range": "bytes=0-9"}) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Accept-Ranges")).To(BeEmpty()) + Expect(w.Body.Bytes()).To(Equal(api.SpecJSON())) + }) + + It("uses different ETags for JSON and YAML", func() { + j := get("/api/v1/openapi.json", nil) + y := get("/api/v1/openapi.yaml", nil) + Expect(j.Header().Get("ETag")).ToNot(Equal(y.Header().Get("ETag"))) + }) + + Describe("with a base path", func() { + decode := func(format string, body []byte) map[string]any { + var doc map[string]any + if format == "json" { + ExpectWithOffset(1, json.Unmarshal(body, &doc)).To(Succeed()) + } else { + ExpectWithOffset(1, yaml.Unmarshal(body, &doc)).To(Succeed()) + } + return doc + } + serverURL := func(doc map[string]any) any { + return doc["servers"].([]any)[0].(map[string]any)["url"] + } + var plainETag string + + BeforeEach(func() { + plainETag = get("/api/v1/openapi.json", nil).Header().Get("ETag") + DeferCleanup(configtest.SetupConfig()) + conf.Server.BasePath = "/music" + router = New(&tests.MockDataStore{}) + }) + + DescribeTable("advertises the server under the base path and changes nothing else", + func(path, format string, bundle []byte) { + w := get(path, nil) + Expect(w.Code).To(Equal(http.StatusOK)) + served, original := decode(format, w.Body.Bytes()), decode(format, bundle) + Expect(serverURL(served)).To(Equal("/music/api/v1")) + delete(served, "servers") + delete(original, "servers") + Expect(served).To(Equal(original)) + }, + Entry("JSON", "/api/v1/openapi.json", "json", api.SpecJSON()), + Entry("YAML", "/api/v1/openapi.yaml", "yaml", api.SpecYAML()), + ) + + It("uses its own ETag, and still revalidates", func() { + etag := get("/api/v1/openapi.json", nil).Header().Get("ETag") + Expect(etag).ToNot(Equal(plainETag)) + Expect(get("/api/v1/openapi.json", map[string]string{"If-None-Match": etag}).Code).To(Equal(http.StatusNotModified)) + }) + + It("escapes base paths that need quoting", func() { + conf.Server.BasePath = "/my music: \"live\"" + router = New(&tests.MockDataStore{}) + Expect(serverURL(decode("json", get("/api/v1/openapi.json", nil).Body.Bytes()))).To(Equal("/my music: \"live\"/api/v1")) + Expect(serverURL(decode("yaml", get("/api/v1/openapi.yaml", nil).Body.Bytes()))).To(Equal("/my music: \"live\"/api/v1")) + }) + }) + + DescribeTable("the bundle advertises the API path exactly once, which the base-path rewrite relies on", + func(bundle []byte) { + Expect(bytes.Count(bundle, []byte(consts.URLPathAPIv1))).To(Equal(1)) + }, + Entry("JSON", api.SpecJSON()), + Entry("YAML", api.SpecYAML()), + ) +}) diff --git a/server/imghttp/headers.go b/server/imghttp/headers.go index 9308bcdab..354b11793 100644 --- a/server/imghttp/headers.go +++ b/server/imghttp/headers.go @@ -4,9 +4,9 @@ package imghttp import ( "net/http" - "strings" "github.com/navidrome/navidrome/core/artwork" + "github.com/navidrome/navidrome/utils/req" ) // WriteImageHeaders applies the artwork caching contract and reports whether a 304 was written @@ -40,28 +40,9 @@ func WriteImageHeaders(w http.ResponseWriter, r *http.Request, img *artwork.Imag h.Set("Cache-Control", "public, no-cache") } - if etag != "" && ifNoneMatch(r.Header.Get("If-None-Match"), etag) { + if etag != "" && req.IfNoneMatch(r, etag) { w.WriteHeader(http.StatusNotModified) return true } return false } - -// ifNoneMatch reports whether If-None-Match asserts hash, using RFC 9110 weak comparison. -func ifNoneMatch(header, hash string) bool { - header = strings.TrimSpace(header) - if header == "" { - return false - } - if header == "*" { - return true - } - for tag := range strings.SplitSeq(header, ",") { - tag = strings.TrimSpace(tag) - tag = strings.TrimPrefix(tag, "W/") - if strings.Trim(tag, `"`) == hash { - return true - } - } - return false -} diff --git a/utils/req/req.go b/utils/req/req.go index 861cca9f7..6a863184c 100644 --- a/utils/req/req.go +++ b/utils/req/req.go @@ -178,3 +178,21 @@ func (r *Values) Float64Or(param string, def float64) float64 { } return f } + +// IfNoneMatch reports whether the request's If-None-Match asserts etag (unquoted), using RFC 9110 weak comparison. +func IfNoneMatch(r *http.Request, etag string) bool { + header := strings.TrimSpace(r.Header.Get("If-None-Match")) + if header == "" { + return false + } + if header == "*" { + return true + } + for tag := range strings.SplitSeq(header, ",") { + tag = strings.TrimPrefix(strings.TrimSpace(tag), "W/") + if strings.Trim(tag, `"`) == etag { + return true + } + } + return false +} diff --git a/utils/req/req_test.go b/utils/req/req_test.go index 5f9de8483..d64f51b02 100644 --- a/utils/req/req_test.go +++ b/utils/req/req_test.go @@ -290,3 +290,21 @@ var _ = Describe("Request Helpers", func() { }) }) }) + +var _ = Describe("IfNoneMatch", func() { + DescribeTable("matches the ETag", + func(header string, expected bool) { + r := httptest.NewRequest("GET", "/", nil) + if header != "" { + r.Header.Set("If-None-Match", header) + } + Expect(req.IfNoneMatch(r, "abc123")).To(Equal(expected)) + }, + Entry("absent header", "", false), + Entry("exact quoted tag", `"abc123"`, true), + Entry("weak tag", `W/"abc123"`, true), + Entry("tag in a list", `"other", W/"abc123"`, true), + Entry("wildcard", "*", true), + Entry("different tag", `"stale"`, false), + ) +}) From 46c432719fdd2fb34c2ec6c4f573a1488fb90215 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sat, 26 Sep 2026 23:30:50 -0400 Subject: [PATCH 18/26] fix(log): redact LastFM keys and Prometheus password in config dump (#6233) The startup Configuration dump is rendered with pretty.Sprintf("%# v"), which pads multi-line struct fields with spaces after the colon. The ApiKey and Secret redaction patterns required the quote right after the colon, so LastFM.ApiKey and LastFM.Secret were logged in clear text even with EnableLogRedacting on. Allow optional whitespace after the colon, like the other config patterns already do. Prometheus.Password had no redaction pattern at all. Add one that also skips escaped quotes, since the password can hold any character and pretty prints it Go-quoted. Add tests for the padded and unpadded forms, plus one that redacts a real pretty.Sprintf dump of LastFM- and Prometheus-shaped structs so a padding change in pretty can't bring the leak back. Reported in https://github.com/navidrome/navidrome/discussions/6232 --- log/log.go | 6 +++-- log/log_test.go | 67 ++++++++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 70 insertions(+), 3 deletions(-) diff --git a/log/log.go b/log/log.go index de2f171b2..da1d7622e 100644 --- a/log/log.go +++ b/log/log.go @@ -27,14 +27,16 @@ var redacted = &Hook{ AcceptedLevels: logrus.AllLevels, RedactionList: []string{ // Keys from the config - "(ApiKey:\")[\\w]*", - "(Secret:\")[\\w]*", + "(ApiKey:[\\s]*\")[\\w]*", + "(Secret:[\\s]*\")[\\w]*", "(PasswordEncryptionKey:[\\s]*\")[^\"]*", "(UserHeader:[\\s]*\")[^\"]*", "(TrustedSources:[\\s]*\")[^\"]*", "(MetricsPath:[\\s]*\")[^\"]*", "(DevAutoCreateAdminPassword:[\\s]*\")[^\"]*", "(DevAutoLoginUsername:[\\s]*\")[^\"]*", + // Prometheus.Password. Any character is allowed, so skip escaped quotes in the value + `(Password:[\s]*")(?:[^"\\]|\\.)*`, // UI appConfig "(subsonicToken:)[\\w]+(\\s)", diff --git a/log/log_test.go b/log/log_test.go index 82207c672..184ff57db 100644 --- a/log/log_test.go +++ b/log/log_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + "github.com/kr/pretty" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" "github.com/sirupsen/logrus" @@ -94,7 +95,7 @@ var _ = Describe("Logger", func() { SetLogSourceLine(true) Error("A crash happened") // NOTE: This assertion breaks if the line number above changes - Expect(hook.LastEntry().Data[" source"]).To(ContainSubstring("/log/log_test.go:95")) + Expect(hook.LastEntry().Data[" source"]).To(ContainSubstring("/log/log_test.go:96")) Expect(hook.LastEntry().Message).To(Equal("A crash happened")) }) @@ -291,5 +292,69 @@ var _ = Describe("Logger", func() { Expect(got).ToNot(ContainSubstring("secret")) Expect(got).To(ContainSubstring(`"User-Agent":["Finamp/1.0"]`)) }) + + // https://github.com/navidrome/navidrome/discussions/6232 + DescribeTable("redacts config keys in the startup Configuration dump", + func(line, expected string) { + Expect(Redact(line)).To(Equal(expected)) + }, + Entry("unpadded ApiKey", `ApiKey:"0123456789abcdef0123456789abcdef"`, `ApiKey:"[REDACTED]"`), + Entry("unpadded Secret", `Secret:"fedcba9876543210fedcba9876543210"`, `Secret:"[REDACTED]"`), + Entry("padded ApiKey", ` ApiKey: "0123456789abcdef0123456789abcdef",`, + ` ApiKey: "[REDACTED]",`), + Entry("padded Secret", ` Secret: "fedcba9876543210fedcba9876543210",`, + ` Secret: "[REDACTED]",`), + Entry("unpadded Prometheus Password", `Password:"p@ss w0rd!"`, `Password:"[REDACTED]"`), + Entry("padded Prometheus Password", ` Password: "p@ss w0rd!",`, ` Password: "[REDACTED]",`), + Entry("Prometheus Password with escaped quotes", ` Password: "a\"b\\\"c",`, + ` Password: "[REDACTED]",`), + ) + + It("redacts secrets in a pretty-printed config struct", func() { + // Mirrors conf.lastfmOptions and conf.prometheusOptions (conf imports log, so it can't be + // used here). pretty only breaks a struct into padded lines when it is long enough, so + // keep all the fields. + type lastfmOptions struct { + Enabled bool + ApiKey string + Secret string + Language string + ScrobbleFirstArtistOnly bool + Languages []string + } + type prometheusOptions struct { + Enabled bool + MetricsPath string + Password string + } + type configOptions struct { + Address string + LastFM lastfmOptions + Prometheus prometheusOptions + } + cfg := configOptions{ + Address: "0.0.0.0", + LastFM: lastfmOptions{ //nolint:gosec + Enabled: true, + ApiKey: "0123456789abcdef0123456789abcdef", + Secret: "fedcba9876543210fedcba9876543210", + Language: "en", + Languages: []string{"en"}, + }, + Prometheus: prometheusOptions{ //nolint:gosec + Enabled: true, + MetricsPath: "/metrics", + Password: `prom"pass-tail`, + }, + } + dump := pretty.Sprintf("Configuration: %# v", cfg) + Expect(dump).To(MatchRegexp(`ApiKey:\s{2,}"`), "the dump must use the padded layout") + + got := Redact(dump) + Expect(got).ToNot(ContainSubstring(cfg.LastFM.ApiKey)) + Expect(got).ToNot(ContainSubstring(cfg.LastFM.Secret)) + Expect(got).ToNot(ContainSubstring("pass-tail")) + Expect(got).To(ContainSubstring(`"en"`)) + }) }) }) From ce484083bf070662e3b976daeca156fd7822c332 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sun, 27 Sep 2026 14:18:24 -0400 Subject: [PATCH 19/26] fix(server): exit with an error code when the server fails to start (#6236) When a startup step failed (for example, the port was already in use), runNavidrome only logged the error and returned. In service mode, service.Run() kept waiting for a stop signal, so the process stayed up serving nothing and the service manager never restarted it. A plain run exited with code 0. runNavidrome now returns the error, unless its context was cancelled by a normal shutdown. Both the plain run and the service goroutine exit with code 1 on that error. The systemd unit no longer lists 1, 2 and 8 in SuccessExitStatus, so Restart=on-failure restarts the service on exit code 1. Fixes #6235 --- cmd/root.go | 20 ++++++++++++-------- cmd/svc.go | 9 +++++++-- 2 files changed, 19 insertions(+), 10 deletions(-) diff --git a/cmd/root.go b/cmd/root.go index edfcbe69c..089f09472 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -44,7 +44,9 @@ Complete documentation is available at https://www.navidrome.org/docs`, preRun() }, Run: func(cmd *cobra.Command, args []string) { - runNavidrome(cmd.Context()) + if err := runNavidrome(cmd.Context()); err != nil { + log.Fatal("Fatal error in Navidrome. Aborting", err) + } }, PostRun: func(cmd *cobra.Command, args []string) { postRun() @@ -76,12 +78,12 @@ func postRun() { } // runNavidrome is the main entry point for the Navidrome server. It starts all the services and blocks. -// If any of the services returns an error, it will log it and exit. If the process receives a signal to exit, -// it will cancel the context and exit gracefully. -func runNavidrome(ctx context.Context) { - defer db.Init(ctx)() +// If any of the services returns an error, it stops the others and returns that error, so the caller can +// exit with a non-zero code. If the context is cancelled (a signal or a service stop), it returns nil. +func runNavidrome(parentCtx context.Context) error { + defer db.Init(parentCtx)() - g, ctx := errgroup.WithContext(ctx) + g, ctx := errgroup.WithContext(parentCtx) g.Go(startServer(ctx)) g.Go(startSignaller(ctx)) g.Go(startScheduler(ctx)) @@ -102,9 +104,11 @@ func runNavidrome(ctx context.Context) { log.Warn(ctx, "Automatic Scanning is DISABLED") } - if err := g.Wait(); err != nil { - log.Error("Fatal error in Navidrome. Aborting", err) + // Errors caused by a normal shutdown are not failures + if err := g.Wait(); err != nil && parentCtx.Err() == nil { + return err } + return nil } // mainContext returns a context that is cancelled when the process receives a signal to exit. diff --git a/cmd/svc.go b/cmd/svc.go index c71f5ef2b..4e8b1fd85 100644 --- a/cmd/svc.go +++ b/cmd/svc.go @@ -53,8 +53,13 @@ func (p *svcControl) Start(service.Service) error { p.done = make(chan struct{}) p.ctx, p.cancel = context.WithCancel(context.Background()) go func() { - runNavidrome(p.ctx) + err := runNavidrome(p.ctx) close(p.done) + // service.Run() only returns when it gets a stop request, so exit here to let the + // service manager see the failure and restart the service + if err != nil { + log.Fatal("Fatal error in Navidrome. Aborting", err) + } }() return nil } @@ -74,7 +79,7 @@ func (p *svcControl) Stop(service.Service) error { var svcInstance = sync.OnceValue(func() service.Service { options := make(service.KeyValue) options["Restart"] = "on-failure" - options["SuccessExitStatus"] = "1 2 8 SIGKILL" + options["SuccessExitStatus"] = "SIGKILL" options["UserService"] = false options["LogDirectory"] = conf.Server.DataFolder.String() options["SystemdScript"] = systemdScript From 4cdffd5633e05cbc29cfd71e34382bf25eae256a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sun, 27 Sep 2026 21:56:58 -0400 Subject: [PATCH 20/26] feat(subsonic): OpenSubsonic API key authentication (#6219) * feat(persistence): store hashed API keys on players * feat(core): refresh key-bound players without renaming them Add Players.Touch, which records usage for a player already identified by an API key without guessing its identity or overwriting its name. Register also stops renaming players that have an API key. Register no longer returns player save errors (or a stale FindMatch ErrNotFound when the save is rate-limited); save failures are only logged, and only the transcoding lookup error is returned, same as Touch. * feat(subsonic): authenticate with OpenSubsonic API keys Co-authored-by: amCap1712 * feat(subsonic): add tokenInfo and advertise apiKeyAuthentication * feat(server): add endpoints to generate and revoke player API keys * feat(ui): manage player API keys Co-authored-by: amCap1712 * fix(subsonic): throttle API keys per key and IP A stale key on one device exhausted the shared per-IP bucket and locked out every valid key from the same IP. The limiter only stores a hash of the bucket string, so the key is not retained. Also adds e2e coverage of API key auth through the real repository, and clarifies the player resolution log message. * fix(ui): keep the new API key dialog open until closed The key is shown only once, so Escape and backdrop clicks no longer dismiss it. Also clarifies when the key can be used as a password. * refactor: simplify API key code paths Share the player refresh tail between Register and Touch, fold the ownership-filtered write tail into execOwned, parse the query once for apiKey conflicts, derive HasAPIKey in the player mock, share the player form inputs between create and edit, and pick the delete button by key state instead of spreading conditional props. * feat(players): set API keys through the player record The key is a write-only apiKey field applied on save: required and owner-only on create, optional on edit, empty to revoke. Replaces the generate/revoke endpoints. * fix(players): reject API keys already in use Creating or editing a player with a key another player already has now returns a validation error instead of a 500, and a create that loses the race no longer leaves a keyless player behind. Ownership is checked before the key on create. * feat(ui): edit player API keys as a form field Replaces the show-once dialog, whose icon-less Close button was invisible on mobile. The key is generated in the browser, required and pre-filled on create. * fix(ui): keep new player API keys out of the record cache The json-server create response echoes the request body, and undoable edits merge the payload into the cache, so the key could reappear on the edit page. Strip it from the create result and save player edits pessimistically. Also fall back to a prompt when the clipboard write fails. * fix(ui): polish player API key field Set userId on the created player record so owner actions show immediately, and show a neutral no-key message to non-owners. * refactor: simplify player API key create and field Write the key hash in the create INSERT so the unique index settles races, re-read the created player instead of hand-building the cached record, reuse isWritable for the revoke check, and collapse the key field's derived state and generate/regenerate buttons. * fix(ui): let the API key field size like other inputs fullWidth is now opt-in instead of forced. * fix(ui): align the API key field with other player inputs Apply react-admin's input className, move the actions (now including Copy) below the field, and use a monospace font so the whole key fits. * fix(ui): redirect to the player list after create Matches the other create pages. * refactor(persistence): name the write-access rule for owned rows Owned-row writes now say which row they target and who may write it: ownedRow(rowID, ownerOrAdmin|ownerOnly) builds the WHERE, updateOwnedRow applies it, and SetAPIKey uses ownerOnly instead of a hand-built user_id filter. updateOwned/deleteOwned keep their signatures. * fix(players): apply an edit's key change and fields atomically Update now runs SetAPIKey and the column update in one transaction. Also shares the key format check, drops FindByAPIKey's unneeded empty-key guard, and sets the context username only on the apiKey path. * fix(subsonic): treat any credential param sent with apiKey as a conflict The spec requires error 43 when u, p, t or s is present with apiKey, even with an empty value. * refactor(subsonic): leave the player cookie code unchanged for key-bound requests Return early instead of wrapping the cookie block, so the diff (and CodeQL's view of it) matches master. * fix(subsonic): don't count key lookup errors as failed logins A database error while checking a key sent as the password now surfaces as a server error instead of a bad password, so it no longer feeds the failed-login limiter. * feat(players): use nds_ as the API key prefix Part of a Navidrome secret prefix family (nd + a letter for the kind), alongside ndg_ for API v1 grants. * feat(ui): make player API keys easier to find Label the Settings menu entry "Players & API keys", add an API key filter to the player list, show the key icon in the mobile list, and add Brazilian Portuguese translations for the new player strings. Signed-off-by: Deluan * feat(ui): always show the player API key filter Signed-off-by: Deluan * fix(ui): hide the unset Last Seen date in the player list Players created by hand have no last_seen yet, which showed as 12/31/1. Signed-off-by: Deluan --------- Signed-off-by: Deluan Co-authored-by: amCap1712 --- consts/consts.go | 1 + core/players.go | 28 +- core/players_test.go | 37 +++ ...20260924010054_add_player_api_key_hash.sql | 8 + model/player.go | 4 + persistence/player_repository.go | 121 +++++++- persistence/player_repository_test.go | 269 +++++++++++++++++- persistence/sql_base_repository.go | 40 ++- persistence/sql_base_repository_test.go | 20 ++ resources/i18n/pt-br.json | 25 +- server/subsonic/api.go | 1 + server/subsonic/e2e/subsonic_apikey_test.go | 54 ++++ server/subsonic/middlewares.go | 118 +++++++- server/subsonic/middlewares_test.go | 162 +++++++++++ server/subsonic/opensubsonic.go | 1 + server/subsonic/opensubsonic_test.go | 66 ++--- .../Responses TokenInfo should match .JSON | 10 + .../Responses TokenInfo should match .XML | 3 + server/subsonic/responses/errors.go | 38 +-- server/subsonic/responses/responses.go | 5 + server/subsonic/responses/responses_test.go | 14 + server/subsonic/system.go | 8 + server/subsonic/system_test.go | 23 ++ tests/mock_data_store.go | 2 +- tests/mock_player_repo.go | 75 +++++ ui/src/App.jsx | 2 +- ui/src/dataProvider/wrapperDataProvider.js | 9 + .../dataProvider/wrapperDataProvider.test.js | 15 + ui/src/i18n/en.json | 25 +- ui/src/layout/AppBar.jsx | 8 +- ui/src/layout/AppBar.test.jsx | 24 +- ui/src/player/ApiKeyInput.jsx | 107 +++++++ ui/src/player/ApiKeyInput.test.jsx | 172 +++++++++++ ui/src/player/PlayerCreate.jsx | 33 +++ ui/src/player/PlayerCreate.test.jsx | 47 +++ ui/src/player/PlayerEdit.jsx | 55 ++-- ui/src/player/PlayerEdit.test.jsx | 25 ++ ui/src/player/PlayerList.jsx | 17 +- ui/src/player/apiKey.js | 15 + ui/src/player/apiKey.test.js | 15 + ui/src/player/index.js | 2 + ui/src/player/playerInputs.jsx | 34 +++ 42 files changed, 1610 insertions(+), 128 deletions(-) create mode 100644 db/migrations/20260924010054_add_player_api_key_hash.sql create mode 100644 server/subsonic/e2e/subsonic_apikey_test.go create mode 100644 server/subsonic/responses/.snapshots/Responses TokenInfo should match .JSON create mode 100644 server/subsonic/responses/.snapshots/Responses TokenInfo should match .XML create mode 100644 server/subsonic/system_test.go create mode 100644 tests/mock_player_repo.go create mode 100644 ui/src/player/ApiKeyInput.jsx create mode 100644 ui/src/player/ApiKeyInput.test.jsx create mode 100644 ui/src/player/PlayerCreate.jsx create mode 100644 ui/src/player/PlayerCreate.test.jsx create mode 100644 ui/src/player/PlayerEdit.test.jsx create mode 100644 ui/src/player/apiKey.js create mode 100644 ui/src/player/apiKey.test.js create mode 100644 ui/src/player/playerInputs.jsx diff --git a/consts/consts.go b/consts/consts.go index 7228299c6..9bdac9125 100644 --- a/consts/consts.go +++ b/consts/consts.go @@ -49,6 +49,7 @@ const ( DefaultEncryptionKey = "just for obfuscation" PasswordsEncryptedKey = "PasswordsEncryptedKey" PasswordAutogenPrefix = "__NAVIDROME_AUTOGEN__" //nolint:gosec + APIKeyPrefix = "nds_" DevInitialUserName = "admin" DevInitialName = "Dev Admin" diff --git a/core/players.go b/core/players.go index e03d8caa2..6fe86fd70 100644 --- a/core/players.go +++ b/core/players.go @@ -17,6 +17,7 @@ import ( type Players interface { Get(ctx context.Context, playerId string) (*model.Player, error) Register(ctx context.Context, id, client, userAgent, ip string) (*model.Player, *model.Transcoding, error) + Touch(ctx context.Context, plr model.Player, client, userAgent, ip string) (*model.Player, *model.Transcoding, error) } func NewPlayers(ds model.DataStore) Players { @@ -33,7 +34,6 @@ type players struct { func (p *players) Register(ctx context.Context, playerID, client, userAgent, ip string) (*model.Player, *model.Transcoding, error) { var plr *model.Player - var trc *model.Transcoding var err error user, _ := request.UserFrom(ctx) if playerID != "" { @@ -58,7 +58,21 @@ func (p *players) Register(ctx context.Context, playerID, client, userAgent, ip log.Info(ctx, "Registering new player", "id", plr.ID, "client", client, "username", username, "type", userAgent) } } - plr.Name = fmt.Sprintf("%s [%s]", client, userAgent) + if !plr.HasAPIKey { + plr.Name = fmt.Sprintf("%s [%s]", client, userAgent) + } + return p.refresh(ctx, plr, userAgent, ip) +} + +// Touch refreshes a player that the request already identified (by API key), without guessing or renaming it. +func (p *players) Touch(ctx context.Context, plr model.Player, client, userAgent, ip string) (*model.Player, *model.Transcoding, error) { + if plr.Client == "" { + plr.Client = client + } + return p.refresh(ctx, &plr, userAgent, ip) +} + +func (p *players) refresh(ctx context.Context, plr *model.Player, userAgent, ip string) (*model.Player, *model.Transcoding, error) { plr.UserAgent = userAgent plr.IP = ip plr.LastSeen = time.Now() @@ -66,14 +80,14 @@ func (p *players) Register(ctx context.Context, playerID, client, userAgent, ip ctx, cancel := context.WithTimeout(ctx, time.Second) defer cancel() - 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 err := p.ds.Player().Put(ctx, plr); err != nil { + log.Warn(ctx, "Could not save player", "id", plr.ID, "client", plr.Client, "username", userName(ctx), "type", plr.UserAgent, err) } }) - if plr.TranscodingId != "" { - trc, err = p.ds.Transcoding().Get(ctx, plr.TranscodingId) + if plr.TranscodingId == "" { + return plr, nil, nil } + trc, err := p.ds.Transcoding().Get(ctx, plr.TranscodingId) return plr, trc, err } diff --git a/core/players_test.go b/core/players_test.go index 302d63157..e452c52ba 100644 --- a/core/players_test.go +++ b/core/players_test.go @@ -114,6 +114,15 @@ var _ = Describe("Players", func() { Expect(trc.ID).To(Equal("1")) }) + It("does not rename a player that has an API key", func() { + plr := &model.Player{ID: "123", Name: "My Phone", Client: "client", UserId: "userid", HasAPIKey: true} + repo.add(plr) + p, _, err := players.Register(ctx, "123", "client", "chrome", "1.2.3.4") + Expect(err).ToNot(HaveOccurred()) + Expect(p.ID).To(Equal("123")) + Expect(p.Name).To(Equal("My Phone")) + }) + Context("bad username casing", func() { ctx := log.NewContext(context.TODO()) ctx = request.WithUser(ctx, model.User{ID: "userid", UserName: "Johndoe"}) @@ -130,6 +139,34 @@ var _ = Describe("Players", func() { }) }) }) + + Describe("Touch", func() { + It("records usage but keeps the name and client", func() { + plr := model.Player{ID: "123", Name: "My Phone", Client: "Symfonium", UserId: "userid", HasAPIKey: true} + p, trc, err := players.Touch(ctx, plr, "OtherClient", "android", "1.2.3.4") + Expect(err).ToNot(HaveOccurred()) + Expect(p.Name).To(Equal("My Phone")) + Expect(p.Client).To(Equal("Symfonium")) + Expect(p.UserAgent).To(Equal("android")) + Expect(p.IP).To(Equal("1.2.3.4")) + Expect(p.LastSeen).To(BeTemporally(">=", beforeRegister)) + Expect(repo.lastSaved).To(Equal(p)) + Expect(trc).To(BeNil()) + }) + + It("fills in the client on first use", func() { + p, _, err := players.Touch(ctx, model.Player{ID: "123", Name: "Manual", UserId: "userid"}, "Symfonium", "android", "1.2.3.4") + Expect(err).ToNot(HaveOccurred()) + Expect(p.Client).To(Equal("Symfonium")) + }) + + It("returns the player's transcoding", func() { + p, trc, err := players.Touch(ctx, model.Player{ID: "123", UserId: "userid", TranscodingId: "1"}, "c", "ua", "1.2.3.4") + Expect(err).ToNot(HaveOccurred()) + Expect(p.ID).To(Equal("123")) + Expect(trc.ID).To(Equal("1")) + }) + }) }) type mockPlayerRepository struct { diff --git a/db/migrations/20260924010054_add_player_api_key_hash.sql b/db/migrations/20260924010054_add_player_api_key_hash.sql new file mode 100644 index 000000000..bbc4cf4d9 --- /dev/null +++ b/db/migrations/20260924010054_add_player_api_key_hash.sql @@ -0,0 +1,8 @@ +-- +goose Up +-- +goose StatementBegin +alter table player add column api_key_hash varchar default null; +create unique index if not exists player_api_key_hash on player(api_key_hash); +-- +goose StatementEnd + +-- +goose Down +SELECT 1; diff --git a/model/player.go b/model/player.go index 2e4484a10..c03058419 100644 --- a/model/player.go +++ b/model/player.go @@ -21,6 +21,8 @@ type Player struct { MaxBitRate int `structs:"max_bit_rate" json:"maxBitRate"` ReportRealPath bool `structs:"report_real_path" json:"reportRealPath"` ScrobbleEnabled bool `structs:"scrobble_enabled" json:"scrobbleEnabled"` + HasAPIKey bool `structs:"-" db:"has_api_key" json:"hasApiKey"` + APIKey *string `structs:"-" json:"apiKey,omitempty"` } type Players []Player @@ -33,4 +35,6 @@ type PlayerRepository interface { 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) + FindByAPIKey(ctx context.Context, key string) (*Player, error) + SetAPIKey(ctx context.Context, playerID, key string) error } diff --git a/persistence/player_repository.go b/persistence/player_repository.go index e46a8d82d..ba5d26794 100644 --- a/persistence/player_repository.go +++ b/persistence/player_repository.go @@ -2,10 +2,16 @@ package persistence import ( "context" + "crypto/sha256" + "encoding/hex" + "regexp" + "strings" . "github.com/Masterminds/squirrel" "github.com/deluan/rest" + "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/id" "github.com/pocketbase/dbx" ) @@ -17,7 +23,8 @@ func NewPlayerRepository(db dbx.Builder) model.PlayerRepository { r := &playerRepository{} r.db = db r.registerModel(&model.Player{}, map[string]filterFunc{ - "name": containsFilter("player.name"), + "name": containsFilter("player.name"), + "hasapikey": hasAPIKeyFilter, }) r.setSortMappings(map[string]string{ "user_name": "username", //TODO rename all user_name and userName to username @@ -25,6 +32,13 @@ func NewPlayerRepository(db dbx.Builder) model.PlayerRepository { return r } +func hasAPIKeyFilter(_ string, value any) Sqlizer { + if v, _ := value.(string); strings.EqualFold(v, "true") { + return NotEq{"player.api_key_hash": nil} + } + return Eq{"player.api_key_hash": nil} +} + func (r *playerRepository) Put(ctx context.Context, p *model.Player) error { _, err := r.put(ctx, p.ID, p) return err @@ -32,7 +46,7 @@ func (r *playerRepository) Put(ctx context.Context, p *model.Player) error { func (r *playerRepository) selectPlayer(ctx context.Context, options ...model.QueryOptions) SelectBuilder { return r.newSelect(ctx, options...). - Columns("player.*"). + Columns("player.*", "player.api_key_hash is not null as has_api_key"). Join("user ON player.user_id = user.id"). Columns("user.user_name username") } @@ -103,32 +117,115 @@ func (r *playerRepository) ReadAll(ctx context.Context, options ...rest.QueryOpt return res, err } -// 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(ctx context.Context, p *model.Player) bool { - u := loggedUser(ctx) - return u.IsAdmin || p.UserId == u.ID +var apiKeyFormat = regexp.MustCompile(`^` + consts.APIKeyPrefix + `[0-9A-Za-z]{22}$`) + +func apiKeyValidationError(msg string) error { + return &rest.ValidationError{Errors: map[string]string{"apiKey": msg}} +} + +func validateAPIKey(key string) error { + if !apiKeyFormat.MatchString(key) { + return apiKeyValidationError("resources.player.validation.apiKeyFormat") + } + return nil } func (r *playerRepository) Save(ctx context.Context, t *model.Player) (string, error) { - if !r.isPermitted(ctx, t) { + u := loggedUser(ctx) + if t.UserId == "" && u.ID != invalidUserId { + t.UserId = u.ID + } + if t.UserId != u.ID { return "", rest.ErrPermissionDenied } - return r.put(ctx, "", t) // Save only creates; edits go through the owner-scoped Update + // Hand-made players are only reachable through a key, so one is required + if t.APIKey == nil || *t.APIKey == "" { + return "", apiKeyValidationError("ra.validation.required") + } + if err := validateAPIKey(*t.APIKey); err != nil { + return "", err + } + values, err := toSQLArgs(t) + if err != nil { + return "", err + } + // Save only creates, so the key hash goes in the same INSERT and the unique index settles races + values["id"] = id.NewRandom() + values["api_key_hash"] = hashAPIKey(*t.APIKey) + _, err = r.executeSQL(ctx, Insert(r.tableName).SetMap(values)) + if isUniqueViolation(err) { + return "", apiKeyValidationError("ra.validation.unique") + } + if err != nil { + return "", err + } + return values["id"].(string), nil } func (r *playerRepository) Update(ctx context.Context, id string, entity model.Player, cols ...string) error { t := &entity t.ID = id - return r.updateOwned(ctx, id, t, cols...) + if t.APIKey == nil { + return r.updateOwned(ctx, id, t, cols...) + } + // The key and the other columns are two writes; commit both or neither + return r.inTx(func(tx *playerRepository) error { + if err := tx.SetAPIKey(ctx, id, *t.APIKey); err != nil { + return err + } + return tx.updateOwned(ctx, id, t, cols...) + }) +} + +func (r *playerRepository) inTx(block func(tx *playerRepository) error) error { + conn, ok := r.db.(*dbx.DB) + if !ok { + return block(r) // already inside a transaction + } + return conn.Transactional(func(tx *dbx.Tx) error { + return block(NewPlayerRepository(tx).(*playerRepository)) + }) } func (r *playerRepository) Delete(ctx context.Context, ids ...string) error { return r.deleteOwnedAll(ctx, ids...) } +// Keys are long random strings, not user-chosen passwords, so a fast unsalted hash is enough and keeps lookups indexed. +func hashAPIKey(key string) string { + sum := sha256.Sum256([]byte(key)) + return hex.EncodeToString(sum[:]) +} + +func (r *playerRepository) FindByAPIKey(ctx context.Context, key string) (*model.Player, error) { + sel := r.selectPlayer(ctx).Where(Eq{"player.api_key_hash": hashAPIKey(key)}) + var res model.Player + if err := r.queryOne(ctx, sel, &res); err != nil { + return nil, err + } + return &res, nil +} + +// SetAPIKey stores the key's hash, or revokes it when key is empty. Setting is owner-only, even for +// admins, so nobody can mint a login for someone else. +func (r *playerRepository) SetAPIKey(ctx context.Context, playerID, key string) error { + if key == "" { + return r.updateOwnedRow(ctx, playerID, ownerOrAdmin, map[string]any{"api_key_hash": nil}) + } + if err := validateAPIKey(key); err != nil { + return err + } + err := r.updateOwnedRow(ctx, playerID, ownerOnly, map[string]any{"api_key_hash": hashAPIKey(key)}) + if isUniqueViolation(err) { + return apiKeyValidationError("ra.validation.unique") + } + return err +} + +func isUniqueViolation(err error) bool { + return err != nil && strings.Contains(err.Error(), "UNIQUE constraint failed") +} + var _ model.PlayerRepository = (*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 f12f3e74e..69afa9556 100644 --- a/persistence/player_repository_test.go +++ b/persistence/player_repository_test.go @@ -2,6 +2,7 @@ package persistence import ( "context" + "errors" "github.com/deluan/rest" "github.com/navidrome/navidrome/log" @@ -12,6 +13,14 @@ import ( "github.com/pocketbase/dbx" ) +const testAPIKey = "nds_0123456789abcdefghijkl" + +func expectAPIKeyError(err error, msg string) { + var verr *rest.ValidationError + ExpectWithOffset(1, errors.As(err, &verr)).To(BeTrue()) + ExpectWithOffset(1, verr.Errors).To(HaveKeyWithValue("apiKey", msg)) +} + var _ = Describe("PlayerRepository", func() { var adminRepo *playerRepository var database *dbx.DB @@ -178,11 +187,12 @@ var _ = Describe("PlayerRepository", func() { clone := player clone.ID = "" clone.IP = "192.168.1.1" + clone.APIKey = new(testAPIKey) id, err := repo.Save(repoCtx, &clone) if clone.UserId == "" { Expect(err).To(HaveOccurred()) - } else if !admin && player.Username == adminPlayer1.Username { + } else if player.UserId != userPlayer.UserId { Expect(err).To(Equal(rest.ErrPermissionDenied)) clone.UserId = "" } else { @@ -202,12 +212,13 @@ var _ = Describe("PlayerRepository", func() { } else { Expect(count).To(Equal(baseCount + 1)) Expect(err).To(BeNil()) + clone.APIKey = nil + clone.HasAPIKey = true Expect(*newItem).To(Equal(clone)) } }, Entry("same user", userPlayer), Entry("other item", otherPlayer), - Entry("fake item", model.Player{}), ) }) @@ -251,6 +262,259 @@ var _ = Describe("PlayerRepository", func() { Entry("regular context", false, model.Players{regularPlayer}, regularPlayer, adminPlayer1), ) + Describe("API keys", func() { + const key = testAPIKey + const otherKey = "nds_ABCDEFGHIJKLMNOPQRSTUV" + var ownerCtx, otherCtx context.Context + + BeforeEach(func() { + ownerCtx = request.WithUser(log.NewContext(GinkgoT().Context()), regularUser) + otherCtx = request.WithUser(log.NewContext(GinkgoT().Context()), thirdUser) + }) + + storedHash := func(id string) string { + var row struct { + Hash string `db:"api_key_hash"` + } + Expect(database.NewQuery("select coalesce(api_key_hash, '') as api_key_hash from player where id = {:id}"). + Bind(dbx.Params{"id": id}).One(&row)).To(Succeed()) + return row.Hash + } + + Describe("SetAPIKey", func() { + It("stores only the hash and finds the player by the key", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + + Expect(storedHash(regularPlayer.ID)).To(Equal(hashAPIKey(key))) + plr, err := adminRepo.FindByAPIKey(ctx, key) + Expect(err).ToNot(HaveOccurred()) + Expect(plr.ID).To(Equal(regularPlayer.ID)) + Expect(plr.HasAPIKey).To(BeTrue()) + Expect(plr.APIKey).To(BeNil()) + }) + + It("replaces the previous key", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, otherKey)).To(Succeed()) + + _, err := adminRepo.FindByAPIKey(ctx, key) + Expect(err).To(MatchError(model.ErrNotFound)) + _, err = adminRepo.FindByAPIKey(ctx, otherKey) + Expect(err).ToNot(HaveOccurred()) + }) + + DescribeTable("rejects malformed keys", + func(bad string) { + err := adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, bad) + expectAPIKeyError(err, "resources.player.validation.apiKeyFormat") + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + }, + Entry("no prefix", "0123456789abcdefghijklmn"), + Entry("too short", "nds_short"), + Entry("too long", key+"x"), + Entry("bad chars", "nds_0123456789abcdefghij-!"), + ) + + It("revokes with an empty key, by the owner or an admin", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, "")).To(Succeed()) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + Expect(adminRepo.SetAPIKey(ctx, regularPlayer.ID, "")).To(Succeed()) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + }) + + It("accepts revoking a player that has no key", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, "")).To(Succeed()) + }) + + It("does not let an admin set a key on another user's player", func() { + Expect(adminRepo.SetAPIKey(ctx, regularPlayer.ID, key)).To(MatchError(rest.ErrPermissionDenied)) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + }) + + It("does not let another user set or revoke", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + Expect(adminRepo.SetAPIKey(otherCtx, regularPlayer.ID, otherKey)).To(MatchError(rest.ErrPermissionDenied)) + Expect(adminRepo.SetAPIKey(otherCtx, regularPlayer.ID, "")).To(MatchError(rest.ErrPermissionDenied)) + Expect(storedHash(regularPlayer.ID)).To(Equal(hashAPIKey(key))) + }) + + It("returns not found for a missing player", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, "missing", key)).To(MatchError(rest.ErrNotFound)) + Expect(adminRepo.SetAPIKey(ownerCtx, "missing", "")).To(MatchError(rest.ErrNotFound)) + }) + + It("does not find unknown or empty keys", func() { + _, err := adminRepo.FindByAPIKey(ctx, otherKey) + Expect(err).To(MatchError(model.ErrNotFound)) + _, err = adminRepo.FindByAPIKey(ctx, "") + Expect(err).To(MatchError(model.ErrNotFound)) + }) + + It("drops the key with the player", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + Expect(adminRepo.Delete(ownerCtx, regularPlayer.ID)).To(Succeed()) + _, err := adminRepo.FindByAPIKey(ctx, key) + Expect(err).To(MatchError(model.ErrNotFound)) + }) + }) + + Describe("hasApiKey filter", func() { + BeforeEach(func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + }) + + filtered := func(value string) []string { + res, err := adminRepo.ReadAll(ctx, rest.QueryOptions{Filters: map[string]any{"hasApiKey": value}}) + Expect(err).ToNot(HaveOccurred()) + var ids []string + for _, p := range res { + ids = append(ids, p.ID) + } + return ids + } + + It("lists only players with a key", func() { + Expect(filtered("true")).To(ConsistOf(regularPlayer.ID)) + count, err := adminRepo.Count(ctx, rest.QueryOptions{Filters: map[string]any{"hasApiKey": "true"}}) + Expect(err).ToNot(HaveOccurred()) + Expect(count).To(Equal(int64(1))) + }) + + It("lists only players without a key", func() { + Expect(filtered("false")).To(ConsistOf(adminPlayer1.ID, adminPlayer2.ID)) + }) + }) + + Describe("Save (create)", func() { + It("creates the player with the key, owned by the logged-in user", func() { + id, err := adminRepo.Save(ownerCtx, &model.Player{Name: "Manual player", APIKey: new(key)}) + Expect(err).ToNot(HaveOccurred()) + + plr, err := adminRepo.FindByAPIKey(ctx, key) + Expect(err).ToNot(HaveOccurred()) + Expect(plr.ID).To(Equal(id)) + Expect(plr.UserId).To(Equal(regularUser.ID)) + }) + + It("requires a key", func() { + count, _ := adminRepo.CountAll(ctx) + _, err := adminRepo.Save(ownerCtx, &model.Player{Name: "No key"}) + expectAPIKeyError(err, "ra.validation.required") + + _, err = adminRepo.Save(ownerCtx, &model.Player{Name: "Empty key", APIKey: new("")}) + expectAPIKeyError(err, "ra.validation.required") + Expect(adminRepo.CountAll(ctx)).To(Equal(count)) + }) + + It("rejects a malformed key without creating the player", func() { + count, _ := adminRepo.CountAll(ctx) + _, err := adminRepo.Save(ownerCtx, &model.Player{Name: "Bad", APIKey: new("nds_bad")}) + expectAPIKeyError(err, "resources.player.validation.apiKeyFormat") + Expect(adminRepo.CountAll(ctx)).To(Equal(count)) + }) + + It("does not let an admin create a keyed player for another user", func() { + count, _ := adminRepo.CountAll(ctx) + _, err := adminRepo.Save(ctx, &model.Player{Name: "For someone", UserId: regularUser.ID, APIKey: new(key)}) + Expect(err).To(MatchError(rest.ErrPermissionDenied)) + _, err = adminRepo.Save(ctx, &model.Player{Name: "For someone", UserId: regularUser.ID}) + Expect(err).To(MatchError(rest.ErrPermissionDenied)) + Expect(adminRepo.CountAll(ctx)).To(Equal(count)) + }) + + It("rejects a key already used by another player without creating the player", func() { + Expect(adminRepo.SetAPIKey(ctx, adminPlayer1.ID, key)).To(Succeed()) + count, _ := adminRepo.CountAll(ctx) + _, err := adminRepo.Save(ownerCtx, &model.Player{Name: "Duplicate", APIKey: new(key)}) + expectAPIKeyError(err, "ra.validation.unique") + Expect(adminRepo.CountAll(ctx)).To(Equal(count)) + }) + }) + + Describe("Update (edit)", func() { + It("rolls back the key change when the rest of the edit fails", func() { + _, err := database.NewQuery(`create trigger fail_player_rename before update of name on player + when new.name = 'boom' begin select raise(abort, 'boom'); end`).Execute() + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { + _, _ = database.NewQuery("drop trigger if exists fail_player_rename").Execute() + }) + + plr := regularPlayer + plr.Name = "boom" + plr.APIKey = new(key) + Expect(adminRepo.Update(ownerCtx, plr.ID, plr, "name", "apiKey")).ToNot(Succeed()) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + }) + + It("keeps the key when apiKey is absent (a normal edit)", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + + plr := regularPlayer + plr.Name = "Renamed" + Expect(adminRepo.Update(ownerCtx, plr.ID, plr, "name", "hasApiKey")).To(Succeed()) + Expect(adminRepo.Update(ownerCtx, plr.ID, plr)).To(Succeed()) + + found, err := adminRepo.FindByAPIKey(ctx, key) + Expect(err).ToNot(HaveOccurred()) + Expect(found.Name).To(Equal("Renamed")) + }) + + It("sets a new key when apiKey has a value", func() { + plr := regularPlayer + plr.APIKey = new(key) + Expect(adminRepo.Update(ownerCtx, plr.ID, plr, "name", "apiKey")).To(Succeed()) + Expect(storedHash(regularPlayer.ID)).To(Equal(hashAPIKey(key))) + }) + + It("revokes the key when apiKey is empty", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + plr := regularPlayer + plr.APIKey = new("") + Expect(adminRepo.Update(ownerCtx, plr.ID, plr, "apiKey")).To(Succeed()) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + }) + + It("lets an admin edit another user's keyed player without touching the key", func() { + Expect(adminRepo.SetAPIKey(ownerCtx, regularPlayer.ID, key)).To(Succeed()) + plr := regularPlayer + plr.MaxBitRate = 192 + Expect(adminRepo.Update(ctx, plr.ID, plr, "maxBitRate", "hasApiKey")).To(Succeed()) + Expect(storedHash(regularPlayer.ID)).To(Equal(hashAPIKey(key))) + }) + + It("refuses an admin setting a key on another user's player and leaves other columns alone", func() { + plr := regularPlayer + plr.Name = "Hijacked" + plr.APIKey = new(key) + Expect(adminRepo.Update(ctx, plr.ID, plr, "name", "apiKey")).To(MatchError(rest.ErrPermissionDenied)) + + got, err := adminRepo.Get(ctx, regularPlayer.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(got.Name).To(Equal(regularPlayer.Name)) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + }) + + It("refuses a key already used by another player and leaves other columns alone", func() { + Expect(adminRepo.SetAPIKey(ctx, adminPlayer1.ID, key)).To(Succeed()) + plr := regularPlayer + plr.Name = "Renamed" + plr.APIKey = new(key) + err := adminRepo.Update(ownerCtx, plr.ID, plr, "name", "apiKey") + expectAPIKeyError(err, "ra.validation.unique") + + got, err := adminRepo.Get(ctx, regularPlayer.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(got.Name).To(Equal(regularPlayer.Name)) + Expect(storedHash(regularPlayer.ID)).To(BeEmpty()) + Expect(storedHash(adminPlayer1.ID)).To(Equal(hashAPIKey(key))) + }) + }) + }) + Describe("Ownership enforcement (cross-tenant write protection)", func() { var regularRepo *playerRepository var regularCtx context.Context @@ -287,6 +551,7 @@ var _ = Describe("PlayerRepository", func() { Name: "HIJACKED", UserId: regularUser.ID, ReportRealPath: true, + APIKey: new(testAPIKey), } id, err := regularRepo.Save(regularCtx, &spoofed) diff --git a/persistence/sql_base_repository.go b/persistence/sql_base_repository.go index be88156d8..03cc6a01b 100644 --- a/persistence/sql_base_repository.go +++ b/persistence/sql_base_repository.go @@ -85,6 +85,22 @@ func (r sqlRepository) addRestriction(ctx context.Context, sql ...Sqlizer) Sqliz return s } +// writeAccess says who may change a row in a table with a user_id column. +type writeAccess int + +const ( + ownerOrAdmin writeAccess = iota // admins may write any row + ownerOnly // even admins may only write their own rows +) + +// ownedRow matches the row rowID only if the logged-in user may write it under access. +func (r sqlRepository) ownedRow(ctx context.Context, rowID string, access writeAccess) Sqlizer { + if access == ownerOnly { + return And{Eq{"id": rowID}, Eq{"user_id": loggedUser(ctx).ID}} + } + return r.addRestriction(ctx, Eq{"id": rowID}) +} + func (r *sqlRepository) registerModel(instance any, filters map[string]filterFunc) { if r.tableName == "" { r.tableName = strings.TrimPrefix(reflect.TypeOf(instance).String(), "*model.") @@ -494,15 +510,12 @@ func (r sqlRepository) updateOwned(ctx context.Context, id string, m any, colsTo } updateValues := filterUpdateValues(values, id, colsToUpdate...) delete(updateValues, "user_id") // ownership is immutable on 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(ctx, id) - } - return nil + return r.updateOwnedRow(ctx, id, ownerOrAdmin, updateValues) +} + +// updateOwnedRow sets values on the row rowID if the logged-in user may write it under access. +func (r sqlRepository) updateOwnedRow(ctx context.Context, rowID string, access writeAccess, values map[string]any) error { + return r.runRowWrite(ctx, rowID, Update(r.tableName).SetMap(values).Where(r.ownedRow(ctx, rowID, access))) } // deleteOwned performs an atomic, ownership-restricted delete of the row identified by id, for @@ -511,12 +524,17 @@ func (r sqlRepository) updateOwned(ctx context.Context, id string, m any, colsTo // 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(ctx context.Context, id string) error { - count, err := r.executeSQL(ctx, Delete(r.tableName).Where(r.addRestriction(ctx, Eq{"id": id}))) + return r.runRowWrite(ctx, id, Delete(r.tableName).Where(r.ownedRow(ctx, id, ownerOrAdmin))) +} + +// runRowWrite executes q, a write already filtered by ownedRow(rowID, …), and classifies a miss. +func (r sqlRepository) runRowWrite(ctx context.Context, rowID string, q Sqlizer) error { + count, err := r.executeSQL(ctx, q) if err != nil { return err } if count == 0 { - return r.classifyOwnedWriteMiss(ctx, id) + return r.classifyOwnedWriteMiss(ctx, rowID) } return nil } diff --git a/persistence/sql_base_repository_test.go b/persistence/sql_base_repository_test.go index 33a8140f8..4b42e7f1d 100644 --- a/persistence/sql_base_repository_test.go +++ b/persistence/sql_base_repository_test.go @@ -19,6 +19,26 @@ var _ = Describe("sqlRepository", func() { r.tableName = "table" }) + Describe("ownedRow", func() { + DescribeTable("matches the row, limited to what the logged-in user may write", + func(user model.User, access writeAccess, expectedSQL string, expectedArgs ...any) { + userCtx := request.WithUser(GinkgoT().Context(), user) + sql, args, err := r.ownedRow(userCtx, "row-1", access).ToSql() + Expect(err).ToNot(HaveOccurred()) + Expect(sql).To(Equal(expectedSQL)) + Expect(args).To(Equal(expectedArgs)) + }, + Entry("admin, ownerOrAdmin: any row", model.User{ID: "admin", IsAdmin: true}, ownerOrAdmin, + "(id = ?)", "row-1"), + Entry("regular, ownerOrAdmin: own rows", model.User{ID: "user"}, ownerOrAdmin, + "(id = ? AND user_id = ?)", "row-1", "user"), + Entry("admin, ownerOnly: own rows", model.User{ID: "admin", IsAdmin: true}, ownerOnly, + "(id = ? AND user_id = ?)", "row-1", "admin"), + Entry("regular, ownerOnly: own rows", model.User{ID: "user"}, ownerOnly, + "(id = ? AND user_id = ?)", "row-1", "user"), + ) + }) + Describe("applyOptions", func() { var sq squirrel.SelectBuilder BeforeEach(func() { diff --git a/resources/i18n/pt-br.json b/resources/i18n/pt-br.json index 6fcde56ca..f0d8a0e06 100644 --- a/resources/i18n/pt-br.json +++ b/resources/i18n/pt-br.json @@ -182,6 +182,7 @@ }, "player": { "name": "Tocador |||| Tocadores", + "menuName": "Tocadores e chaves de API", "fields": { "name": "Nome", "transcodingId": "Conversão", @@ -190,7 +191,29 @@ "userName": "Usuário", "lastSeen": "Últ. acesso", "reportRealPath": "Use paths reais", - "scrobbleEnabled": "Enviar scrobbles para serviços externos" + "scrobbleEnabled": "Enviar scrobbles para serviços externos", + "hasApiKey": "Chave de API" + }, + "actions": { + "generateApiKey": "Gerar chave de API", + "regenerateApiKey": "Gerar nova", + "revokeApiKey": "Revogar", + "copyApiKey": "Copiar" + }, + "message": { + "apiKeyActive": "Este tocador tem uma chave de API. Use-a no seu app como chave de API, ou como senha se o app não usar autenticação por token.", + "apiKeyNone": "Sem chave de API. Gere uma para conectar um app a este tocador.", + "apiKeyNoneOther": "Sem chave de API.", + "apiKeyPending": "Copie esta chave agora. Ela será salva quando você clicar em Salvar e não será exibida novamente.", + "apiKeyRevokePending": "A chave de API será removida quando você salvar.", + "deleteWithKeyTitle": "Excluir tocador", + "deleteWithKeyContent": "Este tocador tem uma chave de API. Os apps que a usam vão parar de funcionar." + }, + "notifications": { + "apiKeyCopied": "Chave de API copiada para o clipboard" + }, + "validation": { + "apiKeyFormat": "Formato de chave de API inválido" } }, "transcoding": { diff --git a/server/subsonic/api.go b/server/subsonic/api.go index deedc46c7..fe724741c 100644 --- a/server/subsonic/api.go +++ b/server/subsonic/api.go @@ -107,6 +107,7 @@ func (api *Router) routes() http.Handler { r.Use(getPlayer(api.players)) h(r, "ping", api.Ping) h(r, "getLicense", api.GetLicense) + h(r, "tokenInfo", api.TokenInfo) }) r.Group(func(r chi.Router) { r.Use(getPlayer(api.players)) diff --git a/server/subsonic/e2e/subsonic_apikey_test.go b/server/subsonic/e2e/subsonic_apikey_test.go new file mode 100644 index 000000000..bb263c15d --- /dev/null +++ b/server/subsonic/e2e/subsonic_apikey_test.go @@ -0,0 +1,54 @@ +package e2e + +import ( + "net/http/httptest" + "net/url" + + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/server/subsonic/responses" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("API key authentication", func() { + var key string + + BeforeEach(func() { + setupTestDB() + userCtx := request.WithUser(ctx, regularUser) + player := &model.Player{ID: "apikey-player", Name: "Phone", UserId: regularUser.ID, Client: "test-client"} + Expect(ds.Player().Put(userCtx, player)).To(Succeed()) + key = "nds_0123456789abcdefghijkl" + Expect(ds.Player().SetAPIKey(userCtx, player.ID, key)).To(Succeed()) + }) + + doKeyReq := func(endpoint, apiKey string) *responses.Subsonic { + q := url.Values{"apiKey": {apiKey}, "v": {"1.16.1"}, "c": {"test-client"}, "f": {"json"}} + w := httptest.NewRecorder() + router.ServeHTTP(w, httptest.NewRequest("GET", "/"+endpoint+"?"+q.Encode(), nil)) + return parseJSONResponse(w) + } + + It("authenticates ping with only the key", func() { + resp := doKeyReq("ping", key) + + Expect(resp.Status).To(Equal(responses.StatusOK)) + }) + + It("reports the key owner in tokenInfo", func() { + resp := doKeyReq("tokenInfo", key) + + Expect(resp.Status).To(Equal(responses.StatusOK)) + Expect(resp.TokenInfo).ToNot(BeNil()) + Expect(resp.TokenInfo.Username).To(Equal(regularUser.UserName)) + }) + + It("rejects an unknown key with error 44", func() { + resp := doKeyReq("ping", "nds_unknown") + + Expect(resp.Status).To(Equal(responses.StatusFailed)) + Expect(resp.Error).ToNot(BeNil()) + Expect(resp.Error.Code).To(Equal(int32(44))) + }) +}) diff --git a/server/subsonic/middlewares.go b/server/subsonic/middlewares.go index 6617661a9..fdb4af28d 100644 --- a/server/subsonic/middlewares.go +++ b/server/subsonic/middlewares.go @@ -65,14 +65,15 @@ func checkRequiredParameters(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { var requiredParameters []string + p := req.Params(r) username, _ := fromInternalOrProxyAuth(r) - if username != "" { + apiKey, _ := p.String("apiKey") + if username != "" || apiKey != "" { requiredParameters = []string{"v", "c"} } else { requiredParameters = []string{"u", "v", "c"} } - p := req.Params(r) for _, param := range requiredParameters { if _, err := p.String(param); err != nil { log.Warn(r, err) @@ -104,10 +105,14 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { ctx := r.Context() var usr *model.User + var keyPlayer *model.Player var err error + p := req.Params(r) + apiKey, _ := p.String("apiKey") username, isInternalAuth := fromInternalOrProxyAuth(r) - if username != "" { + switch { + case username != "": authType := If(isInternalAuth, "internal", "reverse-proxy") usr, err = ds.User().FindByUsername(ctx, username) if errors.Is(err, context.Canceled) { @@ -119,8 +124,16 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { } else if err != nil { log.Error(ctx, "API: Error authenticating username", "auth", authType, "username", username, "remoteAddr", r.RemoteAddr, err) } - } else { - p := req.Params(r) + case apiKey != "": + usr, keyPlayer, err = authenticateAPIKey(ctx, ds, limiter, r, apiKey) + if err != nil { + if ctx.Err() == nil { + sendError(w, r, err) + } + return + } + ctx = request.WithUsername(ctx, usr.UserName) + default: username, _ := p.String("u") pass, _ := p.String("p") token, _ := p.String("t") @@ -142,6 +155,9 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { usr, err = ds.User().FindByUsernameWithPassword(ctx, username) if err == nil { err = validateCredentials(usr, pass, token, salt, jwt) + if errors.Is(err, model.ErrInvalidAuth) && pass != "" && jwt == "" { + keyPlayer, err = playerFromPasswordKey(ctx, ds, usr, pass) + } } invalidLogin := errors.Is(err, model.ErrNotFound) || errors.Is(err, model.ErrInvalidAuth) slot.release(invalidLogin) @@ -162,11 +178,77 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { } ctx = request.WithUser(ctx, *usr) + if keyPlayer != nil { + ctx = request.WithPlayer(ctx, *keyPlayer) + } next.ServeHTTP(w, r.WithContext(ctx)) }) } } +var apiKeyConflicts = []string{"u", "p", "t", "s", "jwt"} + +func authenticateAPIKey(ctx context.Context, ds model.DataStore, limiter *authLimiter, r *http.Request, key string) (*model.User, *model.Player, error) { + query := r.URL.Query() + for _, param := range apiKeyConflicts { + if query.Has(param) { + log.Warn(ctx, "API: apiKey sent with other credentials", "auth", "apikey", "param", param, "remoteAddr", r.RemoteAddr) + return nil, nil, newError(responses.ErrorMultipleAuthMechanismsProvided) + } + } + + // Per key, so a stale key on one device cannot lock out valid keys sharing the IP + slot, allowed := limiter.acquire(ctx, "apikey\x00"+server.ClientIP(r)+"\x00"+key) + if !allowed { + if err := ctx.Err(); err != nil { + return nil, nil, err + } + log.Warn(ctx, "API: Too many failed API key attempts", "auth", "apikey", "remoteAddr", r.RemoteAddr) + return nil, nil, newError(responses.ErrorInvalidAPIKey) + } + + player, err := ds.Player().FindByAPIKey(ctx, key) + var usr *model.User + if err == nil { + usr, err = ds.User().Get(ctx, player.UserId) + } + slot.release(errors.Is(err, model.ErrNotFound)) + switch { + case errors.Is(err, context.Canceled): + return nil, nil, err + case errors.Is(err, model.ErrNotFound): + log.Warn(ctx, "API: Invalid API key", "auth", "apikey", "remoteAddr", r.RemoteAddr) + return nil, nil, newError(responses.ErrorInvalidAPIKey) + case err != nil: + log.Error(ctx, "API: Error authenticating API key", "auth", "apikey", "remoteAddr", r.RemoteAddr, err) + return nil, nil, newError(responses.ErrorAuthenticationFail) + } + return usr, player, nil +} + +// playerFromPasswordKey lets clients that only have a password field log in with an API key. +// It returns ErrInvalidAuth when pass is not a key of usr, so only real failures skip the limiter count. +func playerFromPasswordKey(ctx context.Context, ds model.DataStore, usr *model.User, pass string) (*model.Player, error) { + key := decodePassword(pass) + if !strings.HasPrefix(key, consts.APIKeyPrefix) { + return nil, model.ErrInvalidAuth + } + plr, err := ds.Player().FindByAPIKey(ctx, key) + if errors.Is(err, model.ErrNotFound) || (err == nil && plr.UserId != usr.ID) { + return nil, model.ErrInvalidAuth + } + return plr, err +} + +func decodePassword(pass string) string { + if strings.HasPrefix(pass, "enc:") { + if dec, err := hex.DecodeString(pass[4:]); err == nil { + return string(dec) + } + } + return pass +} + func adminOnly(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { loggedUser, ok := request.UserFrom(r.Context()) @@ -194,12 +276,7 @@ func validateCredentials(user *model.User, pass, token, salt, jwt string) error claims.Subject == user.UserName && auth.CheckClaims(claims, *user, auth.AudienceSubsonic) == nil case pass != "": - if strings.HasPrefix(pass, "enc:") { - if dec, err := hex.DecodeString(pass[4:]); err == nil { - pass = string(dec) - } - } - valid = pass == user.Password + valid = decodePassword(pass) == user.Password case token != "": t := fmt.Sprintf("%x", md5.Sum([]byte(user.Password+salt))) valid = t == token @@ -217,12 +294,20 @@ func getPlayer(players core.Players) func(next http.Handler) http.Handler { ctx := r.Context() userName, _ := request.UsernameFrom(ctx) client, _ := request.ClientFrom(ctx) - playerId := playerIDFromCookie(r, userName) ip, _, _ := net.SplitHostPort(r.RemoteAddr) userAgent := canonicalUserAgent(r) - player, trc, err := players.Register(ctx, playerId, client, userAgent, ip) + + var player *model.Player + var trc *model.Transcoding + var err error + keyPlayer, boundByKey := request.PlayerFrom(ctx) + if boundByKey { + player, trc, err = players.Touch(ctx, keyPlayer, client, userAgent, ip) + } else { + player, trc, err = players.Register(ctx, playerIDFromCookie(r, userName), client, userAgent, ip) + } if err != nil { - log.Error(ctx, "Could not register player", "username", userName, "client", client, err) + log.Error(ctx, "Could not resolve player", "username", userName, "client", client, err) } else { ctx = request.WithPlayer(ctx, *player) if trc != nil { @@ -230,6 +315,11 @@ func getPlayer(players core.Players) func(next http.Handler) http.Handler { } r = r.WithContext(ctx) + // A key already identifies the player, so the cookie would only add a second, weaker signal + if boundByKey { + next.ServeHTTP(w, r) + return + } cookie := &http.Cookie{ //nolint:gosec // Secure omitted: Navidrome may run over plain HTTP Name: playerIDCookieName(userName), Value: player.ID, diff --git a/server/subsonic/middlewares_test.go b/server/subsonic/middlewares_test.go index 0879ee540..b3ad972b7 100644 --- a/server/subsonic/middlewares_test.go +++ b/server/subsonic/middlewares_test.go @@ -3,6 +3,7 @@ package subsonic import ( "context" "crypto/md5" + "encoding/hex" "errors" "fmt" "net/http" @@ -119,6 +120,14 @@ var _ = Describe("Middlewares", func() { Expect(next.called).To(BeTrue()) }) + It("does not require u when apiKey is present", func() { + r := newGetRequest("apiKey=nds_abc", "v=1.15", "c=test") + cp := checkRequiredParameters(next) + cp.ServeHTTP(w, r) + + Expect(next.called).To(BeTrue()) + }) + It("fails when user is missing", func() { r := newGetRequest("v=1.15", "c=test") cp := checkRequiredParameters(next) @@ -311,6 +320,101 @@ var _ = Describe("Middlewares", func() { }) }) + When("using API key authentication", func() { + var key string + serve := func(params ...string) { + authenticate(ds)(next).ServeHTTP(w, newGetRequest(params...)) + } + + BeforeEach(func() { + usr, err := ds.User().FindByUsername(ctx, "admin") + Expect(err).ToNot(HaveOccurred()) + Expect(ds.Player().Put(ctx, &model.Player{ID: "player-1", Name: "My Phone", UserId: usr.ID, Client: "Symfonium"})).To(Succeed()) + key = "nds_0123456789abcdefghijkl" + Expect(ds.Player().SetAPIKey(ctx, "player-1", key)).To(Succeed()) + }) + + It("authenticates the owner and binds the key's player", func() { + serve("apiKey=" + key) + + Expect(next.called).To(BeTrue()) + user, _ := request.UserFrom(next.req.Context()) + Expect(user.UserName).To(Equal("admin")) + username, _ := request.UsernameFrom(next.req.Context()) + Expect(username).To(Equal("admin")) + player, ok := request.PlayerFrom(next.req.Context()) + Expect(ok).To(BeTrue()) + Expect(player.ID).To(Equal("player-1")) + }) + + It("accepts the key in a POST form body", func() { + r := newPostRequest("", "apiKey="+key) + cp := postFormToQueryParams(authenticate(ds)(next)) + cp.ServeHTTP(w, r) + + Expect(next.called).To(BeTrue()) + player, _ := request.PlayerFrom(next.req.Context()) + Expect(player.ID).To(Equal("player-1")) + }) + + It("rejects an unknown key with error 44", func() { + serve("apiKey=nds_unknown") + + Expect(w.Body.String()).To(ContainSubstring(`code="44"`)) + Expect(next.called).To(BeFalse()) + }) + + DescribeTable("rejects apiKey mixed with other credentials with error 43", + func(extra string) { + serve("apiKey="+key, extra) + + Expect(w.Body.String()).To(ContainSubstring(`code="43"`)) + Expect(next.called).To(BeFalse()) + }, + Entry("u", "u=admin"), + Entry("p", "p=wordpass"), + Entry("t", "t=abc"), + Entry("s", "s=abc"), + Entry("jwt", "jwt=abc"), + Entry("empty u", "u="), + Entry("empty p", "p="), + ) + + Context("key sent as the password", func() { + It("authenticates and binds the key's player", func() { + serve("u=admin", "p="+key) + + Expect(next.called).To(BeTrue()) + player, ok := request.PlayerFrom(next.req.Context()) + Expect(ok).To(BeTrue()) + Expect(player.ID).To(Equal("player-1")) + }) + + It("accepts the hex-encoded form", func() { + serve("u=admin", "p=enc:"+hex.EncodeToString([]byte(key))) + + Expect(next.called).To(BeTrue()) + }) + + It("still accepts a real password that starts with the key prefix", func() { + Expect(ds.User().Put(ctx, &model.User{UserName: "prefixed", NewPassword: "nds_secret"})).To(Succeed()) + serve("u=prefixed", "p=nds_secret") + + Expect(next.called).To(BeTrue()) + _, ok := request.PlayerFrom(next.req.Context()) + Expect(ok).To(BeFalse()) + }) + + It("rejects another user's key with error 40", func() { + Expect(ds.User().Put(ctx, &model.User{UserName: "other", NewPassword: "pw"})).To(Succeed()) + serve("u=other", "p="+key) + + Expect(w.Body.String()).To(ContainSubstring(`code="40"`)) + Expect(next.called).To(BeFalse()) + }) + }) + }) + When("failed attempts reach AuthRequestLimit", func() { var cp http.Handler @@ -376,6 +480,21 @@ var _ = Describe("Middlewares", func() { Expect(next.called).To(BeTrue()) }) + It("does not count server errors when a key is sent as the password", func() { + usr, _ := ds.User().FindByUsername(ctx, "admin") + playerRepo := ds.Player().(*tests.MockPlayerRepo) + Expect(playerRepo.Put(ctx, &model.Player{ID: "player-1", UserId: usr.ID})).To(Succeed()) + key := "nds_0123456789abcdefghijkl" + Expect(playerRepo.SetAPIKey(ctx, "player-1", key)).To(Succeed()) + + playerRepo.Error = errors.New("db down") + failTimes(5, "u=admin", "p="+key) + playerRepo.Error = nil + + serve(newGetRequest("u=admin", "p="+key)) + Expect(next.called).To(BeTrue()) + }) + It("does not block other usernames from the same IP", func() { _ = ds.User().Put(ctx, &model.User{UserName: "other", NewPassword: "otherpass"}) failTimes(3, "u=admin", "p=WRONG") @@ -405,6 +524,25 @@ var _ = Describe("Middlewares", func() { Expect(next.called).To(BeTrue()) }) + It("throttles a repeated bad key without locking out valid keys from the same IP", func() { + usr, _ := ds.User().FindByUsername(ctx, "admin") + playerRepo := ds.Player().(*tests.MockPlayerRepo) + Expect(playerRepo.Put(ctx, &model.Player{ID: "player-1", UserId: usr.ID})).To(Succeed()) + key := "nds_0123456789abcdefghijkl" + Expect(playerRepo.SetAPIKey(ctx, "player-1", key)).To(Succeed()) + + for range 3 { + Expect(serve(newGetRequest("apiKey=nds_bad")).Body.String()).To(ContainSubstring(`code="44"`)) + } + playerRepo.APIKeys["nds_bad"] = "player-1" + rec := serve(newGetRequest("apiKey=nds_bad")) + Expect(next.called).To(BeFalse()) + Expect(rec.Body.String()).To(ContainSubstring(`code="44"`)) + + serve(newGetRequest("apiKey=" + key)) + Expect(next.called).To(BeTrue()) + }) + It("is disabled when AuthRequestLimit is 0", func() { conf.Server.AuthRequestLimit = 0 cp = authenticate(ds)(next) @@ -532,6 +670,24 @@ var _ = Describe("Middlewares", func() { Expect(cookieStr).To(BeEmpty()) }) + Context("player bound by an API key", func() { + BeforeEach(func() { + r = r.WithContext(request.WithPlayer(r.Context(), model.Player{ID: "keyed"})) + gp := getPlayer(mockedPlayers)(next) + gp.ServeHTTP(w, r) + }) + + It("uses the key's player", func() { + Expect(mockedPlayers.touched).To(BeTrue()) + player, _ := request.PlayerFrom(next.req.Context()) + Expect(player.ID).To(Equal("keyed")) + }) + + It("does not set the player cookie", func() { + Expect(w.Header().Get("Set-Cookie")).To(BeEmpty()) + }) + }) + Context("PlayerId specified in Cookies", func() { BeforeEach(func() { cookie := &http.Cookie{ @@ -712,6 +868,12 @@ func (mh *mockHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { type mockPlayers struct { core.Players transcoding *model.Transcoding + touched bool +} + +func (mp *mockPlayers) Touch(_ context.Context, plr model.Player, _, _, _ string) (*model.Player, *model.Transcoding, error) { + mp.touched = true + return &plr, mp.transcoding, nil } func (mp *mockPlayers) Get(ctx context.Context, playerId string) (*model.Player, error) { diff --git a/server/subsonic/opensubsonic.go b/server/subsonic/opensubsonic.go index 2b2a31bf3..21c407039 100644 --- a/server/subsonic/opensubsonic.go +++ b/server/subsonic/opensubsonic.go @@ -16,6 +16,7 @@ func (api *Router) GetOpenSubsonicExtensions(_ *http.Request) (*responses.Subson {Name: "transcoding", Versions: []int32{1}}, {Name: "playbackReport", Versions: []int32{1}}, {Name: "topSongsByArtistId", Versions: []int32{1}}, + {Name: "apiKeyAuthentication", Versions: []int32{1}}, } if api.sonic != nil && api.sonic.HasProvider() { extensions = append(extensions, responses.OpenSubsonicExtension{ diff --git a/server/subsonic/opensubsonic_test.go b/server/subsonic/opensubsonic_test.go index 2615a652d..8740c9971 100644 --- a/server/subsonic/opensubsonic_test.go +++ b/server/subsonic/opensubsonic_test.go @@ -44,44 +44,13 @@ var _ = Describe("GetOpenSubsonicExtensions", func() { router = subsonic.New(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) }) - It("should return the base 6 OpenSubsonicExtensions without sonicSimilarity", func() { + It("should return the base 8 OpenSubsonicExtensions without sonicSimilarity", func() { router.ServeHTTP(w, r) // Make sure the endpoint is public, by not passing any authentication Expect(w.Code).To(Equal(http.StatusOK)) Expect(w.Header().Get("Content-Type")).To(Equal("application/json")) - var response responses.JsonWrapper - err := json.Unmarshal(w.Body.Bytes(), &response) - Expect(err).NotTo(HaveOccurred()) - Expect(*response.Subsonic.OpenSubsonicExtensions).To(SatisfyAll( - HaveLen(7), - ContainElement(responses.OpenSubsonicExtension{Name: "transcodeOffset", Versions: []int32{1}}), - ContainElement(responses.OpenSubsonicExtension{Name: "formPost", Versions: []int32{1}}), - ContainElement(responses.OpenSubsonicExtension{Name: "songLyrics", Versions: []int32{1, 2}}), - ContainElement(responses.OpenSubsonicExtension{Name: "indexBasedQueue", Versions: []int32{1}}), - ContainElement(responses.OpenSubsonicExtension{Name: "transcoding", Versions: []int32{1}}), - ContainElement(responses.OpenSubsonicExtension{Name: "playbackReport", Versions: []int32{1}}), - ContainElement(responses.OpenSubsonicExtension{Name: "topSongsByArtistId", Versions: []int32{1}}), - )) - Expect(*response.Subsonic.OpenSubsonicExtensions).NotTo( - ContainElement(responses.OpenSubsonicExtension{Name: "sonicSimilarity", Versions: []int32{1}}), - ) - }) - }) - - Context("with sonic similarity plugin", func() { - BeforeEach(func() { - sonicService := sonicsvc.New(nil, &mockSonicPluginLoader{names: []string{"test-plugin"}}, nil) - router = subsonic.New(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, sonicService) - }) - - It("should return 7 extensions including sonicSimilarity", func() { - router.ServeHTTP(w, r) - - Expect(w.Code).To(Equal(http.StatusOK)) - Expect(w.Header().Get("Content-Type")).To(Equal("application/json")) - var response responses.JsonWrapper err := json.Unmarshal(w.Body.Bytes(), &response) Expect(err).NotTo(HaveOccurred()) @@ -93,8 +62,41 @@ var _ = Describe("GetOpenSubsonicExtensions", func() { ContainElement(responses.OpenSubsonicExtension{Name: "indexBasedQueue", Versions: []int32{1}}), ContainElement(responses.OpenSubsonicExtension{Name: "transcoding", Versions: []int32{1}}), ContainElement(responses.OpenSubsonicExtension{Name: "playbackReport", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "topSongsByArtistId", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "apiKeyAuthentication", Versions: []int32{1}}), + )) + Expect(*response.Subsonic.OpenSubsonicExtensions).NotTo( + ContainElement(responses.OpenSubsonicExtension{Name: "sonicSimilarity", Versions: []int32{1}}), + ) + }) + }) + + Context("with sonic similarity plugin", func() { + BeforeEach(func() { + sonicService := sonicsvc.New(nil, &mockSonicPluginLoader{names: []string{"test-plugin"}}, nil) + router = subsonic.New(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, sonicService) + }) + + It("should return 9 extensions including sonicSimilarity", func() { + router.ServeHTTP(w, r) + + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).To(Equal("application/json")) + + var response responses.JsonWrapper + err := json.Unmarshal(w.Body.Bytes(), &response) + Expect(err).NotTo(HaveOccurred()) + Expect(*response.Subsonic.OpenSubsonicExtensions).To(SatisfyAll( + HaveLen(9), + ContainElement(responses.OpenSubsonicExtension{Name: "transcodeOffset", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "formPost", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "songLyrics", Versions: []int32{1, 2}}), + ContainElement(responses.OpenSubsonicExtension{Name: "indexBasedQueue", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "transcoding", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "playbackReport", Versions: []int32{1}}), ContainElement(responses.OpenSubsonicExtension{Name: "sonicSimilarity", Versions: []int32{1}}), ContainElement(responses.OpenSubsonicExtension{Name: "topSongsByArtistId", Versions: []int32{1}}), + ContainElement(responses.OpenSubsonicExtension{Name: "apiKeyAuthentication", Versions: []int32{1}}), )) }) }) diff --git a/server/subsonic/responses/.snapshots/Responses TokenInfo should match .JSON b/server/subsonic/responses/.snapshots/Responses TokenInfo should match .JSON new file mode 100644 index 000000000..f2e251f49 --- /dev/null +++ b/server/subsonic/responses/.snapshots/Responses TokenInfo should match .JSON @@ -0,0 +1,10 @@ +{ + "status": "ok", + "version": "1.16.1", + "type": "navidrome", + "serverVersion": "v0.55.0", + "openSubsonic": true, + "tokenInfo": { + "username": "deluan" + } +} diff --git a/server/subsonic/responses/.snapshots/Responses TokenInfo should match .XML b/server/subsonic/responses/.snapshots/Responses TokenInfo should match .XML new file mode 100644 index 000000000..7ea786bb9 --- /dev/null +++ b/server/subsonic/responses/.snapshots/Responses TokenInfo should match .XML @@ -0,0 +1,3 @@ + + + diff --git a/server/subsonic/responses/errors.go b/server/subsonic/responses/errors.go index 42e5427b3..9c9dd10f6 100644 --- a/server/subsonic/responses/errors.go +++ b/server/subsonic/responses/errors.go @@ -1,25 +1,29 @@ package responses const ( - ErrorGeneric int32 = 0 - ErrorMissingParameter int32 = 10 - ErrorClientTooOld int32 = 20 - ErrorServerTooOld int32 = 30 - ErrorAuthenticationFail int32 = 40 - ErrorAuthorizationFail int32 = 50 - ErrorTrialExpired int32 = 60 - ErrorDataNotFound int32 = 70 + ErrorGeneric int32 = 0 + ErrorMissingParameter int32 = 10 + ErrorClientTooOld int32 = 20 + ErrorServerTooOld int32 = 30 + ErrorAuthenticationFail int32 = 40 + ErrorMultipleAuthMechanismsProvided int32 = 43 + ErrorInvalidAPIKey int32 = 44 + ErrorAuthorizationFail int32 = 50 + ErrorTrialExpired int32 = 60 + ErrorDataNotFound int32 = 70 ) -var errors = map[int32]string{ - ErrorGeneric: "A generic error", - ErrorMissingParameter: "Required parameter is missing", - ErrorClientTooOld: "Incompatible Subsonic REST protocol version. Client must upgrade", - ErrorServerTooOld: "Incompatible Subsonic REST protocol version. Server must upgrade", - ErrorAuthenticationFail: "Wrong username or password", - ErrorAuthorizationFail: "User is not authorized for the given operation", - ErrorTrialExpired: "The trial period for the Subsonic server is over. Please upgrade to Subsonic Premium. Visit subsonic.org for details", - ErrorDataNotFound: "The requested data was not found", +var errors = map[int32]string{ //nolint:gosec // G101 false positive: error messages, not credentials + ErrorGeneric: "A generic error", + ErrorMissingParameter: "Required parameter is missing", + ErrorClientTooOld: "Incompatible Subsonic REST protocol version. Client must upgrade", + ErrorServerTooOld: "Incompatible Subsonic REST protocol version. Server must upgrade", + ErrorAuthenticationFail: "Wrong username or password", + ErrorMultipleAuthMechanismsProvided: "Multiple conflicting authentication mechanisms provided", + ErrorInvalidAPIKey: "Invalid API key", + ErrorAuthorizationFail: "User is not authorized for the given operation", + ErrorTrialExpired: "The trial period for the Subsonic server is over. Please upgrade to Subsonic Premium. Visit subsonic.org for details", + ErrorDataNotFound: "The requested data was not found", } func ErrorMsg(code int32) string { diff --git a/server/subsonic/responses/responses.go b/server/subsonic/responses/responses.go index 252eee4c6..51d0020b2 100644 --- a/server/subsonic/responses/responses.go +++ b/server/subsonic/responses/responses.go @@ -63,6 +63,7 @@ type Subsonic struct { PlayQueueByIndex *PlayQueueByIndex `xml:"playQueueByIndex,omitempty" json:"playQueueByIndex,omitempty"` TranscodeDecision *TranscodeDecision `xml:"transcodeDecision,omitempty" json:"transcodeDecision,omitempty"` SonicMatches *Array[SonicMatch] `xml:"sonicMatch,omitempty" json:"sonicMatch,omitempty"` + TokenInfo *TokenInfo `xml:"tokenInfo,omitempty" json:"tokenInfo,omitempty"` } const ( @@ -596,6 +597,10 @@ type OpenSubsonicExtension struct { type OpenSubsonicExtensions []OpenSubsonicExtension +type TokenInfo struct { + Username string `xml:"username,attr" json:"username"` +} + type ItemGenre struct { Name string `xml:"name,attr" json:"name"` } diff --git a/server/subsonic/responses/responses_test.go b/server/subsonic/responses/responses_test.go index 586e46b63..027ac11d5 100644 --- a/server/subsonic/responses/responses_test.go +++ b/server/subsonic/responses/responses_test.go @@ -1015,6 +1015,20 @@ var _ = Describe("Responses", func() { }) }) + Describe("TokenInfo", func() { + BeforeEach(func() { + response.OpenSubsonic = true + response.TokenInfo = &TokenInfo{Username: "deluan"} + }) + + It("should match .XML", func() { + Expect(xml.MarshalIndent(response, "", " ")).To(MatchSnapshot()) + }) + It("should match .JSON", func() { + Expect(json.MarshalIndent(response, "", " ")).To(MatchSnapshot()) + }) + }) + Describe("InternetRadioStations", func() { BeforeEach(func() { response.InternetRadioStations = &InternetRadioStations{} diff --git a/server/subsonic/system.go b/server/subsonic/system.go index e14099942..4a59e7898 100644 --- a/server/subsonic/system.go +++ b/server/subsonic/system.go @@ -3,6 +3,7 @@ package subsonic import ( "net/http" + "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server/subsonic/responses" ) @@ -15,3 +16,10 @@ func (api *Router) GetLicense(_ *http.Request) (*responses.Subsonic, error) { response.License = &responses.License{Valid: true} return response, nil } + +func (api *Router) TokenInfo(r *http.Request) (*responses.Subsonic, error) { + user, _ := request.UserFrom(r.Context()) + response := newResponse() + response.TokenInfo = &responses.TokenInfo{Username: user.UserName} + return response, nil +} diff --git a/server/subsonic/system_test.go b/server/subsonic/system_test.go new file mode 100644 index 000000000..a8c85238e --- /dev/null +++ b/server/subsonic/system_test.go @@ -0,0 +1,23 @@ +package subsonic + +import ( + "net/http/httptest" + + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("TokenInfo", func() { + It("returns the authenticated username", func() { + api := &Router{} + r := httptest.NewRequest("GET", "/tokenInfo", nil) + r = r.WithContext(request.WithUser(r.Context(), model.User{UserName: "deluan"})) + + resp, err := api.TokenInfo(r) + + Expect(err).ToNot(HaveOccurred()) + Expect(resp.TokenInfo.Username).To(Equal("deluan")) + }) +}) diff --git a/tests/mock_data_store.go b/tests/mock_data_store.go index db798eece..a5e4126cf 100644 --- a/tests/mock_data_store.go +++ b/tests/mock_data_store.go @@ -228,7 +228,7 @@ func (db *MockDataStore) Player() model.PlayerRepository { if db.RealDS != nil { return db.RealDS.Player() } - db.MockedPlayer = struct{ model.PlayerRepository }{} + db.MockedPlayer = CreateMockPlayerRepo() return db.MockedPlayer } diff --git a/tests/mock_player_repo.go b/tests/mock_player_repo.go new file mode 100644 index 000000000..56835f0cc --- /dev/null +++ b/tests/mock_player_repo.go @@ -0,0 +1,75 @@ +package tests + +import ( + "context" + "maps" + "slices" + + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/id" +) + +func CreateMockPlayerRepo() *MockPlayerRepo { + return &MockPlayerRepo{Data: map[string]*model.Player{}, APIKeys: map[string]string{}} +} + +// MockPlayerRepo keeps API keys in plaintext (key -> player ID); hashing belongs to the real repository. +type MockPlayerRepo struct { + model.PlayerRepository + Error error + Data map[string]*model.Player + APIKeys map[string]string +} + +func (m *MockPlayerRepo) Get(_ context.Context, id string) (*model.Player, error) { + if m.Error != nil { + return nil, m.Error + } + p, ok := m.Data[id] + if !ok { + return nil, model.ErrNotFound + } + cp := *p + cp.HasAPIKey = slices.Contains(slices.Collect(maps.Values(m.APIKeys)), id) + return &cp, nil +} + +func (m *MockPlayerRepo) Put(_ context.Context, p *model.Player) error { + if m.Error != nil { + return m.Error + } + if p.ID == "" { + p.ID = id.NewRandom() + } + cp := *p + m.Data[p.ID] = &cp + return nil +} + +func (m *MockPlayerRepo) FindByAPIKey(ctx context.Context, key string) (*model.Player, error) { + if m.Error != nil { + return nil, m.Error + } + if playerID, ok := m.APIKeys[key]; ok { + return m.Get(ctx, playerID) + } + return nil, model.ErrNotFound +} + +func (m *MockPlayerRepo) SetAPIKey(_ context.Context, playerID, key string) error { + if m.Error != nil { + return m.Error + } + if _, ok := m.Data[playerID]; !ok { + return model.ErrNotFound + } + m.removeKeys(playerID) + if key != "" { + m.APIKeys[key] = playerID + } + return nil +} + +func (m *MockPlayerRepo) removeKeys(playerID string) { + maps.DeleteFunc(m.APIKeys, func(_ string, v string) bool { return v == playerID }) +} diff --git a/ui/src/App.jsx b/ui/src/App.jsx index d10aa5a33..4de369394 100644 --- a/ui/src/App.jsx +++ b/ui/src/App.jsx @@ -141,7 +141,7 @@ const Admin = (props) => { , permissions === 'admin' ? ( { return userResponse } +// ra-data-json-server merges the request body into the result; re-read so the plaintext key is never cached +const createPlayer = async (resource, params) => { + const { data } = await dataProvider.create(resource, params) + return dataProvider.getOne(resource, { id: data.id }) +} + const wrapperDataProvider = { ...dataProvider, getList: (resource, params) => { @@ -194,6 +200,9 @@ const wrapperDataProvider = { return createUser(params) } const [r, p] = mapResource(resource, params) + if (resource === 'player') { + return createPlayer(r, p) + } return dataProvider.create(r, p) }, delete: (resource, params) => { diff --git a/ui/src/dataProvider/wrapperDataProvider.test.js b/ui/src/dataProvider/wrapperDataProvider.test.js index 4225a5a54..1c33aad5d 100644 --- a/ui/src/dataProvider/wrapperDataProvider.test.js +++ b/ui/src/dataProvider/wrapperDataProvider.test.js @@ -88,6 +88,21 @@ describe('wrapperDataProvider', () => { }) }) + describe('create player', () => { + it('returns the server record, never the plaintext API key', async () => { + const data = { name: 'Phone', apiKey: 'nds_0123456789abcdefghijkl' } + const saved = { id: 'p1', name: 'Phone', hasApiKey: true, userId: 'u1' } + mockProvider.create.mockResolvedValue({ data: { ...data, id: 'p1' } }) + mockProvider.getOne.mockResolvedValue({ data: saved }) + + const result = await wrapperDataProvider.create('player', { data }) + + expect(mockProvider.create).toHaveBeenCalledWith('player', { data }) + expect(mockProvider.getOne).toHaveBeenCalledWith('player', { id: 'p1' }) + expect(result.data).toEqual(saved) + }) + }) + describe('refreshMetadata', () => { it('posts to the album metadata refresh endpoint', () => { mockHttpClient.mockResolvedValue({ json: {} }) diff --git a/ui/src/i18n/en.json b/ui/src/i18n/en.json index f5af05b7b..851b47ee3 100644 --- a/ui/src/i18n/en.json +++ b/ui/src/i18n/en.json @@ -182,6 +182,7 @@ }, "player": { "name": "Player |||| Players", + "menuName": "Players & API keys", "fields": { "name": "Name", "transcodingId": "Transcoding", @@ -190,7 +191,29 @@ "userName": "Username", "lastSeen": "Last Seen At", "reportRealPath": "Report Real Path", - "scrobbleEnabled": "Send Scrobbles to external services" + "scrobbleEnabled": "Send Scrobbles to external services", + "hasApiKey": "API Key" + }, + "actions": { + "generateApiKey": "Generate API key", + "regenerateApiKey": "Regenerate", + "revokeApiKey": "Revoke", + "copyApiKey": "Copy" + }, + "message": { + "apiKeyActive": "This player has an API key. Use it in your app as the API key, or as the password if the app does not use token authentication.", + "apiKeyNone": "No API key. Generate one to connect an app to this player.", + "apiKeyNoneOther": "No API key.", + "apiKeyPending": "Copy this key now. It is saved when you click Save and will not be shown again.", + "apiKeyRevokePending": "The API key will be removed when you save.", + "deleteWithKeyTitle": "Delete player", + "deleteWithKeyContent": "This player has an API key. Apps using it will stop working." + }, + "notifications": { + "apiKeyCopied": "API key copied to clipboard" + }, + "validation": { + "apiKeyFormat": "Invalid API key format" } }, "transcoding": { diff --git a/ui/src/layout/AppBar.jsx b/ui/src/layout/AppBar.jsx index 7de111e67..460d33bb9 100644 --- a/ui/src/layout/AppBar.jsx +++ b/ui/src/layout/AppBar.jsx @@ -102,9 +102,11 @@ const CustomUserMenu = ({ onClick, ...rest }) => { } const renderSettingsMenuItemLink = (resource, id) => { - const label = translate(`resources.${resource.name}.name`, { - smart_count: id ? 1 : 2, - }) + const label = resource.options.label + ? translate(resource.options.label) + : translate(`resources.${resource.name}.name`, { + smart_count: id ? 1 : 2, + }) const link = id ? `/${resource.name}/${id}` : `/${resource.name}` return ( ({ resources: [] })) + vi.mock('react-admin', () => ({ AppBar: ({ userMenu }) =>
{userMenu}
, + MenuItemLink: ({ primaryText }) =>
{primaryText}
, useTranslate: () => (x) => x, usePermissions: () => ({ permissions: 'admin' }), - getResources: () => [], + getResources: () => mocks.resources, })) vi.mock('./NowPlayingPanel', () => ({ @@ -41,6 +44,7 @@ describe('', () => { config.devActivityPanel = true config.enableNowPlaying = true config.enableQuickConnect = false + mocks.resources = [] store = createStore(combineReducers({ activity: activityReducer }), { activity: { nowPlayingCount: 0 }, }) @@ -84,4 +88,22 @@ describe('', () => { expect(screen.queryAllByText('menu.quickConnect.name')).toHaveLength(0) expect(screen.queryAllByText('menu.about')).not.toHaveLength(0) }) + + it('uses the resource label for settings items when set', () => { + mocks.resources = [ + { + name: 'player', + hasList: true, + options: { subMenu: 'settings', label: 'resources.player.menuName' }, + }, + { name: 'transcoding', hasList: true, options: { subMenu: 'settings' } }, + ] + render( + + + , + ) + expect(screen.getByText('resources.player.menuName')).toBeInTheDocument() + expect(screen.getByText('resources.transcoding.name')).toBeInTheDocument() + }) }) diff --git a/ui/src/player/ApiKeyInput.jsx b/ui/src/player/ApiKeyInput.jsx new file mode 100644 index 000000000..1f2dd1ed2 --- /dev/null +++ b/ui/src/player/ApiKeyInput.jsx @@ -0,0 +1,107 @@ +import React from 'react' +import PropTypes from 'prop-types' +import { useInput, useNotify, useTranslate } from 'react-admin' +import { Button, TextField } from '@material-ui/core' +import { FaKey } from 'react-icons/fa' +import { MdContentCopy, MdDelete, MdRefresh } from 'react-icons/md' +import { isWritable } from '../common/playlistUtils' +import { generateApiKey } from './apiKey' + +const identity = (v) => v +const MASK = '•'.repeat(26) + +const ApiKeyInput = ({ record, isCreate, fullWidth, className, ...props }) => { + const translate = useTranslate() + const notify = useNotify() + // Identity format/parse keep "" (revoke) distinct from undefined (untouched) + const { + input: { value, onChange }, + meta: { error, touched }, + } = useInput({ ...props, format: identity, parse: identity }) + + const isOwner = isCreate || record?.userId === localStorage.getItem('userId') + const pending = !!value + const revoking = value === '' && !!record?.hasApiKey + const saved = value == null && !!record?.hasApiKey + const hasKey = pending || saved + + const copy = () => { + const fallback = () => + prompt(translate('message.shareCopyToClipboard'), value) + if (navigator.clipboard && window.isSecureContext) { + navigator.clipboard + .writeText(value) + .then( + () => notify('resources.player.notifications.apiKeyCopied'), + fallback, + ) + } else { + fallback() + } + } + + const helperText = pending + ? 'resources.player.message.apiKeyPending' + : revoking + ? 'resources.player.message.apiKeyRevokePending' + : saved + ? 'resources.player.message.apiKeyActive' + : isOwner + ? 'resources.player.message.apiKeyNone' + : 'resources.player.message.apiKeyNoneOther' + + return ( +
+ +
+ {pending && ( + + )} + {isOwner && ( + + )} + {saved && isWritable(record?.userId) && ( + + )} +
+
+ ) +} + +ApiKeyInput.propTypes = { + source: PropTypes.string.isRequired, + record: PropTypes.object, + isCreate: PropTypes.bool, + fullWidth: PropTypes.bool, + className: PropTypes.string, + validate: PropTypes.oneOfType([PropTypes.func, PropTypes.array]), +} + +export default ApiKeyInput diff --git a/ui/src/player/ApiKeyInput.test.jsx b/ui/src/player/ApiKeyInput.test.jsx new file mode 100644 index 000000000..dd4269d6e --- /dev/null +++ b/ui/src/player/ApiKeyInput.test.jsx @@ -0,0 +1,172 @@ +import * as React from 'react' +import { render, screen, fireEvent, waitFor } from '@testing-library/react' +import { Form } from 'react-final-form' +import { describe, it, expect, vi, beforeEach } from 'vitest' +import ApiKeyInput from './ApiKeyInput' + +const hooks = vi.hoisted(() => ({ notify: vi.fn() })) +const KEY = 'nds_0123456789abcdefghijkl' +const KEY_FORMAT = /^nds_[0-9A-Za-z]{22}$/ + +vi.mock('react-admin', async () => { + const actual = await vi.importActual('react-admin') + return { + ...actual, + useNotify: () => hooks.notify, + useTranslate: () => (key) => key, + } +}) + +const renderInput = ({ + record, + initialValues = {}, + isCreate = false, + fullWidth, +}) => { + let values + const utils = render( +
{}} + initialValues={initialValues} + render={({ values: v }) => { + values = v + return ( + + ) + }} + />, + ) + return { ...utils, values: () => values } +} + +const text = (key) => screen.queryByText(key) + +describe('ApiKeyInput', () => { + beforeEach(() => { + vi.clearAllMocks() + localStorage.setItem('userId', 'owner') + localStorage.setItem('role', 'regular') + }) + + it('shows a pending key with copy and regenerate on create', () => { + const { values } = renderInput({ + record: {}, + isCreate: true, + initialValues: { apiKey: KEY }, + }) + expect(screen.getByDisplayValue(KEY)).toBeInTheDocument() + expect(text('resources.player.message.apiKeyPending')).toBeInTheDocument() + expect( + screen.getByRole('button', { + name: 'resources.player.actions.copyApiKey', + }), + ).toBeInTheDocument() + expect( + text('resources.player.actions.revokeApiKey'), + ).not.toBeInTheDocument() + + fireEvent.click( + screen.getByText('resources.player.actions.regenerateApiKey'), + ) + expect(values().apiKey).toMatch(KEY_FORMAT) + expect(values().apiKey).not.toBe(KEY) + }) + + it('masks a saved key and lets the owner regenerate or revoke', () => { + const { values } = renderInput({ + record: { id: 'p1', userId: 'owner', hasApiKey: true }, + }) + expect(screen.queryByDisplayValue(/^nds_/)).not.toBeInTheDocument() + expect(text('resources.player.message.apiKeyActive')).toBeInTheDocument() + expect( + text('resources.player.actions.regenerateApiKey'), + ).toBeInTheDocument() + + fireEvent.click(screen.getByText('resources.player.actions.revokeApiKey')) + expect(values().apiKey).toBe('') + expect( + text('resources.player.message.apiKeyRevokePending'), + ).toBeInTheDocument() + expect(text('resources.player.actions.generateApiKey')).toBeInTheDocument() + }) + + it('lets the owner generate a key when there is none', () => { + const { values } = renderInput({ + record: { id: 'p1', userId: 'owner', hasApiKey: false }, + }) + expect(values().apiKey).toBeUndefined() + expect(text('resources.player.message.apiKeyNone')).toBeInTheDocument() + + fireEvent.click(screen.getByText('resources.player.actions.generateApiKey')) + expect(values().apiKey).toMatch(KEY_FORMAT) + expect(text('resources.player.message.apiKeyPending')).toBeInTheDocument() + }) + + it('lets an admin revoke but not set a key on another user player', () => { + localStorage.setItem('role', 'admin') + renderInput({ record: { id: 'p1', userId: 'someone', hasApiKey: true } }) + expect( + text('resources.player.actions.regenerateApiKey'), + ).not.toBeInTheDocument() + expect(text('resources.player.actions.revokeApiKey')).toBeInTheDocument() + }) + + it('falls back to a prompt when the clipboard write fails', async () => { + vi.stubGlobal('isSecureContext', true) + vi.stubGlobal('prompt', vi.fn()) + Object.defineProperty(navigator, 'clipboard', { + configurable: true, + value: { writeText: vi.fn().mockRejectedValue(new Error('denied')) }, + }) + try { + renderInput({ + record: {}, + isCreate: true, + initialValues: { apiKey: KEY }, + }) + fireEvent.click( + screen.getByRole('button', { + name: 'resources.player.actions.copyApiKey', + }), + ) + await waitFor(() => + expect(window.prompt).toHaveBeenCalledWith( + 'message.shareCopyToClipboard', + KEY, + ), + ) + expect(hooks.notify).not.toHaveBeenCalled() + } finally { + vi.unstubAllGlobals() + delete navigator.clipboard + } + }) + + it('shows a neutral message to an admin viewing another user player with no key', () => { + localStorage.setItem('role', 'admin') + renderInput({ record: { id: 'p1', userId: 'someone', hasApiKey: false } }) + expect(text('resources.player.message.apiKeyNoneOther')).toBeInTheDocument() + expect(text('resources.player.message.apiKeyNone')).not.toBeInTheDocument() + expect(screen.queryAllByRole('button')).toHaveLength(0) + }) + + it('shows no actions to another regular user', () => { + renderInput({ record: { id: 'p1', userId: 'someone', hasApiKey: true } }) + expect(screen.queryAllByRole('button')).toHaveLength(0) + }) + + it('is not full width unless asked', () => { + const record = { id: 'p1', userId: 'owner', hasApiKey: true } + const { container, unmount } = renderInput({ record }) + expect(container.querySelector('.MuiFormControl-fullWidth')).toBeNull() + unmount() + + const { container: wide } = renderInput({ record, fullWidth: true }) + expect(wide.querySelector('.MuiFormControl-fullWidth')).not.toBeNull() + }) +}) diff --git a/ui/src/player/PlayerCreate.jsx b/ui/src/player/PlayerCreate.jsx new file mode 100644 index 000000000..bf02cde1b --- /dev/null +++ b/ui/src/player/PlayerCreate.jsx @@ -0,0 +1,33 @@ +import React, { useMemo } from 'react' +import { Create, SimpleForm, required, useTranslate } from 'react-admin' +import { Title } from '../common' +import { playerInputs } from './playerInputs' +import ApiKeyInput from './ApiKeyInput' +import { generateApiKey } from './apiKey' + +const PlayerCreateTitle = () => { + const translate = useTranslate() + const resourceName = translate('resources.player.name', { smart_count: 1 }) + return ( + + ) +} + +const PlayerCreate = (props) => { + // Memoized so re-renders don't swap the key the user may have already copied + const initialValues = useMemo(() => ({ apiKey: generateApiKey() }), []) + return ( + <Create title={<PlayerCreateTitle />} {...props}> + <SimpleForm + variant="outlined" + redirect="list" + initialValues={initialValues} + > + {playerInputs()} + <ApiKeyInput source="apiKey" isCreate validate={required()} /> + </SimpleForm> + </Create> + ) +} + +export default PlayerCreate diff --git a/ui/src/player/PlayerCreate.test.jsx b/ui/src/player/PlayerCreate.test.jsx new file mode 100644 index 000000000..dc1f817db --- /dev/null +++ b/ui/src/player/PlayerCreate.test.jsx @@ -0,0 +1,47 @@ +import * as React from 'react' +import { render } from '@testing-library/react' +import { describe, it, expect, vi, beforeEach } from 'vitest' +import PlayerCreate from './PlayerCreate' +import ApiKeyInput from './ApiKeyInput' + +const hooks = vi.hoisted(() => ({ forms: [] })) + +vi.mock('react-admin', async () => { + const actual = await vi.importActual('react-admin') + return { + ...actual, + Create: ({ children }) => children, + SimpleForm: (props) => { + hooks.forms.push(props) + return null + }, + } +}) + +describe('PlayerCreate', () => { + beforeEach(() => { + hooks.forms = [] + }) + + it('pre-fills one generated API key that survives re-renders', () => { + const { rerender } = render(<PlayerCreate resource="player" />) + rerender(<PlayerCreate resource="player" />) + + const [first, second] = hooks.forms.map((f) => f.initialValues) + expect(hooks.forms).toHaveLength(2) + expect(first.apiKey).toMatch(/^nds_[0-9A-Za-z]{22}$/) + expect(second).toBe(first) + }) + + it('requires the API key', () => { + render(<PlayerCreate resource="player" />) + + const input = React.Children.toArray(hooks.forms[0].children).find( + (child) => child.type === ApiKeyInput, + ) + expect(input.props.source).toBe('apiKey') + expect(input.props.isCreate).toBe(true) + expect(input.props.validate('')).toBeTruthy() + expect(input.props.validate('nds_0123456789abcdefghijkl')).toBeUndefined() + }) +}) diff --git a/ui/src/player/PlayerEdit.jsx b/ui/src/player/PlayerEdit.jsx index 1826500bd..08f389698 100644 --- a/ui/src/player/PlayerEdit.jsx +++ b/ui/src/player/PlayerEdit.jsx @@ -1,17 +1,17 @@ import { - TextInput, - BooleanInput, TextField, Edit, - required, SimpleForm, - SelectInput, - ReferenceInput, useTranslate, + DeleteButton, + DeleteWithConfirmButton, + SaveButton, + Toolbar, } from 'react-admin' +import { makeStyles } from '@material-ui/core/styles' import { Title } from '../common' -import config from '../config' -import { BITRATE_CHOICES } from '../consts' +import ApiKeyInput from './ApiKeyInput' +import { playerInputs } from './playerInputs' const PlayerTitle = ({ record }) => { const translate = useTranslate() @@ -19,24 +19,35 @@ const PlayerTitle = ({ record }) => { return <Title subTitle={`${resourceName} ${record ? record.name : ''}`} /> } +const useToolbarStyles = makeStyles({ + toolbar: { + display: 'flex', + justifyContent: 'space-between', + }, +}) + +const PlayerEditToolbar = (props) => ( + <Toolbar {...props} classes={useToolbarStyles()}> + <SaveButton /> + {props.record?.hasApiKey ? ( + <DeleteWithConfirmButton + mutationMode="pessimistic" + confirmTitle="resources.player.message.deleteWithKeyTitle" + confirmContent="resources.player.message.deleteWithKeyContent" + /> + ) : ( + <DeleteButton /> + )} + </Toolbar> +) + const PlayerEdit = (props) => ( - <Edit title={<PlayerTitle />} {...props}> - <SimpleForm variant={'outlined'}> - <TextInput source="name" validate={[required()]} /> - <ReferenceInput - source="transcodingId" - reference="transcoding" - sort={{ field: 'name', order: 'ASC' }} - > - <SelectInput source="name" resettable /> - </ReferenceInput> - <SelectInput source="maxBitRate" resettable choices={BITRATE_CHOICES} /> - <BooleanInput source="reportRealPath" fullWidth /> - {(config.lastFMEnabled || config.listenBrainzEnabled) && ( - <BooleanInput source="scrobbleEnabled" fullWidth /> - )} + <Edit title={<PlayerTitle />} mutationMode="pessimistic" {...props}> + <SimpleForm variant={'outlined'} toolbar={<PlayerEditToolbar />}> + {playerInputs()} <TextField source="client" /> <TextField source="userName" /> + <ApiKeyInput source="apiKey" /> </SimpleForm> </Edit> ) diff --git a/ui/src/player/PlayerEdit.test.jsx b/ui/src/player/PlayerEdit.test.jsx new file mode 100644 index 000000000..38fa2cc8e --- /dev/null +++ b/ui/src/player/PlayerEdit.test.jsx @@ -0,0 +1,25 @@ +import * as React from 'react' +import { render } from '@testing-library/react' +import { describe, it, expect, vi } from 'vitest' +import PlayerEdit from './PlayerEdit' + +const hooks = vi.hoisted(() => ({ editProps: null })) + +vi.mock('react-admin', async () => { + const actual = await vi.importActual('react-admin') + return { + ...actual, + Edit: (props) => { + hooks.editProps = props + return null + }, + } +}) + +describe('PlayerEdit', () => { + // An optimistic or undoable save would put the new key in react-admin's cache + it('saves pessimistically', () => { + render(<PlayerEdit resource="player" id="p1" />) + expect(hooks.editProps.mutationMode).toBe('pessimistic') + }) +}) diff --git a/ui/src/player/PlayerList.jsx b/ui/src/player/PlayerList.jsx index a2b009bad..c7bd32570 100644 --- a/ui/src/player/PlayerList.jsx +++ b/ui/src/player/PlayerList.jsx @@ -2,18 +2,20 @@ import React from 'react' import { Datagrid, TextField, - DateField, FunctionField, ReferenceField, Filter, SearchInput, + NullableBooleanInput, } from 'react-admin' import { useMediaQuery } from '@material-ui/core' -import { SimpleList, List } from '../common' +import { FaKey } from 'react-icons/fa' +import { SimpleList, List, DateField } from '../common' const PlayerFilter = (props) => ( <Filter {...props} variant={'outlined'}> <SearchInput id="search" source="name" alwaysOn /> + <NullableBooleanInput source="hasApiKey" alwaysOn /> </Filter> ) @@ -30,7 +32,11 @@ const PlayerList = ({ permissions, ...props }) => { <SimpleList primaryText={(r) => r.name} secondaryText={(r) => r.userName} - tertiaryText={(r) => (r.maxBitRate ? r.maxBitRate : '-')} + tertiaryText={(r) => ( + <> + {r.hasApiKey && <FaKey />} {r.maxBitRate ? r.maxBitRate : '-'} + </> + )} /> ) : ( <Datagrid rowClick="edit"> @@ -43,6 +49,11 @@ const PlayerList = ({ permissions, ...props }) => { source="maxBitRate" render={(r) => (r.maxBitRate ? r.maxBitRate : '-')} /> + <FunctionField + source="hasApiKey" + sortable={false} + render={(r) => (r.hasApiKey ? <FaKey /> : null)} + /> <DateField source="lastSeen" showTime sortByOrder={'DESC'} /> </Datagrid> )} diff --git a/ui/src/player/apiKey.js b/ui/src/player/apiKey.js new file mode 100644 index 000000000..27d683052 --- /dev/null +++ b/ui/src/player/apiKey.js @@ -0,0 +1,15 @@ +const API_KEY_PREFIX = 'nds_' +const ALPHABET = + '0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz' +const KEY_LENGTH = 22 + +// Bytes >= 248 are dropped so every character is equally likely (248 = 4 * 62). +export const generateApiKey = () => { + let key = '' + while (key.length < KEY_LENGTH) { + for (const b of window.crypto.getRandomValues(new Uint8Array(32))) { + if (b < 248 && key.length < KEY_LENGTH) key += ALPHABET[b % 62] + } + } + return API_KEY_PREFIX + key +} diff --git a/ui/src/player/apiKey.test.js b/ui/src/player/apiKey.test.js new file mode 100644 index 000000000..4c41871e7 --- /dev/null +++ b/ui/src/player/apiKey.test.js @@ -0,0 +1,15 @@ +import { describe, it, expect } from 'vitest' +import { generateApiKey } from './apiKey' + +describe('generateApiKey', () => { + it('returns the prefix plus 22 base62 characters', () => { + for (let i = 0; i < 50; i++) { + expect(generateApiKey()).toMatch(/^nds_[0-9A-Za-z]{22}$/) + } + }) + + it('returns a different key each time', () => { + const keys = new Set(Array.from({ length: 100 }, generateApiKey)) + expect(keys.size).toBe(100) + }) +}) diff --git a/ui/src/player/index.js b/ui/src/player/index.js index aaa3d58d7..c28cfa414 100644 --- a/ui/src/player/index.js +++ b/ui/src/player/index.js @@ -1,9 +1,11 @@ import { BsFillMusicPlayerFill } from 'react-icons/bs' import PlayerList from './PlayerList' import PlayerEdit from './PlayerEdit' +import PlayerCreate from './PlayerCreate' export default { list: PlayerList, edit: PlayerEdit, + create: PlayerCreate, icon: BsFillMusicPlayerFill, } diff --git a/ui/src/player/playerInputs.jsx b/ui/src/player/playerInputs.jsx new file mode 100644 index 000000000..a8b719a81 --- /dev/null +++ b/ui/src/player/playerInputs.jsx @@ -0,0 +1,34 @@ +import React from 'react' +import { + BooleanInput, + ReferenceInput, + SelectInput, + TextInput, + required, +} from 'react-admin' +import config from '../config' +import { BITRATE_CHOICES } from '../consts' + +// Returned as an array, not a component, so SimpleForm still injects its props into each input. +export const playerInputs = () => + [ + <TextInput key="name" source="name" validate={[required()]} />, + <ReferenceInput + key="transcodingId" + source="transcodingId" + reference="transcoding" + sort={{ field: 'name', order: 'ASC' }} + > + <SelectInput source="name" resettable /> + </ReferenceInput>, + <SelectInput + key="maxBitRate" + source="maxBitRate" + resettable + choices={BITRATE_CHOICES} + />, + <BooleanInput key="reportRealPath" source="reportRealPath" fullWidth />, + (config.lastFMEnabled || config.listenBrainzEnabled) && ( + <BooleanInput key="scrobbleEnabled" source="scrobbleEnabled" fullWidth /> + ), + ].filter(Boolean) From ab5869123dad0d435bda3040849a9b9c0fc11888 Mon Sep 17 00:00:00 2001 From: DawidKrynski <dawid_krynski.64@wp.pl> Date: Mon, 28 Sep 2026 14:07:05 +0200 Subject: [PATCH 21/26] fix(playlists): include co-credited album artists when adding an artist to a playlist - #6240 (#6241) * fix(persistence): include co-credited album artists when adding an artist to a playlist - #6240 Signed-off-by: Dawid Krynski <188586034+DawidKrynski@users.noreply.github.com> * test(persistence): cover first album artist and track-artist-only in AddArtists The joint track now uses a track artist that is not an album artist, and the AddArtists specs check all three cases: the first album artist still matches, a co-credited album artist matches, and a track-artist-only ID adds nothing. The last case guards against widening the role filter. --------- Signed-off-by: Dawid Krynski <188586034+DawidKrynski@users.noreply.github.com> Co-authored-by: Dawid Krynski <188586034+DawidKrynski@users.noreply.github.com> Co-authored-by: Deluan <deluan@navidrome.org> --- persistence/playlist_track_repository.go | 4 +- persistence/playlist_track_repository_test.go | 39 +++++++++++++++++++ 2 files changed, 42 insertions(+), 1 deletion(-) diff --git a/persistence/playlist_track_repository.go b/persistence/playlist_track_repository.go index 341aab8af..392446cef 100644 --- a/persistence/playlist_track_repository.go +++ b/persistence/playlist_track_repository.go @@ -234,7 +234,9 @@ func (r *playlistTrackRepository) AddAlbums(ctx context.Context, albumIds []stri } func (r *playlistTrackRepository) AddArtists(ctx context.Context, artistIds []string) (int, error) { - return r.addMediaFileIds(ctx, Eq{"album_artist_id": artistIds}) + // Match by album-artist participation, not the deprecated album_artist_id + // column, which only holds the first album artist. + return r.addMediaFileIds(ctx, ParticipantIDFilter("media_file", artistIds, model.RoleAlbumArtist)) } func (r *playlistTrackRepository) AddDiscs(ctx context.Context, discs []model.DiscID) (int, error) { diff --git a/persistence/playlist_track_repository_test.go b/persistence/playlist_track_repository_test.go index 3c532c405..119c57116 100644 --- a/persistence/playlist_track_repository_test.go +++ b/persistence/playlist_track_repository_test.go @@ -219,6 +219,45 @@ var _ = Describe("PlaylistTrackRepository", func() { }) }) + Describe("AddArtists", func() { + var tracks model.PlaylistTrackRepository + var joint model.MediaFile + + BeforeEach(func() { + mfRepo := NewMediaFileRepository(GetDBXBuilder()) + joint = mf(model.MediaFile{ID: "pls-coartist-track", Title: "Joint Track", ArtistID: artistPunctuation.ID, + Artist: artistPunctuation.Name, AlbumID: "pls-coartist-album", Album: "Joint Album", + AlbumArtistID: artistKraftwerk.ID, AlbumArtist: artistKraftwerk.Name, Path: p("joint/track.mp3")}) + joint.Participants[model.RoleAlbumArtist] = model.ParticipantList{ + {Artist: artistKraftwerk}, + {Artist: artistBeatles}, + } + Expect(mfRepo.Put(ctx, &joint)).To(Succeed()) + DeferCleanup(func() { _ = mfRepo.Delete(ctx, joint.ID) }) + + plsRepo := NewPlaylistRepository(GetDBXBuilder()) + pls := model.Playlist{Name: "Co-album-artist", OwnerID: adminUser.ID, OwnerName: adminUser.UserName} + Expect(plsRepo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { _ = plsRepo.Delete(ctx, pls.ID) }) + tracks = plsRepo.Tracks(ctx, pls.ID, false) + }) + + It("adds tracks where the artist is the first album artist", func() { + Expect(tracks.AddArtists(ctx, []string{artistKraftwerk.ID})).To(Equal(1)) + Expect(tracks.GetMediaFileIDs(ctx)).To(ConsistOf(joint.ID)) + }) + + It("adds tracks where the artist is not the first album artist", func() { + Expect(tracks.AddArtists(ctx, []string{artistBeatles.ID})).To(Equal(1)) + Expect(tracks.GetMediaFileIDs(ctx)).To(ConsistOf(joint.ID)) + }) + + It("does not add tracks where the artist is only the track artist", func() { + Expect(tracks.AddArtists(ctx, []string{artistPunctuation.ID})).To(Equal(0)) + Expect(tracks.GetMediaFileIDs(ctx)).To(BeEmpty()) + }) + }) + Describe("library access", func() { var otherLib model.Library var restrictedUser model.User From 612c8290f956c6a35a75c4adad6aa12f10da9c3a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= <deluan@navidrome.org> Date: Mon, 28 Sep 2026 08:48:36 -0400 Subject: [PATCH 22/26] feat(ui): show read-only values in edit forms as dimmed, themable inputs (#6238) --- ui/src/common/ReadOnlyFields.jsx | 99 +++++++++++++++ ui/src/common/ReadOnlyFields.test.jsx | 121 +++++++++++++++++++ ui/src/common/index.js | 1 + ui/src/library/LibraryEdit.jsx | 166 +++++++------------------- ui/src/player/PlayerEdit.jsx | 7 +- ui/src/playlist/PlaylistEdit.jsx | 12 +- ui/src/radio/RadioEdit.jsx | 12 +- ui/src/share/ShareEdit.jsx | 31 +++-- ui/src/user/UserEdit.jsx | 11 +- ui/src/user/UserEdit.test.jsx | 6 +- 10 files changed, 306 insertions(+), 160 deletions(-) create mode 100644 ui/src/common/ReadOnlyFields.jsx create mode 100644 ui/src/common/ReadOnlyFields.test.jsx diff --git a/ui/src/common/ReadOnlyFields.jsx b/ui/src/common/ReadOnlyFields.jsx new file mode 100644 index 000000000..c56643fbd --- /dev/null +++ b/ui/src/common/ReadOnlyFields.jsx @@ -0,0 +1,99 @@ +import React from 'react' +import PropTypes from 'prop-types' +import get from 'lodash/get' +import { FieldTitle, useRecordContext } from 'react-admin' +import { TextField } from '@material-ui/core' +import { makeStyles } from '@material-ui/core/styles' +import { useDateLocale } from '../i18n/useDateLocale' +import { + formatBytes, + formatDateTime, + formatDuration2, + formatNumber, +} from '../utils/formatters' +import { isDateSet } from '../utils/validations' + +const useStyles = makeStyles( + (theme) => ({ + inputRoot: { + '&:hover $notchedOutline': { + borderColor: theme.palette.divider, + }, + }, + notchedOutline: { + borderColor: theme.palette.divider, + }, + }), + { name: 'NDReadOnlyField' }, +) + +const identity = (v) => v + +// Renders a record value as a dimmed, non-editable input, so it lines up with the inputs in a form +export const ReadOnlyTextField = ({ + source, + label, + resource, + className, + fullWidth, + format = identity, + ...props +}) => { + const classes = useStyles(props) + const record = useRecordContext(props) + const value = get(record, source) + + return ( + <TextField + id={source} + className={className} + label={<FieldTitle label={label} source={source} resource={resource} />} + value={value == null ? '' : format(value)} + variant="outlined" + margin="dense" + fullWidth={fullWidth} + focused={false} + helperText=" " + InputProps={{ + readOnly: true, + classes: { + root: classes.inputRoot, + notchedOutline: classes.notchedOutline, + }, + }} + inputProps={{ tabIndex: -1 }} + /> + ) +} + +ReadOnlyTextField.propTypes = { + source: PropTypes.string.isRequired, + label: PropTypes.oneOfType([PropTypes.string, PropTypes.bool]), + record: PropTypes.object, + resource: PropTypes.string, + className: PropTypes.string, + classes: PropTypes.object, + fullWidth: PropTypes.bool, + format: PropTypes.func, +} + +export const ReadOnlyDateField = (props) => { + const locale = useDateLocale() + const format = (v) => (isDateSet(v) ? formatDateTime(v, locale) : '') + return <ReadOnlyTextField format={format} {...props} /> +} + +export const ReadOnlyNumberField = (props) => { + const locale = useDateLocale() + return ( + <ReadOnlyTextField format={(v) => formatNumber(v, locale)} {...props} /> + ) +} + +export const ReadOnlySizeField = (props) => ( + <ReadOnlyTextField format={formatBytes} {...props} /> +) + +export const ReadOnlyDurationField = (props) => ( + <ReadOnlyTextField format={formatDuration2} {...props} /> +) diff --git a/ui/src/common/ReadOnlyFields.test.jsx b/ui/src/common/ReadOnlyFields.test.jsx new file mode 100644 index 000000000..00543a786 --- /dev/null +++ b/ui/src/common/ReadOnlyFields.test.jsx @@ -0,0 +1,121 @@ +import React from 'react' +import { render, screen } from '@testing-library/react' +import { describe, it, expect, beforeEach, vi } from 'vitest' +import { + ReadOnlyDateField, + ReadOnlyDurationField, + ReadOnlyNumberField, + ReadOnlySizeField, + ReadOnlyTextField, +} from './ReadOnlyFields' + +vi.mock('react-admin', async (importOriginal) => ({ + ...(await importOriginal()), + useLocale: vi.fn(), +})) + +describe('ReadOnlyFields', () => { + const record = { + id: '1', + client: 'NavidromeUI', + createdAt: '2026-09-17T14:30:00Z', + lastVisitedAt: '0001-01-01T00:00:00Z', + count: 1234567, + zero: 0, + size: 1536000, + duration: 3725, + } + + beforeEach(async () => { + vi.clearAllMocks() + vi.spyOn(navigator, 'languages', 'get').mockReturnValue([]) + const { useLocale } = await import('react-admin') + vi.mocked(useLocale).mockReturnValue('de') + }) + + const renderField = (Field, props) => + render(<Field record={record} resource="player" {...props} />) + + describe('<ReadOnlyTextField>', () => { + it('shows the record value with the translated field label', () => { + renderField(ReadOnlyTextField, { source: 'client' }) + const input = screen.getByLabelText('resources.player.fields.client') + expect(input).toHaveValue('NavidromeUI') + }) + + it('cannot be edited or reached with the Tab key', () => { + renderField(ReadOnlyTextField, { source: 'client' }) + const input = screen.getByRole('textbox') + expect(input).toHaveAttribute('readonly') + expect(input).toHaveAttribute('tabindex', '-1') + }) + + it('shows an empty value when the record has no value', () => { + renderField(ReadOnlyTextField, { source: 'userName' }) + expect(screen.getByRole('textbox')).toHaveValue('') + }) + + it('uses an explicit label when given', () => { + renderField(ReadOnlyTextField, { source: 'client', label: 'Custom' }) + expect(screen.getByLabelText('Custom')).toBeInTheDocument() + }) + + it('applies a custom format', () => { + renderField(ReadOnlyTextField, { + source: 'client', + format: (v) => v.toUpperCase(), + }) + expect(screen.getByRole('textbox')).toHaveValue('NAVIDROMEUI') + }) + + it('exposes theme-overridable class names', () => { + const { container } = renderField(ReadOnlyTextField, { + source: 'client', + }) + expect( + container.querySelector('[class*="NDReadOnlyField-inputRoot"]'), + ).toBeInTheDocument() + expect( + container.querySelector('[class*="NDReadOnlyField-notchedOutline"]'), + ).toBeInTheDocument() + }) + }) + + describe('<ReadOnlyDateField>', () => { + it('formats the date using the selected language', () => { + renderField(ReadOnlyDateField, { source: 'createdAt' }) + expect(screen.getByRole('textbox').value).toMatch(/^17\.9\.2026, /) + }) + + it('shows an empty value when the date is not set', () => { + renderField(ReadOnlyDateField, { source: 'lastVisitedAt' }) + expect(screen.getByRole('textbox')).toHaveValue('') + }) + }) + + describe('<ReadOnlyNumberField>', () => { + it('formats the number using the selected language', () => { + renderField(ReadOnlyNumberField, { source: 'count' }) + expect(screen.getByRole('textbox')).toHaveValue('1.234.567') + }) + + it('shows zero', () => { + renderField(ReadOnlyNumberField, { source: 'zero' }) + expect(screen.getByRole('textbox')).toHaveValue('0') + }) + }) + + describe('<ReadOnlySizeField>', () => { + it('formats bytes as a human-readable size', () => { + renderField(ReadOnlySizeField, { source: 'size' }) + expect(screen.getByRole('textbox')).toHaveValue('1.46 MB') + }) + }) + + describe('<ReadOnlyDurationField>', () => { + it('formats seconds as a human-readable duration', () => { + renderField(ReadOnlyDurationField, { source: 'duration' }) + expect(screen.getByRole('textbox')).toHaveValue('1h 2m 5s') + }) + }) +}) diff --git a/ui/src/common/index.js b/ui/src/common/index.js index 0177df326..fb8f40f00 100644 --- a/ui/src/common/index.js +++ b/ui/src/common/index.js @@ -15,6 +15,7 @@ export * from './perPageStore' export * from './PlayButton' export * from './QuickFilter' export * from './RangeField' +export * from './ReadOnlyFields' export * from './ShuffleAllButton' export * from './SimpleList' export * from './SizeField' diff --git a/ui/src/library/LibraryEdit.jsx b/ui/src/library/LibraryEdit.jsx index 7e89c892c..53d17ac7f 100644 --- a/ui/src/library/LibraryEdit.jsx +++ b/ui/src/library/LibraryEdit.jsx @@ -6,7 +6,6 @@ import { BooleanInput, required, SaveButton, - DateField, useTranslate, useMutation, useNotify, @@ -16,8 +15,13 @@ import { import { Typography, Box } from '@material-ui/core' import { makeStyles } from '@material-ui/core/styles' import DeleteLibraryButton from './DeleteLibraryButton' -import { Title } from '../common' -import { formatBytes, formatDuration2, formatNumber } from '../utils/index.js' +import { + ReadOnlyDateField, + ReadOnlyDurationField, + ReadOnlyNumberField, + ReadOnlySizeField, + Title, +} from '../common' const useStyles = makeStyles({ toolbar: { @@ -26,6 +30,8 @@ const useStyles = makeStyles({ }, }) +const readOnlyProps = { resource: 'library', fullWidth: true } + const LibraryTitle = ({ record }) => { const translate = useTranslate() const resourceName = translate('resources.library.name', { smart_count: 1 }) @@ -125,132 +131,40 @@ const LibraryEdit = (props) => { {translate('resources.library.sections.statistics')} </Typography> - <Box display="flex"> - <Box flex={1} mr="0.5em"> - <TextInput - InputProps={{ readOnly: true }} - resource={'library'} - source={'totalSongs'} - label={translate('resources.library.fields.totalSongs')} - fullWidth - variant="outlined" - /> - </Box> - <Box flex={1} ml="0.5em"> - <TextInput - InputProps={{ readOnly: true }} - resource={'library'} - source={'totalAlbums'} - label={translate( - 'resources.library.fields.totalAlbums', - )} - fullWidth - variant="outlined" - /> - </Box> - </Box> - - <Box display="flex"> - <Box flex={1} mr="0.5em"> - <TextInput - InputProps={{ readOnly: true }} - resource={'library'} - source={'totalArtists'} - label={translate( - 'resources.library.fields.totalArtists', - )} - fullWidth - variant="outlined" - /> - </Box> - <Box flex={1} ml="0.5em"> - <TextInput - InputProps={{ readOnly: true }} - resource={'library'} - source={'totalSize'} - label={translate('resources.library.fields.totalSize')} - format={(v) => formatBytes(v, 2)} - fullWidth - variant="outlined" - /> - </Box> - </Box> - - <Box display="flex"> - <Box flex={1} mr="0.5em"> - <TextInput - InputProps={{ readOnly: true }} - resource={'library'} - source={'totalDuration'} - label={translate( - 'resources.library.fields.totalDuration', - )} - format={formatDuration2} - fullWidth - variant="outlined" - /> - </Box> - <Box flex={1} ml="0.5em"> - <TextInput - InputProps={{ readOnly: true }} - resource={'library'} - source={'totalMissingFiles'} - label={translate( - 'resources.library.fields.totalMissingFiles', - )} - fullWidth - variant="outlined" - /> - </Box> - </Box> - - {/* Timestamps Section */} - <Box mb="1em"> - <Typography - variant="body2" - color="textSecondary" - gutterBottom - > - {translate('resources.library.fields.lastScanAt')} - </Typography> - <DateField - variant="body1" - source="lastScanAt" - showTime - record={formProps.record} + <Box + display="grid" + gridTemplateColumns="1fr 1fr" + gridColumnGap="1em" + > + <ReadOnlyNumberField + source="totalSongs" + {...readOnlyProps} /> - </Box> - - <Box mb="1em"> - <Typography - variant="body2" - color="textSecondary" - gutterBottom - > - {translate('resources.library.fields.updatedAt')} - </Typography> - <DateField - variant="body1" - source="updatedAt" - showTime - record={formProps.record} + <ReadOnlyNumberField + source="totalAlbums" + {...readOnlyProps} /> - </Box> - - <Box mb="2em"> - <Typography - variant="body2" - color="textSecondary" - gutterBottom - > - {translate('resources.library.fields.createdAt')} - </Typography> - <DateField - variant="body1" - source="createdAt" - showTime - record={formProps.record} + <ReadOnlyNumberField + source="totalArtists" + {...readOnlyProps} /> + <ReadOnlySizeField source="totalSize" {...readOnlyProps} /> + <ReadOnlyDurationField + source="totalDuration" + {...readOnlyProps} + /> + <ReadOnlyNumberField + source="totalMissingFiles" + {...readOnlyProps} + /> + <Box gridColumn="1 / -1"> + <ReadOnlyDateField + source="lastScanAt" + {...readOnlyProps} + /> + </Box> + <ReadOnlyDateField source="updatedAt" {...readOnlyProps} /> + <ReadOnlyDateField source="createdAt" {...readOnlyProps} /> </Box> </Box> </Box> diff --git a/ui/src/player/PlayerEdit.jsx b/ui/src/player/PlayerEdit.jsx index 08f389698..d785eb04e 100644 --- a/ui/src/player/PlayerEdit.jsx +++ b/ui/src/player/PlayerEdit.jsx @@ -1,5 +1,4 @@ import { - TextField, Edit, SimpleForm, useTranslate, @@ -9,7 +8,7 @@ import { Toolbar, } from 'react-admin' import { makeStyles } from '@material-ui/core/styles' -import { Title } from '../common' +import { ReadOnlyTextField, Title } from '../common' import ApiKeyInput from './ApiKeyInput' import { playerInputs } from './playerInputs' @@ -45,8 +44,8 @@ const PlayerEdit = (props) => ( <Edit title={<PlayerTitle />} mutationMode="pessimistic" {...props}> <SimpleForm variant={'outlined'} toolbar={<PlayerEditToolbar />}> {playerInputs()} - <TextField source="client" /> - <TextField source="userName" /> + <ReadOnlyTextField source="client" /> + <ReadOnlyTextField source="userName" /> <ApiKeyInput source="apiKey" /> </SimpleForm> </Edit> diff --git a/ui/src/playlist/PlaylistEdit.jsx b/ui/src/playlist/PlaylistEdit.jsx index f6882e366..c7594fa13 100644 --- a/ui/src/playlist/PlaylistEdit.jsx +++ b/ui/src/playlist/PlaylistEdit.jsx @@ -3,7 +3,6 @@ import { FormDataConsumer, SimpleForm, TextInput, - TextField, BooleanInput, required, useTranslate, @@ -11,13 +10,14 @@ import { ReferenceInput, SelectInput, } from 'react-admin' -import { isWritable, Title } from '../common' +import { isWritable, ReadOnlyTextField, Title } from '../common' const SyncFragment = ({ formData, variant, ...rest }) => { + if (!formData.path) return null return ( <> - {formData.path && <BooleanInput source="sync" {...rest} />} - {formData.path && <TextField source="path" {...rest} />} + <BooleanInput source="sync" {...rest} /> + <ReadOnlyTextField source="path" {...rest} /> </> ) } @@ -56,10 +56,10 @@ const PlaylistEditForm = (props) => { /> </ReferenceInput> ) : ( - <TextField source="ownerName" /> + <ReadOnlyTextField source="ownerName" /> )} <BooleanInput source="public" disabled={!isWritable(record.ownerId)} /> - <FormDataConsumer> + <FormDataConsumer fullWidth> {(formDataProps) => <SyncFragment {...formDataProps} />} </FormDataConsumer> </SimpleForm> diff --git a/ui/src/radio/RadioEdit.jsx b/ui/src/radio/RadioEdit.jsx index bbe001e6f..af879deaa 100644 --- a/ui/src/radio/RadioEdit.jsx +++ b/ui/src/radio/RadioEdit.jsx @@ -1,5 +1,4 @@ import { - DateField, Edit, required, SimpleForm, @@ -9,7 +8,12 @@ import { import { CardMedia } from '@material-ui/core' import { makeStyles } from '@material-ui/core/styles' import { urlValidate } from '../utils/validations' -import { Title, ImageUploadOverlay, useImageLoadingState } from '../common' +import { + Title, + ImageUploadOverlay, + ReadOnlyDateField, + useImageLoadingState, +} from '../common' import subsonic from '../subsonic' import config from '../config' import { RADIO_PLACEHOLDER_IMAGE } from '../consts' @@ -65,8 +69,8 @@ const RadioEdit = (props) => { fullWidth validate={[urlValidate]} /> - <DateField variant="body1" source="updatedAt" showTime /> - <DateField variant="body1" source="createdAt" showTime /> + <ReadOnlyDateField source="updatedAt" /> + <ReadOnlyDateField source="createdAt" /> </SimpleForm> </Edit> ) diff --git a/ui/src/share/ShareEdit.jsx b/ui/src/share/ShareEdit.jsx index 2cf7f2df7..a222d3369 100644 --- a/ui/src/share/ShareEdit.jsx +++ b/ui/src/share/ShareEdit.jsx @@ -2,13 +2,16 @@ import { DateTimeInput, BooleanInput, Edit, - NumberField, SimpleForm, TextInput, } from 'react-admin' import { sharePlayerUrl } from '../utils' import { Link } from '@material-ui/core' -import { DateField } from '../common' +import { + ReadOnlyDateField, + ReadOnlyNumberField, + ReadOnlyTextField, +} from '../common' import config from '../config' export const ShareEdit = (props) => { @@ -16,20 +19,26 @@ export const ShareEdit = (props) => { const url = sharePlayerUrl(id) return ( <Edit {...props}> - <SimpleForm {...rest}> - <Link source="URL" href={url} target="_blank" rel="noopener noreferrer"> + <SimpleForm variant={'outlined'} {...rest}> + <Link + source="URL" + href={url} + target="_blank" + rel="noopener noreferrer" + variant="inherit" + > {url} </Link> <TextInput source="description" /> {config.enableDownloads && <BooleanInput source="downloadable" />} <DateTimeInput source="expiresAt" /> - <TextInput source="contents" disabled /> - <TextInput source="format" disabled /> - <TextInput source="maxBitRate" disabled /> - <TextInput source="username" disabled /> - <NumberField source="visitCount" disabled /> - <DateField source="lastVisitedAt" disabled showTime /> - <DateField source="createdAt" disabled showTime /> + <ReadOnlyTextField source="contents" /> + <ReadOnlyTextField source="format" /> + <ReadOnlyTextField source="maxBitRate" /> + <ReadOnlyTextField source="username" /> + <ReadOnlyNumberField source="visitCount" /> + <ReadOnlyDateField source="lastVisitedAt" /> + <ReadOnlyDateField source="createdAt" /> </SimpleForm> </Edit> ) diff --git a/ui/src/user/UserEdit.jsx b/ui/src/user/UserEdit.jsx index c5d9c75a4..feadafff1 100644 --- a/ui/src/user/UserEdit.jsx +++ b/ui/src/user/UserEdit.jsx @@ -3,7 +3,6 @@ import { makeStyles } from '@material-ui/core/styles' import { TextInput, BooleanInput, - DateField, PasswordInput, Edit, required, @@ -21,7 +20,7 @@ import { useRecordContext, } from 'react-admin' import { Typography } from '@material-ui/core' -import { Title } from '../common' +import { ReadOnlyDateField, Title } from '../common' import DeleteUserButton from './DeleteUserButton' import { LibrarySelectionField } from './LibrarySelectionField.jsx' import { validateUserForm } from './userValidation' @@ -183,10 +182,10 @@ const UserEdit = (props) => { helperText={translate('resources.user.helperTexts.scrobbleFilter')} /> - <DateField variant="body1" source="lastLoginAt" showTime /> - <DateField variant="body1" source="lastAccessAt" showTime /> - <DateField variant="body1" source="updatedAt" showTime /> - <DateField variant="body1" source="createdAt" showTime /> + <ReadOnlyDateField source="lastLoginAt" /> + <ReadOnlyDateField source="lastAccessAt" /> + <ReadOnlyDateField source="updatedAt" /> + <ReadOnlyDateField source="createdAt" /> </SimpleForm> </Edit> ) diff --git a/ui/src/user/UserEdit.test.jsx b/ui/src/user/UserEdit.test.jsx index 74405cc13..837b25ad6 100644 --- a/ui/src/user/UserEdit.test.jsx +++ b/ui/src/user/UserEdit.test.jsx @@ -51,9 +51,6 @@ vi.mock('react-admin', () => ({ BooleanInput: ({ source }) => ( <input type="checkbox" data-testid={`boolean-input-${source}`} /> ), - DateField: ({ source }) => ( - <div data-testid={`date-field-${source}`}>Date</div> - ), PasswordInput: ({ source }) => ( <input type="password" data-testid={`password-input-${source}`} /> ), @@ -82,6 +79,9 @@ vi.mock('./DeleteUserButton', () => ({ vi.mock('../common', () => ({ Title: ({ subTitle }) => <div data-testid="title">{subTitle}</div>, + ReadOnlyDateField: ({ source }) => ( + <div data-testid={`date-field-${source}`}>Date</div> + ), })) // Mock Material-UI From a62d1629f0ff56208460f17aa407b44a41569f0c Mon Sep 17 00:00:00 2001 From: Matt Van Horn <mvanhorn@users.noreply.github.com> Date: Mon, 28 Sep 2026 15:06:44 -0700 Subject: [PATCH 23/26] fix(ui): honor EnableCoverAnimation for theme cover animations (#6234) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: honor cover animation setting in Squiddies Glass Fixes #5170 * fix(ui): move cover animation check into AlbumDetails Apply a noCoverAnimation class from AlbumDetails when enableCoverAnimation is off, so every theme gets the fix. Drop the Squiddies Glass theme changes and its test, and cover the class in AlbumDetails.test.jsx. --------- Co-authored-by: Matt Van Horn <455140+mvanhorn@users.noreply.github.com> Co-authored-by: Deluan Quintão <deluan@navidrome.org> --- ui/src/album/AlbumDetails.jsx | 11 ++++- ui/src/album/AlbumDetails.test.jsx | 70 +++++++++++++++++++++++++++++- 2 files changed, 79 insertions(+), 2 deletions(-) diff --git a/ui/src/album/AlbumDetails.jsx b/ui/src/album/AlbumDetails.jsx index 7fbc479db..8d636a047 100644 --- a/ui/src/album/AlbumDetails.jsx +++ b/ui/src/album/AlbumDetails.jsx @@ -16,6 +16,7 @@ import { useRecordContext, useTranslate, } from 'react-admin' +import clsx from 'clsx' import Lightbox from 'react-image-lightbox' import config from '../config' import 'react-image-lightbox/style.css' @@ -79,6 +80,9 @@ const useStyles = makeStyles( alignItems: 'center', justifyContent: 'center', }, + noCoverAnimation: { + '&, &::before, &::after': { animation: 'none' }, + }, cover: { objectFit: 'contain', cursor: 'pointer', @@ -250,7 +254,12 @@ const AlbumDetails = (props) => { return ( <Card className={classes.root}> <div className={classes.cardContents}> - <div className={classes.coverParent}> + <div + className={clsx( + classes.coverParent, + !config.enableCoverAnimation && classes.noCoverAnimation, + )} + > <Artwork record={record} fit="contain" diff --git a/ui/src/album/AlbumDetails.test.jsx b/ui/src/album/AlbumDetails.test.jsx index 484045444..afc4bf0a0 100644 --- a/ui/src/album/AlbumDetails.test.jsx +++ b/ui/src/album/AlbumDetails.test.jsx @@ -3,7 +3,29 @@ import { describe, test, expect, beforeEach, afterEach } from 'vitest' import { render } from '@testing-library/react' import { RecordContextProvider } from 'react-admin' import { useMediaQuery } from '@material-ui/core' -import { Details } from './AlbumDetails' +import { createTheme, ThemeProvider } from '@material-ui/core/styles' +import config from '../config' +import AlbumDetails, { Details } from './AlbumDetails' + +vi.mock('../subsonic', () => ({ + default: { + getAlbumInfo: () => + Promise.resolve({ + json: { 'subsonic-response': { status: 'ok', albumInfo: {} } }, + }), + getCoverArtUrl: () => '', + }, +})) + +vi.mock('react-admin', async () => { + const actual = await vi.importActual('react-admin') + return { + ...actual, + useDataProvider: () => ({ getOne: vi.fn() }), + useNotify: () => vi.fn(), + useRefresh: () => vi.fn(), + } +}) // Mock useMediaQuery vi.mock('@material-ui/core', async () => { @@ -343,3 +365,49 @@ describe('Details component', () => { }) }) }) + +describe('AlbumDetails cover animation', () => { + const albumRecord = { + id: '123', + name: 'Test Album', + songCount: 12, + duration: 3600, + size: 102400, + } + const originalEnableCoverAnimation = config.enableCoverAnimation + + beforeEach(() => { + vi.mocked(useMediaQuery).mockReturnValue(false) + }) + + afterEach(() => { + config.enableCoverAnimation = originalEnableCoverAnimation + }) + + const renderAlbum = () => + render( + <ThemeProvider theme={createTheme()}> + <RecordContextProvider value={albumRecord}> + <AlbumDetails /> + </RecordContextProvider> + </ThemeProvider>, + ) + + test('applies noCoverAnimation when enableCoverAnimation is false', () => { + config.enableCoverAnimation = false + const { container } = renderAlbum() + const cover = container.querySelector('[class*="coverParent"]') + + expect(cover).not.toBeNull() + expect(cover.className).toMatch(/noCoverAnimation/) + }) + + test('omits noCoverAnimation when enableCoverAnimation is true', () => { + config.enableCoverAnimation = true + const { container } = renderAlbum() + const cover = container.querySelector('[class*="coverParent"]') + + expect(cover).not.toBeNull() + expect(cover.className).not.toMatch(/noCoverAnimation/) + }) +}) From 3a31f702b529fe3d9843894e024d3a0316c11650 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20Dav=C3=B3?= <david@ddavo.me> Date: Tue, 29 Sep 2026 04:50:13 +0200 Subject: [PATCH 24/26] feat(ui): add played filter to album list (#6207) --- persistence/album_repository.go | 1 + persistence/album_repository_test.go | 74 ++++++++++++++++++++++++++++ persistence/sql_annotations_test.go | 2 + ui/src/album/AlbumList.jsx | 1 + ui/src/i18n/en.json | 3 +- 5 files changed, 80 insertions(+), 1 deletion(-) diff --git a/persistence/album_repository.go b/persistence/album_repository.go index 8e7662b8f..6a940a320 100644 --- a/persistence/album_repository.go +++ b/persistence/album_repository.go @@ -131,6 +131,7 @@ var albumFilters = sync.OnceValue(func() map[string]filterFunc { "recently_played": recentlyPlayedFilter, "starred": annotationBoolFilter("starred"), "has_rating": annotationBoolFilter("rating"), + "played": annotationBoolFilter("play_count"), "missing": booleanFilter, "genre_id": genreFilter(AlbumGenres), "role_total_id": allRolesFilter, diff --git a/persistence/album_repository_test.go b/persistence/album_repository_test.go index 7652e703d..0235805c0 100644 --- a/persistence/album_repository_test.go +++ b/persistence/album_repository_test.go @@ -417,6 +417,80 @@ var _ = Describe("AlbumRepository", func() { } }) }) + + Describe("played", func() { + var playedAlbum model.Album + + BeforeEach(func() { + playedAlbum = model.Album{ID: "played-album", Name: "Played Album", LibraryID: 1, SongCount: 1} + Expect(albumRepo.Put(ctx, &playedAlbum)).To(Succeed()) + Expect(albumRepo.IncPlayCount(ctx, playedAlbum.ID, time.Now())).To(Succeed()) + }) + + AfterEach(func() { + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("annotation").Where(squirrel.Eq{"item_id": playedAlbum.ID})) + _, _ = albumRepo.executeSQL(ctx, squirrel.Delete("album").Where(squirrel.Eq{"id": playedAlbum.ID})) + }) + + It("false includes items without annotations", func() { + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ + Filters: map[string]any{"played": "false"}, + }) + Expect(err).ToNot(HaveOccurred()) + albums := res + + var found bool + for _, a := range albums { + if a.ID == albumWithoutAnnotation.ID { + found = true + break + } + } + Expect(found).To(BeTrue(), "Album without annotation should be included in played=false filter") + }) + + It("true excludes items without annotations", func() { + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ + Filters: map[string]any{"played": "true"}, + }) + Expect(err).ToNot(HaveOccurred()) + albums := res + + for _, a := range albums { + Expect(a.ID).ToNot(Equal(albumWithoutAnnotation.ID)) + } + }) + + It("true includes items with play count", func() { + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ + Filters: map[string]any{"played": "true"}, + }) + Expect(err).ToNot(HaveOccurred()) + albums := res + + var found bool + for _, a := range albums { + if a.ID == playedAlbum.ID { + found = true + Expect(a.PlayCount).To(BeNumerically(">", 0)) + break + } + } + Expect(found).To(BeTrue(), "Album with play count should be included in played=true filter") + }) + + It("false excludes items with play count", func() { + res, err := albumRepo.ReadAll(ctx, rest.QueryOptions{ + Filters: map[string]any{"played": "false"}, + }) + Expect(err).ToNot(HaveOccurred()) + albums := res + + for _, a := range albums { + Expect(a.ID).ToNot(Equal(playedAlbum.ID)) + } + }) + }) }) Describe("Album.PlayCount", func() { diff --git a/persistence/sql_annotations_test.go b/persistence/sql_annotations_test.go index ec28c610a..6761331d9 100644 --- a/persistence/sql_annotations_test.go +++ b/persistence/sql_annotations_test.go @@ -91,6 +91,8 @@ var _ = Describe("Annotation Filters", func() { Entry("starred=false", "starred", "false", "COALESCE(starred, 0) = 0", []any(nil)), Entry("starred=True (case insensitive)", "starred", "True", "COALESCE(starred, 0) > 0", []any(nil)), Entry("rating=true", "rating", "true", "COALESCE(rating, 0) > 0", []any(nil)), + Entry("play_count=true", "play_count", "true", "COALESCE(play_count, 0) > 0", []any(nil)), + Entry("play_count=false", "play_count", "false", "COALESCE(play_count, 0) = 0", []any(nil)), ) It("returns nil if value is not a string", func() { diff --git a/ui/src/album/AlbumList.jsx b/ui/src/album/AlbumList.jsx index 3cb8a648e..81e4dcd14 100644 --- a/ui/src/album/AlbumList.jsx +++ b/ui/src/album/AlbumList.jsx @@ -163,6 +163,7 @@ const AlbumFilter = (props) => { /> </ReferenceInput> <NullableBooleanInput source="compilation" /> + <NullableBooleanInput source="played" defaultValue={false} /> <NumberInput source="year" /> {config.enableFavourites && ( <NullableBooleanInput diff --git a/ui/src/i18n/en.json b/ui/src/i18n/en.json index 851b47ee3..a04b6e311 100644 --- a/ui/src/i18n/en.json +++ b/ui/src/i18n/en.json @@ -83,7 +83,8 @@ "grouping": "Grouping", "media": "Media", "mood": "Mood", - "missing": "Missing" + "missing": "Missing", + "played": "Played" }, "actions": { "playAll": "Play", From e03ce871773c434b95ec2a623243c6397df160f5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= <deluan@navidrome.org> Date: Tue, 29 Sep 2026 11:21:14 -0400 Subject: [PATCH 25/26] fix(ui): make Undo and AMusic Save buttons readable (#6249) --- ui/src/layout/Notification.jsx | 25 ++++++++++++++++++++----- ui/src/themes/amusic.js | 15 +++++++++++++++ ui/src/themes/nuclear.js | 7 +++++++ 3 files changed, 42 insertions(+), 5 deletions(-) diff --git a/ui/src/layout/Notification.jsx b/ui/src/layout/Notification.jsx index 001d3fc01..ff641ec22 100644 --- a/ui/src/layout/Notification.jsx +++ b/ui/src/layout/Notification.jsx @@ -1,11 +1,26 @@ import React from 'react' import { Notification as RANotification } from 'react-admin' +import { makeStyles } from '@material-ui/core/styles' -const Notification = (props) => ( - <RANotification - {...props} - anchorOrigin={{ vertical: 'top', horizontal: 'center' }} - /> +// RA's primary.light Undo is unreadable on the light snackbar of dark themes +const useStyles = makeStyles( + { + undo: { + color: 'inherit', + }, + }, + { name: 'NDNotification' }, ) +const Notification = (props) => { + const classes = useStyles() + return ( + <RANotification + {...props} + classes={{ undo: classes.undo }} + anchorOrigin={{ vertical: 'top', horizontal: 'center' }} + /> + ) +} + export default Notification diff --git a/ui/src/themes/amusic.js b/ui/src/themes/amusic.js index 55205baf3..eef7f90d6 100644 --- a/ui/src/themes/amusic.js +++ b/ui/src/themes/amusic.js @@ -79,6 +79,16 @@ export default { color: '#eee', backgroundColor: '#ff4e6b', }, + containedPrimary: { + color: '#fff', + backgroundColor: '#D60017', + '&:hover': { + backgroundColor: '#a30011', + '@media (hover: none)': { + backgroundColor: '#D60017', + }, + }, + }, textSizeSmall: { fontSize: '0.8rem', paddingRight: '0.5rem', @@ -192,6 +202,11 @@ export default { paddingBottom: '1rem', }, }, + NDNotification: { + undo: { + color: '#fff', + }, + }, RaConfirm: { confirmPrimary: { color: '#fff', diff --git a/ui/src/themes/nuclear.js b/ui/src/themes/nuclear.js index b34896c97..bf804139e 100644 --- a/ui/src/themes/nuclear.js +++ b/ui/src/themes/nuclear.js @@ -86,6 +86,13 @@ export default { }, }, }, + NDNotification: { + undo: { + '& .MuiButton-label': { + color: 'inherit', + }, + }, + }, MuiChip: { root: { backgroundColor: nukeCol['accent'], From 0e1893530b844898cbf87825fdefd4b13e1ff3cc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= <deluan@navidrome.org> Date: Tue, 29 Sep 2026 13:23:36 -0400 Subject: [PATCH 26/26] fix(ui): don't crash the playlist list when rows lose their record (#6250) * fix(ui): don't crash playlist list rows that lost their record react-admin 3 evicts records fetched more than 10 minutes ago whenever another getList for the same resource completes, but the list keeps its cached ids. The Datagrid then renders those rows with an undefined record, and the Public and Auto-import switches crashed reading record.id. This happened when the playlist list was left open and the sidebar or the add to playlist dialog reloaded a smaller set of playlists. Both switches now render nothing when the row has no record; the next list refresh fills the row in again. * refactor(ui): merge playlist list toggles into one ToggleField The Public and Auto-import switches were copies that differed only in the field they flip. ToggleField now flips its source field, and ToggleAutoImport just shows it for playlists that have a file path. The tests render inside TestContext, so they use react-admin's real hooks instead of mocks. --- ui/src/playlist/PlaylistList.jsx | 46 ++++++--------------------- ui/src/playlist/PlaylistList.test.jsx | 25 ++++++++++++++- 2 files changed, 34 insertions(+), 37 deletions(-) diff --git a/ui/src/playlist/PlaylistList.jsx b/ui/src/playlist/PlaylistList.jsx index d2b17b108..14d819a4e 100644 --- a/ui/src/playlist/PlaylistList.jsx +++ b/ui/src/playlist/PlaylistList.jsx @@ -67,15 +67,15 @@ const PlaylistFilter = (props) => { ) } -const TogglePublicInput = ({ resource, source }) => { +export const ToggleField = ({ resource, source }) => { const record = useRecordContext() const notify = useNotify() - const [togglePublic] = useUpdate( + const [toggle] = useUpdate( resource, - record.id, + record?.id, { ...record, - public: !record.public, + [source]: !record?.[source], }, { undoable: false, @@ -86,10 +86,12 @@ const TogglePublicInput = ({ resource, source }) => { ) const handleClick = (e) => { - togglePublic() + toggle() e.stopPropagation() } + if (!record) return null + return ( <Switch checked={record[source]} @@ -99,35 +101,9 @@ const TogglePublicInput = ({ resource, source }) => { ) } -const ToggleAutoImport = ({ resource, source }) => { +export const ToggleAutoImport = (props) => { const record = useRecordContext() - const notify = useNotify() - const [ToggleAutoImport] = useUpdate( - resource, - record.id, - { - ...record, - sync: !record.sync, - }, - { - undoable: false, - onFailure: (error) => { - notify('ra.page.error', 'warning') - }, - }, - ) - const handleClick = (e) => { - ToggleAutoImport() - e.stopPropagation() - } - - return record.path ? ( - <Switch - checked={record[source]} - onClick={handleClick} - disabled={!isWritable(record.ownerId)} - /> - ) : null + return record?.path ? <ToggleField {...props} /> : null } const PlaylistListBulkActions = (props) => { @@ -169,9 +145,7 @@ const PlaylistList = (props) => { updatedAt: isDesktop && ( <DateField source="updatedAt" sortByOrder={'DESC'} /> ), - public: !isXsmall && ( - <TogglePublicInput source="public" sortByOrder={'DESC'} /> - ), + public: !isXsmall && <ToggleField source="public" sortByOrder={'DESC'} />, comment: <TextField source="comment" />, sync: !isXsmall && ( <ToggleAutoImport source="sync" sortByOrder={'DESC'} /> diff --git a/ui/src/playlist/PlaylistList.test.jsx b/ui/src/playlist/PlaylistList.test.jsx index 4fbc6d516..6c714b827 100644 --- a/ui/src/playlist/PlaylistList.test.jsx +++ b/ui/src/playlist/PlaylistList.test.jsx @@ -1,7 +1,8 @@ import React from 'react' import { render, screen } from '@testing-library/react' import { describe, it, expect, vi } from 'vitest' -import { PlaylistLove } from './PlaylistList' +import { TestContext } from 'ra-test' +import { PlaylistLove, ToggleField, ToggleAutoImport } from './PlaylistList' vi.mock('../config', () => ({ default: { enableFavourites: true }, @@ -32,3 +33,25 @@ describe('<PlaylistLove />', () => { }) }) }) + +// react-admin evicts records older than 10 minutes while the list still holds +// their ids, so rows can render with no record. +describe('playlist toggles without a record', () => { + it('<ToggleField /> renders nothing', () => { + const { container } = render( + <TestContext> + <ToggleField resource="playlist" source="public" /> + </TestContext>, + ) + expect(container.innerHTML).toBe('') + }) + + it('<ToggleAutoImport /> renders nothing', () => { + const { container } = render( + <TestContext> + <ToggleAutoImport resource="playlist" source="sync" /> + </TestContext>, + ) + expect(container.innerHTML).toBe('') + }) +})